jaxmodel.constraints

jaxmodel.constraints.wrap_vector_constraint(var_spec: VariableSpec, fun: Callable[[Dict[str, Array], Dict[str, Array]], Array])[source]
class jaxmodel.constraints.ConstraintEntry(name: 'str', kind: 'str', fun: 'callable', structure: 'str' = 'nonlinear', jac_y_fun: 'Optional[callable]' = None, metadata: 'Optional[dict[str, Any]]' = None)[source]

Bases: object

name: str
kind: str
fun: callable
structure: str = 'nonlinear'
jac_y_fun: callable | None = None
metadata: dict[str, Any] | None = None
jaxmodel.constraints.eq_scalar(fun, var_spec: VariableSpec, name: str = 'eq_scalar', structure: str = 'nonlinear', jac_y_fun: callable | None = None, metadata: dict[str, Any] | None = None) ConstraintEntry[source]
jaxmodel.constraints.ineq_scalar(fun, var_spec: VariableSpec, name: str = 'ineq_scalar', structure: str = 'nonlinear', jac_y_fun: callable | None = None, metadata: dict[str, Any] | None = None) ConstraintEntry[source]
jaxmodel.constraints.eq_block(fun, var_spec: VariableSpec, name: str = 'eq_block', structure: str = 'nonlinear', jac_y_fun: callable | None = None, metadata: dict[str, Any] | None = None) ConstraintEntry[source]
jaxmodel.constraints.ineq_block(fun, var_spec: VariableSpec, name: str = 'ineq_block', structure: str = 'nonlinear', jac_y_fun: callable | None = None, metadata: dict[str, Any] | None = None) ConstraintEntry[source]
jaxmodel.constraints.quadratic_eq_scalar(var_spec: VariableSpec, Q: Array, c: Array, rhs_const: float | Array = 0.0, x_coeff: Array | None = None, x_name: str | None = None, name: str = 'quadratic_eq') ConstraintEntry[source]
jaxmodel.constraints.quadratic_ineq_scalar(var_spec: VariableSpec, Q: Array, c: Array, rhs_const: float | Array = 0.0, x_coeff: Array | None = None, x_name: str | None = None, name: str = 'quadratic_ineq') ConstraintEntry[source]
jaxmodel.constraints.aggregate_constraints(entries: ~typing.Sequence[~jaxmodel.constraints.ConstraintEntry], dtype=<class 'jax.numpy.float64'>)[source]
jaxmodel.constraints.aggregate_constraint_jacobian(entries: ~typing.Sequence[~jaxmodel.constraints.ConstraintEntry], n_vars: int, dtype=<class 'jax.numpy.float64'>)[source]