essos.augmented_lagrangian¶
ALM (Augmented Lagrangian Method) using JAX and optimizers from OPTAX/JAXOPT/OPTIMISTIX inspired by mdmm_jax github repository
Classes¶
A class containing constrain parameters for Augmented Lagrangian Method |
|
A minimal mutable container holding init and loss callables for a constraint. |
|
Mutable composite constraint container. |
|
Wraps a constraint with selective named dependencies, similar to custom_loss. |
|
Wraps a constraint with automatic or user-specified scaling. |
|
Functions¶
|
|
|
Represents an equality constraint, g(x) = 0. |
|
Represents an inequality constraint, h(x) >= 0, which uses a slack |
|
Combines constraints with selective named dependencies, mirroring losses.py. |
|
|
|
|
|
|
|
|
|
Apply Mu update rule for vectorized constraints. |
|
Apply Mu_Tolerance across all constraints (JAX-compatible, no Python loops or dict indexing). |
|
Module Contents¶
- class essos.augmented_lagrangian.LagrangeMultiplier¶
Bases:
NamedTupleA class containing constrain parameters for Augmented Lagrangian Method
- value: Any¶
- penalty: Any¶
- omega: Any¶
- eta: Any¶
- sq_grad: Any¶
- essos.augmented_lagrangian._multiplier_like(out, multiplier, penalty, omega, eta, sq_grad)¶
- class essos.augmented_lagrangian.BaseConstraint(init: Callable, loss: Callable)¶
A minimal mutable container holding init and loss callables for a constraint.
This mirrors the simple tuple-like behavior used elsewhere but allows attribute access and matches the base_loss style in losses.py.
- init¶
- loss¶
- class essos.augmented_lagrangian.CompositeConstraint(init_fn: Callable, loss_fn: Callable, selective_map=None, arg_names=None)¶
Mutable composite constraint container.
Exposes init and loss callables (same as Constraint) while allowing attaching metadata like arg_names, _dependencies, and selective_map. set_dependencies will propagate dependencies to any contained SelectiveConstraint instances.
- init¶
- loss¶
- selective_map¶
- arg_names¶
- _dependencies¶
- _starting_dofs = None¶
- _dofs_to_pytree = None¶
- clear_cache()¶
- property dependencies¶
- set_dependencies(deps)¶
- property starting_dofs¶
- property dofs_to_pytree¶
- class essos.augmented_lagrangian.SelectiveConstraint(constraint: BaseConstraint, *arg_names, **kwargs)¶
Wraps a constraint with selective named dependencies, similar to custom_loss.
This allows constraints to only depend on a subset of the available degrees of freedom by name, enabling combination of constraints with different argument requirements. No need to specify indices - just the dependency names!
Why filtering is necessary: - Different constraints may require different subsets of arguments (e.g., one depends on
‘field’ only, another on ‘coil’ only, a third on both).
Named filtering ensures each constraint receives only its required arguments, avoiding: * Unnecessary computations with unused data * Constraints expecting different signatures from breaking * Wasteful memory transfer of irrelevant arrays
Enables flexible constraint composition where DOF dependencies vary.
- constraint¶
The underlying Constraint (init_fn, loss_fn) tuple
- arg_names¶
Tuple of argument names that this constraint depends on
- dependencies¶
Dictionary mapping dependency names to their arrays/objects
Example
# Create constraints on different DOF subsets - no indices needed! field_constraint = alm.eq(lambda field: jnp.sum(field**2)) surface_constraint = alm.eq(lambda surface: jnp.sum(surface**2))
selective1 = SelectiveConstraint(field_constraint, ‘field’) selective2 = SelectiveConstraint(surface_constraint, ‘surface’)
combined = alm.combine(selective1, selective2) # Set dependencies by name combined.dependencies = {‘field’: field_array, ‘surface’: surface_array}
- constraint¶
- arg_names = ()¶
- kwargs¶
- _dependencies¶
- _starting_dofs = None¶
- _dofs_to_pytree = None¶
- clear_cache()¶
- property dependencies¶
Get the dependencies dictionary.
- _get_filtered_args()¶
Extract only the required arguments from dependencies by name.
- property starting_dofs¶
- property dofs_to_pytree¶
- init(*args, **kwargs)¶
Initialize constraint parameters using current dependencies.
- loss(params, *args, **kwargs)¶
Compute loss using current dependencies.
- Parameters:
params – Constraint parameters from init()
- Returns:
(loss_value, constraint_info) tuple
- class essos.augmented_lagrangian.ScaledConstraint(constraint: BaseConstraint, scale_factor=1.0, auto_scale=True, elementwise_scale=False)¶
Wraps a constraint with automatic or user-specified scaling.
Useful when combining constraints with vastly different magnitudes. Automatically normalizes constraint output by its initial norm or element-wise absolute values.
- constraint¶
The underlying Constraint (init_fn, loss_fn) tuple
- scale_factor¶
Scalar to multiply constraint output (auto-computed or user-specified)
- auto_scale¶
Whether to auto-compute scale_factor from initial constraint norm
- elementwise_scale¶
If True, scale by 1/abs(constraint) element-wise; if False, scale by 1/norm
Example
# Auto-scale based on initial constraint norm (default) constraint = alm.eq(lambda x: jnp.sum(x**2)) scaled = ScaledConstraint(constraint, auto_scale=True, elementwise_scale=False)
# Auto-scale element-wise scaled = ScaledConstraint(constraint, auto_scale=True, elementwise_scale=True)
# Or use fixed scaling scaled = ScaledConstraint(constraint, scale_factor=1.0, auto_scale=False)
- constraint¶
- scale_factor = 1.0¶
- auto_scale = True¶
- elementwise_scale = False¶
- _initial_scale = None¶
- _initial_scale_scalar = None¶
- init(*args, **kwargs)¶
Initialize constraint parameters and compute scale factor if needed.
- loss(params, *args, **kwargs)¶
Compute scaled constraint loss.
- Parameters:
params – Constraint parameters from init()
*args – Arguments to constraint
**kwargs – Keyword arguments
- Returns:
(scaled_loss_value, constraint_info) tuple
- essos.augmented_lagrangian.eq(fun, model_lagrangian='Standard', multiplier=0.0, penalty=1.0, omega=1.0, eta=1.0, sq_grad=0.0, weight=1.0, reduction=jnp.sum)¶
Represents an equality constraint, g(x) = 0.
- Parameters:
fun – The constraint function, a differentiable function of your parameters which should output zero when satisfied and smoothly increasingly far from zero values for increasing levels of constraint violation.
damping – Sets the damping (oscillation reduction) strength.
weight – Weights the loss from the constraint relative to the primary loss function’s value.
reduction – The function that is used to aggregate the constraints if the constraint function outputs more than one element.
- Returns:
An (init_fn, loss_fn) constraint tuple for the equality constraint.
- essos.augmented_lagrangian.ineq(fun, model_lagrangian='Standard', multiplier=0.0, penalty=1.0, omega=1.0, eta=1.0, sq_grad=0.0, weight=1.0, reduction=jnp.sum)¶
Represents an inequality constraint, h(x) >= 0, which uses a slack variable internally to convert it to an equality constraint.
- Parameters:
fun – The constraint function, a differentiable function of your parameters which should output greater than or equal to zero when satisfied and smoothly increasingly negative values for increasing levels of constraint violation.
damping – Sets the damping (oscillation reduction) strength.
weight – Weights the loss from the constraint relative to the primary loss function’s value.
reduction – The function that is used to aggregate the constraints if the constraint function outputs more than one element.
- Returns:
An (init_fn, loss_fn) constraint tuple for the inequality constraint.
- essos.augmented_lagrangian.combine(*args)¶
Combines constraints with selective named dependencies, mirroring losses.py.
Each SelectiveConstraint specifies which named arguments it depends on. Regular Constraints still work (they receive all arguments positionally).
The returned combined constraint supports both: - Old style: combined.init(arg1, arg2); combined.loss(params, arg1, arg2) - New style: Set combined.dependencies = {…} and call combined.init/loss()
Implementation optimizes for JIT by pre-wrapping constraint functions at combination time, avoiding dynamic control flow inside the loss/init functions.
- Validation:
All constraints must be Constraint or SelectiveConstraint objects
Constraint outputs must be compatible (all scalars, all same shape, etc.)
- Parameters:
*args – A series of constraint (init_fn, loss_fn) tuples or SelectiveConstraint objects.
- Returns:
A combined Constraint with optional .dependencies and .arg_names attributes.
- Raises:
ValueError – If no constraints provided or constraint types are invalid
TypeError – If constraint objects are not Constraint or SelectiveConstraint
- essos.augmented_lagrangian.total_infeasibility(tree)¶
- essos.augmented_lagrangian.norm_constraints(tree)¶
- essos.augmented_lagrangian.infty_norm_constraints(tree)¶
- essos.augmented_lagrangian.penalty_average(tree)¶
- essos.augmented_lagrangian.apply_mu_tolerance_per_constraint(constraint_dict, grad_dict, constraint_info, constraint_info_prev=None, model_lagrangian='Standard', model_mu='Mu_Adaptative_1', beta=2.0, mu_max=10000.0, alpha=0.99, gamma=0.01, epsilon=1e-08, eta_tol=0.0001, omega_tol=1e-06, decrease_tol=0.75)¶
Apply Mu update rule for vectorized constraints.
Supports three strategies: - Mu_Tolerance: All elements updated together based on norm, uses average penalty - Mu_Adaptative_1: Element-wise updates using eta parameter - Mu_Adaptative_2: Element-wise updates using decrease tolerance criterion
- Parameters:
constraint_dict – dict with ‘lambda’ and optionally ‘slack’ keys containing LagrangeMultiplier objects
grad_dict – dict with gradients matching the constraint_dict structure
constraint_info – current constraint violation information (array or scalar)
constraint_info_prev – previous constraint violation info (needed for Mu_Adaptative_2)
model_lagrangian – ‘Standard’ or ‘Squared’ lagrangian formulation
model_mu – ‘Mu_Tolerance’ (global norm, avg penalty), ‘Mu_Adaptative_1’ (element-wise eta), or ‘Mu_Adaptative_2’ (element-wise decrease)
decrease_tol – tolerance for decrease criterion in Mu_Adaptative_2 (default 0.75 = 25% decrease)
- essos.augmented_lagrangian.apply_mu_tolerance_all_constraints(constraint_dicts_map, grad_dicts_map=None, constraint_infos_map=None, model_lagrangian='Standard', beta=2.0, mu_max=10000.0, alpha=0.99, gamma=0.01, epsilon=1e-08, eta_tol=0.0001, omega_tol=1e-06)¶
Apply Mu_Tolerance across all constraints (JAX-compatible, no Python loops or dict indexing).
This mirrors the per-constraint updater but computes a single global decision using all constraint infos, then applies the same update to every constraint’s Lagrange multipliers. :param constraint_dicts_map: pytree mapping keys -> per-constraint lagrange dicts :param grad_dicts_map: optional pytree of matching shapes with gradient/info dicts :param constraint_infos_map: optional pytree mapping keys -> constraint info pytrees
- Returns:
pytree with updated LagrangeMultiplier dicts matching constraint_dicts_map.
- essos.augmented_lagrangian.ALM_model_jaxopt_lbfgsb(constraints: BaseConstraint, loss=lambda x: ..., model_lagrangian='Standard', model_mu='Mu_Tolerance', beta=2.0, mu_max=10000.0, alpha=0.99, gamma=0.01, epsilon=1e-08, eta_tol=0.0001, omega_tol=1e-06, **kargs)¶