Constraints
A constraint describes the space a value lives in, three ways: its physical bounds,
a bijector from the unconstrained real line (raw space) for unconstrained optimizers,
and a base space (base_bounds and base_bijector) for bounded optimizers.
Closedness. Each bound is closed (the value may sit on it exactly) or open (the value
may only approach it), and every constraint reports which as closed, a
(lower, upper) pair. An infinite bound is closed unless asked otherwise, so ±∞ is
inside: RealLine is [−∞, ∞], Positive is (0, ∞], NonNegative is [0, ∞], and
likewise for Negative and NonPositive. Interval, GreaterThan and LessThan
are closed by default and take closed as one bool for both ends or a
(lower, upper) pair, so GreaterThan(1.0, closed=(True, False)) is [1, ∞). An
Interval may have infinite ends, per element, and behaves there like the matching
half-line or real line; Interval(-inf, inf, closed=False) is the open real line.
NaN is outside every constraint.
is_outside respects closedness, and intersect keeps the open side where two
constraints share a bound, finite or infinite. Closedness follows the bounds through a
Transformed bijector, so Transformed(RealLine(), Sigmoid()) is [0, 1]. A
distribution's support is closed at a finite endpoint where its density is positive
and finite, and open at ±∞, so Uniform(a, b) gives [a, b], while LogNormal gives
(0, ∞).
A value at ±∞ has an infinite raw value and no usable gradient, as on any closed bound: it suits a fixed parameter, not a free one's starting point.
Raw space. The bijector maps the real line onto the open interior, so a value exactly on a closed bound has an infinite raw value. Starting a solver from such a value (by nudging it inward, say) is up to the solver.
Base space. The base space is where a bounded optimizer works, and it reaches the
closed bounds exactly. With two finite bounds it is the unit box, mapped affinely onto
the physical bounds, so every parameter presents the same scale whatever its units.
Otherwise it is the physical space itself. For a constraint inferred from a
distribution, the base comes from the support, not the prior: a Normal prior gives
the real line with an identity map, while its raw space is still whitened by the prior.
import jax.numpy as jnp
from parax.constraints import GreaterThan, Interval, NonNegative, Positive, intersect
assert not Interval(0.0, 1.0).is_outside(jnp.array(0.0))
assert Interval(0.0, 1.0, closed=(False, True)).is_outside(jnp.array(0.0))
assert Positive().is_outside(jnp.array(0.0))
assert not NonNegative().is_outside(jnp.array(0.0))
assert intersect(Interval(0.0, 10.0), Positive()).closed == (False, True)
assert not Positive().is_outside(jnp.array(jnp.inf))
assert GreaterThan(1.0, closed=(True, False)).is_outside(jnp.array(jnp.inf))
assert Interval(0.0, jnp.inf).is_outside(jnp.array(jnp.nan))
parax.constraints.AbstractConstraint
Bases: Module
The base class for all physical constraints in Parax.
Constraints are a higher-level concept that provide bounds and bijectors over constrained domains. This is useful for use with unconstrained solvers (which require a bijector from the unconstrained real line to the constrained domain) and bounded solvers (which accept lower and upper bounds directly).
Attributes:
| Name | Type | Description |
|---|---|---|
bounds |
AbstractVar[tuple[PyTree, PyTree]]
|
A tuple containing the physical lower and upper bounds of the constrained space. |
closed |
AbstractVar[tuple[PyTree, PyTree]]
|
A tuple saying whether each of |
bijector |
AbstractVar[AbstractBijector]
|
A |
base_bounds |
AbstractVar[tuple[PyTree, PyTree]]
|
A tuple containing the foundational, un-skewed orthogonal bounds, which a
bounded optimizer works in directly. Where the physical space has two finite bounds
this is the unit box, so that a step of a given size carries the same meaning along
every axis whatever the physical units. For transformed constraints, this isolates
the safe topological box before any dense correlations or skews are applied. It falls
back to |
base_bijector |
AbstractVar[AbstractBijector]
|
A |
clip(value)
Clip a value to lie within this constraint.
Source code in parax/constraints.py
75 76 77 78 79 | |
is_outside(value)
Returns if another value is outside the constraint.
A value on a closed bound is inside; a value on an open bound is outside. NaN is outside every constraint.
Source code in parax/constraints.py
81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 | |
midpoint()
Returns the midpoint of the constraint.
Note that non-finite constraints may return infinity.
Source code in parax/constraints.py
98 99 100 101 102 103 104 | |
parax.constraints.AbstractConstrained
Bases: AbstractBounded[T]
The abstract interface for a constrained PyTree.
Used as a type check for parax.is_constrained.
Implies that the PyTree has associated constraints (and therefore bounds), but does not necessarily enforce that the PyTree follows those constraints.
Attributes:
| Name | Type | Description |
|---|---|---|
constraint |
AbstractVar[AbstractConstraint]
|
Returns the active constraint of the PyTree. |
bounds |
AbstractVar[tuple[T, T]]
|
Returns the current PyTree bounds. Each must have a matching PyTree structure as |
parax.constraints.AbstractConstrainable
Bases: AbstractConstrained[T]
The abstract interface for a constrainable PyTree.
Variables implementing this interface support the dynamic injection and updating of constraints.
Used as a type check for parax.is_constrainable.
constrain(constraint)
abstractmethod
Returns a new instance of the PyTree with the updated constraint, ensuring internal state (like unconstrained raw values) is recalculated if necessary.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
constraint
|
AbstractConstraint
|
The new constraint to apply. |
required |
Returns:
| Type | Description |
|---|---|
Self
|
A new instance of the constrainable PyTree. |
Source code in parax/constraints.py
936 937 938 939 940 941 942 943 944 945 946 947 948 949 | |
parax.constraints.RealLine(shape=())
Bases: AbstractUncorrelatedConstraint
Represents a value that can span the entire real number line.
Effectively a structural no-op constraint using an Identity bijector,
useful for maintaining consistent types in mixed parameter sets. Both infinite
ends are closed, so ±∞ is inside; for an open end, use Interval(-inf, inf, closed=...).
Attributes:
| Name | Type | Description |
|---|---|---|
shape |
Any
|
The expected shape of the unconstrained parameter. |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Any
|
The expected shape of the unconstrained parameter. |
()
|
Source code in parax/constraints.py
157 158 159 160 161 162 | |
parax.constraints.GreaterThan(lower, closed=True)
Bases: AbstractUncorrelatedConstraint
Represents a value greater than, or equal to, a lower bound.
Attributes:
| Name | Type | Description |
|---|---|---|
lower |
ndarray
|
The lower bound array or scalar. |
closed |
tuple[bool, bool]
|
Whether each bound is included, as a |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
lower
|
Union[float, Array]
|
The lower bound. |
required |
closed
|
bool | tuple[bool, bool]
|
Whether the bounds themselves are included: one bool for both
|
True
|
Source code in parax/constraints.py
192 193 194 195 196 197 198 199 200 201 202 203 204 | |
parax.constraints.LessThan(upper, closed=True)
Bases: AbstractUncorrelatedConstraint
Represents a value less than, or equal to, an upper bound.
Attributes:
| Name | Type | Description |
|---|---|---|
upper |
ndarray
|
The upper bound array or scalar. |
closed |
tuple[bool, bool]
|
Whether each bound is included, as a |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
upper
|
Union[float, Array]
|
The upper bound. |
required |
closed
|
bool | tuple[bool, bool]
|
Whether the bounds themselves are included: one bool for both
-∞ and |
True
|
Source code in parax/constraints.py
227 228 229 230 231 232 233 234 235 236 237 238 239 | |
parax.constraints.Interval(lower, upper, closed=True)
Bases: AbstractUncorrelatedConstraint
Represents a value bounded between a lower and upper value.
Either end may be infinite, per element. Where it is, the element behaves like the
matching half-line or real line (bijector and base space), keeping its own closed.
Attributes:
| Name | Type | Description |
|---|---|---|
lower |
ndarray
|
The lower bound. |
upper |
ndarray
|
The upper bound. |
closed |
tuple[bool, bool]
|
Whether |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
lower
|
Union[float, Array]
|
The lower bound. |
required |
upper
|
Union[float, Array]
|
The upper bound. |
required |
closed
|
bool | tuple[bool, bool]
|
Whether the bounds themselves are included: one bool for both,
or a |
True
|
Source code in parax/constraints.py
291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 | |
base_bijector
property
Maps the base space onto the physical interval: affinely from the unit box where both ends are finite, and unchanged elsewhere.
base_bounds
property
The unit box where both ends are finite, and the physical space elsewhere.
Overrides the Identity default from AbstractUncorrelatedConstraint, which
would hand a bounded optimizer the physical box. Normalising to [0, 1]
is generally numerically better during optimization.
parax.constraints.Positive(shape=(), dtype=None)
Bases: GreaterThan
Convenience constraint for values that must be strictly positive: (0, ∞].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Any
|
The shape of the parameter array. |
()
|
dtype
|
Any
|
The JAX data type of the parameter array. |
None
|
Source code in parax/constraints.py
370 371 372 373 374 375 376 | |
parax.constraints.NonNegative(shape=(), dtype=None)
Bases: GreaterThan
Convenience constraint for values that must be non-negative: [0, ∞].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Any
|
The shape of the parameter array. |
()
|
dtype
|
Any
|
The JAX data type of the parameter array. |
None
|
Source code in parax/constraints.py
381 382 383 384 385 386 387 | |
parax.constraints.Negative(shape=(), dtype=None)
Bases: LessThan
Convenience constraint for values that must be strictly negative: [-∞, 0).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Any
|
The shape of the parameter array. |
()
|
dtype
|
Any
|
The JAX data type of the parameter array. |
None
|
Source code in parax/constraints.py
392 393 394 395 396 397 398 | |
parax.constraints.NonPositive(shape=(), dtype=None)
Bases: LessThan
Convenience constraint for values that must be non-positive: [-∞, 0].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Any
|
The shape of the parameter array. |
()
|
dtype
|
Any
|
The JAX data type of the parameter array. |
None
|
Source code in parax/constraints.py
403 404 405 406 407 408 409 | |
parax.constraints.Leafwise(tree)
Bases: AbstractConstraint
Represents a PyTree of constraints mapping over a PyTree of inputs.
Useful for applying heterogeneous constraints to complex nested structures
(like equinox.Module instances) simultaneously.
Attributes:
| Name | Type | Description |
|---|---|---|
tree |
PyTree[AbstractConstraint]
|
The PyTree containing |
Source code in parax/constraints.py
530 531 532 533 534 535 536 537 | |
parax.constraints.Custom(bijector, bounds=(jnp.array(-jnp.inf), jnp.array(jnp.inf)), closed=True, base_bounds=None, base_bijector=None)
Bases: AbstractConstraint
An escape hatch for power users who need a specific distreqx bijector mapping with predefined physical bounds.
Attributes:
| Name | Type | Description |
|---|---|---|
bijector |
AbstractBijector
|
The internal, user-defined distreqx bijector mapping from the unconstrained real line to the physical space. |
bounds |
tuple[PyTree, PyTree]
|
The manually defined physical boundaries |
closed |
tuple[PyTree, PyTree]
|
Whether each of |
base_bounds |
tuple[PyTree, PyTree]
|
The orthogonal base boundaries. Defaults to |
base_bijector |
AbstractBijector
|
The bijector mapping from |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
bijector
|
AbstractBijector
|
The custom |
required |
bounds
|
tuple[PyTree, PyTree]
|
A tuple of |
(array(-inf), array(inf))
|
closed
|
bool | tuple[PyTree, PyTree]
|
Whether the bounds themselves are included: one bool for every
bound, or a |
True
|
base_bounds
|
tuple[PyTree, PyTree] | None
|
Optional. A tuple of |
None
|
base_bijector
|
AbstractBijector | None
|
Optional. The bijector handling spatial skew/correlation.
If None, defaults to |
None
|
Source code in parax/constraints.py
600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 | |
parax.constraints.tree_constraints(tree)
Extracts the individual constraints of a PyTree.
Standard arrays default to parax.constraints.RealLine.
Note that this function does not allow non-array/constrainable leaf nodes.
If you have leaves in your tree that are neither arrays nor derive
from parax.constraints.AbstractConstrainable, be sure to mark
them as static or filter them out using e.g. eqx.filter first.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
tree
|
PyTree
|
The PyTree model to extract constraints from. |
required |
Returns:
| Type | Description |
|---|---|
PyTree
|
A PyTree representing the active constraints. |
Source code in parax/constraints.py
878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 | |
parax.constraints.tree_leafwise_constraint(tree)
Extracts the single leafwise constraint of a PyTree.
Wraps the output of parax.constraints.tree_constraints
in a parax.constraints.Leafwise constraint to define
a single constraint that matches the shape of tree.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
tree
|
PyTree
|
The PyTree model containing probabilistic nodes or standard arrays. |
required |
Returns:
| Type | Description |
|---|---|
Leafwise
|
A single constraint whose shape matches the structure of |
Source code in parax/constraints.py
910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 | |
parax.constraints.tree_constrain(tree, constraints)
Applies a PyTree of constraints to a PyTree of constrainable PyTrees.
Standard arrays will be returned untouched if the matching constraint
is a RealLine. Attempting to apply a bounded constraint directly
to a standard array will raise an error.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
tree
|
PyTree
|
The PyTree model to update. Must have a matching PyTree structure
to |
required |
constraints
|
PyTree
|
A PyTree of |
required |
Returns:
| Type | Description |
|---|---|
PyTree
|
A new PyTree with the constraints applied. |
Source code in parax/constraints.py
965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 | |
parax.constraints.intersect(a, b)
Calculates the intersection of two constraints. Returns the most specific constraint class possible.
Each bound is the tighter of the two, keeping its closedness. Where both
constraints share a bound, finite or infinite, open wins. The result is a named
class (RealLine, Positive, ...) only when that class's closedness matches.
Source code in parax/constraints.py
1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 | |