essos.augmented_lagrangian

ALM (Augmented Lagrangian Method) using JAX and optimizers from OPTAX/JAXOPT/OPTIMISTIX inspired by mdmm_jax github repository

Classes

LagrangeMultiplier

A class containing constrain parameters for Augmented Lagrangian Method

BaseConstraint

A minimal mutable container holding init and loss callables for a constraint.

CompositeConstraint

Mutable composite constraint container.

SelectiveConstraint

Wraps a constraint with selective named dependencies, similar to custom_loss.

ScaledConstraint

Wraps a constraint with automatic or user-specified scaling.

ALM

Functions

_multiplier_like(out, multiplier, penalty, omega, eta, ...)

eq(fun[, model_lagrangian, multiplier, penalty, ...])

Represents an equality constraint, g(x) = 0.

ineq(fun[, model_lagrangian, multiplier, penalty, ...])

Represents an inequality constraint, h(x) >= 0, which uses a slack

combine(*args)

Combines constraints with selective named dependencies, mirroring losses.py.

total_infeasibility(tree)

norm_constraints(tree)

infty_norm_constraints(tree)

penalty_average(tree)

apply_mu_tolerance_per_constraint(constraint_dict, ...)

Apply Mu update rule for vectorized constraints.

apply_mu_tolerance_all_constraints(constraint_dicts_map)

Apply Mu_Tolerance across all constraints (JAX-compatible, no Python loops or dict indexing).

ALM_model_jaxopt_lbfgsb(constraints[, loss, ...])

Module Contents

class essos.augmented_lagrangian.LagrangeMultiplier

Bases: NamedTuple

A 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.

class essos.augmented_lagrangian.ALM

Bases: NamedTuple

init: Callable
update: Callable
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)