Skip to content

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 bounds is included in the constrained space. Each side matches the structure of its bound, with a bool (or bool array) for each leaf: True for an edge a value may sit on exactly, False for one it may only approach. An infinite bound is closed unless asked otherwise, so ±∞ is inside.

bijector AbstractVar[AbstractBijector]

A distreqx.bijectors.AbstractBijector mapping from the unconstrained real line to the physical space. It maps onto the open interior, so a value exactly on a closed bound has an infinite raw value.

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 bounds where there is no finite extent to normalise against, such as an unbounded or half-bounded domain.

base_bijector AbstractVar[AbstractBijector]

A distreqx.bijectors.AbstractBijector mapping from the orthogonal base_bounds space into the physical bounds space, edge onto edge. Defaults to Identity unless geometric skews are present.

clip(value)

Clip a value to lie within this constraint.

Source code in parax/constraints.py
75
76
77
78
79
def clip(self, value: PyTree) -> PyTree:
    """
    Clip a value to lie within this constraint.
    """
    return jax.tree.map(jnp.clip, value, self.bounds[0], self.bounds[1])

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
def is_outside(self, value: PyTree) -> PyTree:
    """
    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.
    """
    lower, upper = self.bounds
    lower_closed, upper_closed = self.closed

    def _is_outside(x, l, u, lc, uc):
        above_lower = jnp.where(lc, x >= l, x > l)
        below_upper = jnp.where(uc, x <= u, x < u)
        return jnp.logical_not(jnp.logical_and(above_lower, below_upper))

    return jax.tree.map(_is_outside, value, lower, upper, lower_closed, upper_closed)

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
def midpoint(self) -> PyTree:
    """
    Returns the midpoint of the constraint.

    Note that non-finite constraints may return infinity.
    """
    return jax.tree.map(lambda a, b: (a + b) / 2.0, self.bounds[0], self.bounds[1])

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

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
@abstractmethod
def constrain(self, constraint: AbstractConstraint) -> Self:
    """
    Returns a new instance of the PyTree with the updated constraint,
    ensuring internal state (like unconstrained raw values) is 
    recalculated if necessary.

    Args:
        constraint: The new constraint to apply.

    Returns:
        A new instance of the constrainable PyTree.
    """
    raise NotImplementedError

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
def __init__(self, shape: Any = ()):
    """
    Args:
        shape: The expected shape of the unconstrained parameter.
    """
    self.shape = shape

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 (lower, upper) pair. The upper bound is +∞.

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 lower and +∞, or a (lower, upper) pair.

True
Source code in parax/constraints.py
192
193
194
195
196
197
198
199
200
201
202
203
204
def __init__(
    self,
    lower: Union[float, Array],
    closed: bool | tuple[bool, bool] = True,
):
    """
    Args:
        lower: The lower bound.
        closed: Whether the bounds themselves are included: one bool for both
            `lower` and +∞, or a `(lower, upper)` pair.
    """
    self.lower = jnp.asarray(lower, dtype=float)
    self.closed = _as_pair(closed)

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 (lower, upper) pair. The lower bound is -∞.

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 upper, or a (lower, upper) pair.

True
Source code in parax/constraints.py
227
228
229
230
231
232
233
234
235
236
237
238
239
def __init__(
    self,
    upper: Union[float, Array],
    closed: bool | tuple[bool, bool] = True,
):
    """
    Args:
        upper: The upper bound.
        closed: Whether the bounds themselves are included: one bool for both
            -∞ and `upper`, or a `(lower, upper)` pair.
    """
    self.upper = jnp.asarray(upper, dtype=float)
    self.closed = _as_pair(closed)

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 lower and upper themselves are included, as a (lower, upper) pair.

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 (lower, upper) pair.

True
Source code in parax/constraints.py
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
def __init__(
    self,
    lower: Union[float, Array],
    upper: Union[float, Array],
    closed: bool | tuple[bool, bool] = True,
):
    """
    Args:
        lower: The lower bound.
        upper: The upper bound.
        closed: Whether the bounds themselves are included: one bool for both,
            or a `(lower, upper)` pair.
    """
    self.lower = jnp.asarray(lower, dtype=float)
    self.upper = jnp.asarray(upper, dtype=float)
    self.closed = _as_pair(closed)

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
def __init__(self, shape: Any = (), dtype: Any = None):
    """
    Args:
        shape: The shape of the parameter array.
        dtype: The JAX data type of the parameter array.
    """
    super().__init__(lower=jnp.zeros(shape, dtype=dtype), closed=(False, True))

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
def __init__(self, shape: Any = (), dtype: Any = None):
    """
    Args:
        shape: The shape of the parameter array.
        dtype: The JAX data type of the parameter array.
    """
    super().__init__(lower=jnp.zeros(shape, dtype=dtype), closed=True)

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
def __init__(self, shape: Any = (), dtype: Any = None):
    """
    Args:
        shape: The shape of the parameter array.
        dtype: The JAX data type of the parameter array.
    """
    super().__init__(upper=jnp.zeros(shape, dtype=dtype), closed=(True, False))

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
def __init__(self, shape: Any = (), dtype: Any = None):
    """
    Args:
        shape: The shape of the parameter array.
        dtype: The JAX data type of the parameter array.
    """
    super().__init__(upper=jnp.zeros(shape, dtype=dtype), closed=True)

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 AbstractConstraint leaves.

Source code in parax/constraints.py
530
531
532
533
534
535
536
537
def __init__(
    self, 
    tree: PyTree[AbstractConstraint],
):
    leaves = jax.tree.leaves(tree, is_leaf=is_constraint)
    if not leaves:
        raise ValueError("The pytree of `tree` cannot be empty.")
    self.tree = tree

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 (lower, upper), each a PyTree matching the constrained value's structure.

closed tuple[PyTree, PyTree]

Whether each of bounds is included, as a (lower, upper) pair of PyTrees matching bounds, with a bool (or bool array) for each leaf.

base_bounds tuple[PyTree, PyTree]

The orthogonal base boundaries. Defaults to bounds if omitted.

base_bijector AbstractBijector

The bijector mapping from base_bounds to bounds. Defaults to Identity if omitted.

Parameters:

Name Type Description Default
bijector AbstractBijector

The custom distreqx bijector.

required
bounds tuple[PyTree, PyTree]

A tuple of (lower, upper) defining the physical boundaries of the constrained space, each a PyTree matching the constrained value's structure. Defaults to (-inf, inf).

(array(-inf), array(inf))
closed bool | tuple[PyTree, PyTree]

Whether the bounds themselves are included: one bool for every bound, or a (lower, upper) pair, each side a bool for all of that side's leaves or a PyTree of them matching the bound. Defaults to closed, infinite bounds included.

True
base_bounds tuple[PyTree, PyTree] | None

Optional. A tuple of (lower, upper) defining the orthogonal base boundaries. If None, defaults to bounds.

None
base_bijector AbstractBijector | None

Optional. The bijector handling spatial skew/correlation. If None, defaults to distreqx.bijectors.Identity.

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
def __init__(
    self,
    bijector: AbstractBijector,
    bounds: tuple[PyTree, PyTree] = (jnp.array(-jnp.inf), jnp.array(jnp.inf)),
    closed: bool | tuple[PyTree, PyTree] = True,
    base_bounds: tuple[PyTree, PyTree] | None = None,
    base_bijector: AbstractBijector | None = None
):
    """
    Args:
        bijector: The custom `distreqx` bijector.
        bounds: A tuple of `(lower, upper)` defining the physical
            boundaries of the constrained space, each a PyTree matching
            the constrained value's structure. Defaults to `(-inf, inf)`.
        closed: Whether the bounds themselves are included: one bool for every
            bound, or a `(lower, upper)` pair, each side a bool for all of that
            side's leaves or a PyTree of them matching the bound. Defaults to
            closed, infinite bounds included.
        base_bounds: Optional. A tuple of `(lower, upper)` defining the orthogonal
            base boundaries. If None, defaults to `bounds`.
        base_bijector: Optional. The bijector handling spatial skew/correlation.
            If None, defaults to `distreqx.bijectors.Identity`.
    """
    self.bijector = bijector
    self.bounds = tuple(jax.tree.map(jnp.asarray, b) for b in bounds)
    self.closed = tuple(
        jax.tree.map(lambda _: side, bound) if isinstance(side, bool) else side
        for side, bound in zip(_as_pair(closed), self.bounds)
    )

    # Default base_bounds to physical bounds if not provided
    if base_bounds is None:
        self.base_bounds = self.bounds
    else:
        self.base_bounds = tuple(jax.tree.map(jnp.asarray, b) for b in base_bounds)

    # Default base_bijector to Identity if not provided
    if base_bijector is None:
        self.base_bijector = Identity()
    else:
        self.base_bijector = base_bijector

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
def tree_constraints(tree: PyTree) -> PyTree:
    """
    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.    

    Args:
        tree: The PyTree model to extract constraints from.

    Returns:
        A PyTree representing the active constraints.
    """
    from parax.wrappers import as_unwrapped

    def _get_constraint(x):
        if is_constrained(x):
            return as_unwrapped(x.constraint)
        if eqx.is_inexact_array(x):
            return RealLine(shape=x.shape)
        raise ValueError(
            f"Found a leaf node of type {type(x)} that is neither constrained "
            f"nor an array in `parax.constraints.tree_constraints`. Value: {x}"
        )

    return jax.tree_util.tree_map(_get_constraint, tree, is_leaf=is_constrained)

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

Source code in parax/constraints.py
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
def tree_leafwise_constraint(tree: PyTree) -> Leafwise:
    """
    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`.

    Args:
        tree: The PyTree model containing probabilistic nodes or standard arrays.

    Returns:
        A single constraint whose shape matches the structure of `tree`.
    """
    return Leafwise(tree_constraints(tree)) 

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

required
constraints PyTree

A PyTree of parax.AbstractConstraint objects.

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
def tree_constrain(tree: PyTree, constraints: PyTree) -> PyTree:
    """
    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.

    Args:
        tree: The PyTree model to update. Must have a matching PyTree structure 
            to `constraints`.
        constraints: A PyTree of `parax.AbstractConstraint` objects.

    Returns:
        A new PyTree with the constraints applied.
    """
    def _apply_constraint(x, c):
        if is_constrainable(x):
            return x.constrain(c)
        if eqx.is_inexact_array(x):
            if isinstance(c, RealLine):
                return x
            raise TypeError(
                "Cannot apply a bounded constraint to a raw JAX array directly. "
                "Ensure the array is wrapped in a `parax.Constrained` variable first."
            )
        raise ValueError(
            f"Found a leaf node of type {type(x)} that is neither constrainable "
            f"nor an array in `parax.constraints.tree_constrain`. Value: {x}"
        )

    return jax.tree_util.tree_map(
        _apply_constraint, tree, constraints, is_leaf=is_constrainable
    )

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
def intersect(a: AbstractConstraint, b: AbstractConstraint) -> AbstractConstraint:
    """
    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.
    """
    a_lower, a_upper = a.bounds
    b_lower, b_upper = b.bounds
    a_lower_closed, a_upper_closed = a.closed
    b_lower_closed, b_upper_closed = b.closed

    lower = jnp.maximum(a_lower, b_lower)
    upper = jnp.minimum(a_upper, b_upper)

    def _closedness(a_bound, b_bound, a_closed, b_closed, a_tighter):
        closed = jnp.where(
            a_tighter,
            a_closed,
            jnp.where(a_bound == b_bound, jnp.logical_and(a_closed, b_closed), b_closed),
        )
        # One flag per bound: open wins wherever the elements disagree.
        return bool(jnp.all(closed))

    lower_closed = _closedness(a_lower, b_lower, a_lower_closed, b_lower_closed, a_lower > b_lower)
    upper_closed = _closedness(a_upper, b_upper, a_upper_closed, b_upper_closed, a_upper < b_upper)

    # Convert to concrete numpy arrays for boolean checks during init
    np_lower = jnp.asarray(lower)
    np_upper = jnp.asarray(upper)

    np_lower, np_upper = eqx.error_if(
        (np_lower, np_upper),
        jnp.any(jnp.greater_equal(np_lower, np_upper)),
        f"Constraint intersection is empty or invalid."
    )

    is_neginf_lower = jnp.all(jnp.isneginf(np_lower))
    is_posinf_upper = jnp.all(jnp.isposinf(np_upper))
    is_zero_lower = jnp.all(jnp.equal(np_lower, 0.0))
    is_zero_upper = jnp.all(jnp.equal(np_upper, 0.0))

    # Resolve to the most specific constraint class whose closedness matches
    closed = (lower_closed, upper_closed)
    if is_neginf_lower and is_posinf_upper and closed == (True, True):
        return RealLine()
    elif is_zero_lower and is_posinf_upper and closed == (True, True):
        return NonNegative()
    elif is_zero_lower and is_posinf_upper and closed == (False, True):
        return Positive()
    elif is_neginf_lower and is_zero_upper and closed == (True, True):
        return NonPositive()
    elif is_neginf_lower and is_zero_upper and closed == (True, False):
        return Negative()
    elif is_posinf_upper and not is_neginf_lower:
        return GreaterThan(lower, closed=closed)
    elif is_neginf_lower and not is_posinf_upper:
        return LessThan(upper, closed=closed)
    else:
        return Interval(lower, upper, closed=closed)