Skip to content

Core API

The pyrox._core subpackage owns the Equinox-to-NumPyro bridge that every other subpackage composes on top of.

Public surface

Core: Equinox-to-NumPyro bridge primitives.

Public surface:

  • PyroxModule — Equinox module with pyrox_param / pyrox_sample
  • PyroxParam — declarative parameter descriptor
  • PyroxSample — declarative sample descriptor
  • Parameterized — param registry with priors, guides, and modes
  • pyrox_method — decorator that activates the per-call context

PyroxParam

Bases: NamedTuple

Lightweight metadata container for a parameter site.

Bundles init value, constraint, and optional event dimension as a single descriptor. This type is a plain value object — higher-level APIs that consume it (for example a future declarative registration helper) live elsewhere; PyroxModule.pyrox_param takes the fields individually as keyword arguments.

Attributes:

Name Type Description
init_value Any

Initial value, lazy callable, or None to look up an existing param site by name.

constraint Any

NumPyro constraint on the parameter domain; None means unconstrained real.

event_dim int | None

Number of rightmost event dimensions, or None.

Source code in packages/pyrox/src/pyrox/_core/descriptors.py
class PyroxParam(NamedTuple):
    """Lightweight metadata container for a parameter site.

    Bundles init value, constraint, and optional event dimension as a
    single descriptor. This type is a plain value object — higher-level
    APIs that consume it (for example a future declarative registration
    helper) live elsewhere; `PyroxModule.pyrox_param` takes the
    fields individually as keyword arguments.

    Attributes:
        init_value: Initial value, lazy callable, or ``None`` to look up
            an existing param site by name.
        constraint: NumPyro constraint on the parameter domain; ``None``
            means unconstrained real.
        event_dim: Number of rightmost event dimensions, or ``None``.
    """

    init_value: Any = None
    constraint: Any = None
    event_dim: int | None = None

PyroxSample dataclass

Lightweight metadata container for a random sample site.

Wraps the prior — either a numpyro.distributions.Distribution or a callable (self) -> Distribution for dependent priors that reference other sampled values on the same module. Like PyroxParam, this is a plain value object; call PyroxModule.pyrox_sample with the underlying prior directly.

Source code in packages/pyrox/src/pyrox/_core/descriptors.py
@dataclass(frozen=True)
class PyroxSample:
    """Lightweight metadata container for a random sample site.

    Wraps the prior — either a `numpyro.distributions.Distribution`
    or a callable ``(self) -> Distribution`` for dependent priors that
    reference other sampled values on the same module. Like
    `PyroxParam`, this is a plain value object; call
    `PyroxModule.pyrox_sample` with the underlying prior directly.
    """

    prior: Any | Callable[[Any], Any]

Parameterized

Bases: PyroxModule

Shared base for modules with priors, constraints, and mode switching.

Subclasses typically declare parameters inside setup, which is invoked automatically after __init__ completes. Use register_param to declare a parameter, set_prior to attach a prior, autoguide to pick a guide type, and set_mode to switch between sampling from the prior and sampling from the guide.

Per-instance state (params, priors, guides, mode) lives in a class-level registry keyed by id(self). Cleanup happens via weakref.finalize when the instance is collected; call _teardown for explicit cleanup.

Source code in packages/pyrox/src/pyrox/_core/parameterized.py
class Parameterized(PyroxModule):
    """Shared base for modules with priors, constraints, and mode switching.

    Subclasses typically declare parameters inside `setup`, which is
    invoked automatically after ``__init__`` completes. Use
    `register_param` to declare a parameter, `set_prior` to
    attach a prior, `autoguide` to pick a guide type, and
    `set_mode` to switch between sampling from the prior and
    sampling from the guide.

    Per-instance state (params, priors, guides, mode) lives in a
    class-level registry keyed by ``id(self)``. Cleanup happens via
    `weakref.finalize` when the instance is collected; call
    `_teardown` for explicit cleanup.
    """

    _registry: ClassVar[dict[int, _State]] = {}

    def __post_init__(self) -> None:
        setup = getattr(self, "setup", None)
        if callable(setup):
            setup()

    def _state(self) -> _State:
        key = id(self)
        state = Parameterized._registry.get(key)
        if state is None:
            state = _State()
            Parameterized._registry[key] = state
            with contextlib.suppress(TypeError):
                weakref.finalize(self, Parameterized._registry.pop, key, None)
        return state

    def _entry(self, name: str) -> _Entry:
        entry = self._state().params.get(name)
        if entry is None:
            raise KeyError(
                f"parameter {name!r} not registered; call register_param "
                "first. If this module previously worked and was then "
                "rebuilt by a functional pytree update (eqx.tree_at, "
                "eqx.apply_updates, jit/filter_jit with the module as an "
                "argument, flatten/unflatten, or checkpoint load), note "
                "that reconstruction skips __init__/setup(), so the "
                "rebuilt copy has an empty registry — keep the original "
                "instance for probabilistic calls, or reconstruct via "
                "__init__ so setup() re-registers its parameters."
            )
        return entry

    def register_param(
        self,
        name: str,
        init_value: Any,
        constraint: Any = None,
    ) -> None:
        self._state().params[name] = _Entry(
            init_value=init_value, constraint=constraint
        )

    def set_prior(self, name: str, prior: Any) -> None:
        self._entry(name).prior = prior

    def autoguide(self, name: str, guide_type: GuideType) -> None:
        if guide_type not in _VALID_GUIDES:
            raise ValueError(
                f"guide_type must be one of {sorted(_VALID_GUIDES)!r}, "
                f"got {guide_type!r}"
            )
        self._entry(name).guide_type = guide_type

    def set_mode(self, mode: Mode) -> None:
        if mode not in ("model", "guide"):
            raise ValueError(f"mode must be 'model' or 'guide', got {mode!r}")
        self._state().mode = mode

    def get_param(self, name: str) -> Any:
        entry = self._entry(name)
        state = self._state()
        if state.mode == "model" and entry.prior is not None:
            return self.pyrox_sample(name, entry.prior)
        if state.mode == "guide" and entry.prior is not None:
            return self._guide_param(name, entry)
        return self.pyrox_param(name, entry.init_value, constraint=entry.constraint)

    def load_pyro_samples(self) -> None:
        for name in list(self._state().params):
            self.get_param(name)

    def _teardown(self) -> None:
        Parameterized._registry.pop(id(self), None)
        super()._teardown()

    def _guide_param(self, name: str, entry: _Entry) -> Any:
        guide = entry.guide_type
        if guide == "delta":
            self._reserve_aux(name, ("_loc",))
            return self._guide_delta(name, entry)
        if guide == "normal":
            self._reserve_aux(name, ("_loc", "_scale"))
            return self._guide_normal(name, entry)
        raise NotImplementedError(
            f"guide_type {guide!r} is not yet supported at the "
            "get_param level; materialize via a dedicated guide layer."
        )

    def _reserve_aux(self, name: str, suffixes: tuple[str, ...]) -> None:
        """Fail loudly if a guide's auxiliary site name collides with a
        user-registered parameter.

        The ``normal`` / ``delta`` guides materialize backing sites named
        ``{name}_loc`` (and ``{name}_scale``). If the user also registered a
        parameter with that exact name, the two would share a fully-qualified
        site name and silently clobber each other via the per-call cache (or
        trip NumPyro's duplicate-site check). Reject it at guide time with an
        actionable message rather than let it corrupt inference.
        """
        params = self._state().params
        for suffix in suffixes:
            aux = f"{name}{suffix}"
            if aux in params:
                raise ValueError(
                    f"the guide for parameter {name!r} needs an auxiliary "
                    f"site named {aux!r}, but a parameter with that name is "
                    "already registered; rename one of them to avoid the "
                    "collision."
                )

    def _guide_delta(self, name: str, entry: _Entry) -> Any:
        """Point-estimate (MAP) guide — a constrained param replayed as a Delta.

        Mirrors `numpyro.infer.autoguide.AutoDelta`: the value is a
        `numpyro.param` in the **prior's** support, wrapped in a
        `numpyro.distributions.Delta` **sample** site. NumPyro's ``replay``
        handler conditions the model only on guide *sample* sites, so a bare
        param under the latent's name is invisible to SVI and raises
        ``RuntimeError: Site ... must be sampled in trace``.

        Support and event rank are taken from the resolved prior — matching
        the model site — rather than from the registered ``constraint``:

        * The point estimate must live in the model prior's support, or
          ``replay`` evaluates the prior's log-density outside its domain and
          SVI diverges (a positive prior on an unconstrained-registered param
          would otherwise let ``{name}_loc`` slide to <= 0).
        * The Delta's ``event_dim`` is the prior's, so batched priors keep
          their plate dimensions instead of collapsing the whole value into a
          single event (which would mis-expand under ``numpyro.plate`` /
          block subsampling of the backing param).
        """
        prior = entry.prior
        if callable(prior) and not isinstance(prior, dist.Distribution):
            prior = prior(self)
        if isinstance(prior, dist.Distribution):
            support = prior.support
            event_dim = prior.event_dim
        else:  # non-distribution prior (deterministic value): fall back.
            support = entry.constraint
            event_dim = 0
        loc = self.pyrox_param(
            f"{name}_loc",
            entry.init_value,
            constraint=support,
            event_dim=event_dim,
        )
        loc = jnp.asarray(loc)
        return self.pyrox_sample(name, dist.Delta(loc, event_dim=event_dim))

    def _guide_normal(self, name: str, entry: _Entry) -> Any:
        """Mean-field normal guide in unconstrained space.

        When ``entry.constraint`` is non-trivial, the latent site is a
        ``TransformedDistribution`` wrapping ``Normal(loc, scale)`` with
        the constraint's bijection, so guide draws always land in the
        prior's support. ``loc`` is initialized by the inverse transform
        of ``init_value`` so guide and prior agree at step zero.
        """
        init = jnp.asarray(entry.init_value)
        # A vector parameter usually carries a joint prior (e.g.
        # ``LogNormal(...).expand([D]).to_event(1)``). The guide site has to
        # declare the same event rank or Trace_ELBO rejects the model/guide
        # pair, so read it off the prior rather than assuming element-wise.
        # Callable (dependent) priors must be resolved first, exactly as
        # `_guide_delta` and `pyrox_sample` do — reading `event_dim` off the
        # callable itself would silently report rank 0.
        prior = entry.prior
        if callable(prior) and not isinstance(prior, dist.Distribution):
            prior = prior(self)
        prior_event_dim = prior.event_dim if isinstance(prior, dist.Distribution) else 0
        if _is_real_support(entry.constraint):
            loc = self.pyrox_param(f"{name}_loc", init)
            scale = self.pyrox_param(
                f"{name}_scale",
                jnp.ones_like(init) * 0.1,
                constraint=dist.constraints.positive,
            )
            return self.pyrox_sample(
                name, dist.Normal(loc, scale).to_event(prior_event_dim)
            )
        transform = _biject_to(entry.constraint)
        # loc AND scale live in unconstrained space. For shape-changing
        # transforms (e.g. StickBreaking for a simplex: K -> K-1) the
        # unconstrained shape differs from ``init``; sizing scale from
        # ``init`` would fail to broadcast against ``transform.inv(init)``.
        unconstrained = transform.inv(init)
        loc = self.pyrox_param(f"{name}_loc", unconstrained)
        scale = self.pyrox_param(
            f"{name}_scale",
            jnp.full_like(unconstrained, 0.1),
            constraint=dist.constraints.positive,
        )
        # Promote the base to the transform's domain event rank so the
        # TransformedDistribution's event structure matches the constrained
        # support (a no-op for element-wise transforms like Exp/Sigmoid).
        # The base lives in the transform's *domain*, while the prior's
        # event rank is stated in its codomain, so subtract the rank the
        # transform itself adds. For an element-wise transform the delta is
        # zero and this is just the prior's rank; for a shape-changing one
        # (corr_cholesky: vector domain -> matrix codomain) taking the
        # prior's rank directly would over-promote a base that has fewer
        # dimensions than that, and construction would fail.
        delta = transform.codomain.event_dim - transform.domain.event_dim
        event_dim = max(transform.domain.event_dim, prior_event_dim - delta)
        base = dist.Normal(loc, scale).to_event(event_dim)
        return self.pyrox_sample(name, dist.TransformedDistribution(base, transform))

PyroxModule

Bases: Module

Equinox module with NumPyro site registration and per-call caching.

Subclasses register deterministic parameters via pyrox_param and random variables via pyrox_sample. Wrap the method that drives registration (typically __call__) with pyrox_method so the per-call _Context is active for the duration of the call.

Without the decorator the cache is inactive and duplicate references to the same site within one trace will hit NumPyro's uniqueness check.

Source code in packages/pyrox/src/pyrox/_core/pyrox_module.py
class PyroxModule(eqx.Module):
    """Equinox module with NumPyro site registration and per-call caching.

    Subclasses register deterministic parameters via `pyrox_param`
    and random variables via `pyrox_sample`. Wrap the method that
    drives registration (typically ``__call__``) with `pyrox_method`
    so the per-call ``_Context`` is active for the duration of the call.

    Without the decorator the cache is inactive and duplicate references
    to the same site within one trace will hit NumPyro's uniqueness check.
    """

    _contexts: ClassVar[dict[int, _Context]] = {}

    def _get_context(self) -> _Context:
        key = id(self)
        ctx = PyroxModule._contexts.get(key)
        if ctx is None:
            ctx = _Context()
            PyroxModule._contexts[key] = ctx
            with contextlib.suppress(TypeError):
                weakref.finalize(self, PyroxModule._contexts.pop, key, None)
        return ctx

    def _pyrox_scope_name(self) -> str:
        """Per-instance scope used when building fully-qualified site names.

        Uses an explicit ``pyrox_name`` attribute if the module defines one
        (as a field or class variable); otherwise falls back to the **class
        name**. Both are deterministic, so site names are stable across
        Python runs, checkpoint round-trips, and — critically — Equinox
        pytree reconstruction (``eqx.tree_at``, ``eqx.filter_jit`` with the
        module as an argument, ``jax.tree.unflatten``, deserialization).
        The previous ``{ClassName}_{id}`` fallback changed on every
        reconstruction, silently desynchronizing site names between e.g.
        an MCMC run and a later ``Predictive`` on a rebuilt copy.

        The scope must be **unique among the instances participating in a
        single trace**. With the class-name fallback, two *unnamed*
        instances of the same class collide: under ``handlers.trace`` this
        raises loudly — NumPyro's uniqueness assertion for sample sites,
        and pyrox's duplicate-scope guard in `pyrox_param` for param
        sites (which NumPyro would otherwise silently alias). Under a
        bare ``handlers.seed`` (no trace) there is no uniqueness check,
        so the collision is silent. When stacking several instances of
        one class in a model, give each a distinct ``pyrox_name`` (a
        per-instance field or constructor argument).
        """
        name = getattr(self, "pyrox_name", None)
        if isinstance(name, str) and name:
            return name
        return type(self).__name__

    def _pyrox_fullname(self, name: str) -> str:
        return f"{self._pyrox_scope_name()}.{name}"

    def pyrox_param(
        self,
        name: str,
        init_value: Any,
        *,
        constraint: Any = None,
        event_dim: int | None = None,
    ) -> Any:
        ctx = self._get_context()
        fullname = self._pyrox_fullname(name)
        if ctx.active:
            cached = ctx.get(fullname)
            if cached is not _MISSING:
                return cached
        kwargs: dict[str, Any] = {}
        if constraint is not None:
            kwargs["constraint"] = constraint
        if event_dim is not None:
            kwargs["event_dim"] = event_dim
        # NumPyro's trace asserts uniqueness for sample sites but silently
        # tolerates duplicate param registrations (last write wins). Without
        # a guard, two instances sharing a scope (e.g. unnamed siblings of
        # one class under the class-name fallback) would silently alias
        # their parameters under SVI/MAP. The guard pre-checks the
        # predicted recorded name against every trace this message will
        # actually reach (see `_visible_traces` for the handler-stack
        # semantics: nested traces, scope prefixes incl. hide_types, and
        # block visibility). Same-instance re-registration across calls in
        # one trace (weight sharing) stays allowed via the per-trace
        # ownership sets; ownership does not leak across traces.
        visible = _visible_traces(fullname)
        for tr, recorded in visible:
            if recorded in tr and recorded not in ctx.trace_owned(tr):
                raise ValueError(
                    f"param site {recorded!r} was already registered "
                    "in this trace by a different module instance. Two "
                    f"instances of {type(self).__name__} are sharing "
                    f"the scope {self._pyrox_scope_name()!r} — give "
                    "each a distinct pyrox_name."
                )
        value = numpyro.param(fullname, init_value, **kwargs)
        for tr, recorded in visible:
            ctx.trace_owned(tr).add(recorded)
        return ctx.set(fullname, value)

    def pyrox_sample(self, name: str, prior: Any) -> Any:
        ctx = self._get_context()
        fullname = self._pyrox_fullname(name)
        if ctx.active:
            cached = ctx.get(fullname)
            if cached is not _MISSING:
                return cached
        resolved = (
            prior(self)
            if callable(prior) and not isinstance(prior, dist.Distribution)
            else prior
        )
        if isinstance(resolved, dist.Distribution):
            value = numpyro.sample(fullname, resolved)
        else:
            value = numpyro.deterministic(fullname, resolved)
        return ctx.set(fullname, value)

    def _teardown(self) -> None:
        """Remove this instance's cached context.

        Class-level registries are keyed by ``id(self)``. Equinox modules
        are typically weak-referenceable, so cleanup normally happens via
        `weakref.finalize`. Call this explicitly in environments where
        weak refs are not available or when you need deterministic cleanup.
        """
        PyroxModule._contexts.pop(id(self), None)

pyrox_method(fn: Callable[..., Any]) -> Callable[..., Any]

Wrap a method so its body runs inside the module's per-call context.

Apply to __call__ (and any other method that registers pyrox sites) so the _Context cache is active for the duration of the call. The cache is cleared when the outermost decorated call returns.

Source code in packages/pyrox/src/pyrox/_core/pyrox_module.py
def pyrox_method(fn: Callable[..., Any]) -> Callable[..., Any]:
    """Wrap a method so its body runs inside the module's per-call context.

    Apply to ``__call__`` (and any other method that registers pyrox sites)
    so the ``_Context`` cache is active for the duration of the call. The
    cache is cleared when the outermost decorated call returns.
    """

    @functools.wraps(fn)
    def wrapper(self: PyroxModule, *args: Any, **kwargs: Any) -> Any:
        with self._get_context():
            return fn(self, *args, **kwargs)

    return wrapper

PyroxModule

PyroxModule

Bases: Module

Equinox module with NumPyro site registration and per-call caching.

Subclasses register deterministic parameters via pyrox_param and random variables via pyrox_sample. Wrap the method that drives registration (typically __call__) with pyrox_method so the per-call _Context is active for the duration of the call.

Without the decorator the cache is inactive and duplicate references to the same site within one trace will hit NumPyro's uniqueness check.

Source code in packages/pyrox/src/pyrox/_core/pyrox_module.py
class PyroxModule(eqx.Module):
    """Equinox module with NumPyro site registration and per-call caching.

    Subclasses register deterministic parameters via `pyrox_param`
    and random variables via `pyrox_sample`. Wrap the method that
    drives registration (typically ``__call__``) with `pyrox_method`
    so the per-call ``_Context`` is active for the duration of the call.

    Without the decorator the cache is inactive and duplicate references
    to the same site within one trace will hit NumPyro's uniqueness check.
    """

    _contexts: ClassVar[dict[int, _Context]] = {}

    def _get_context(self) -> _Context:
        key = id(self)
        ctx = PyroxModule._contexts.get(key)
        if ctx is None:
            ctx = _Context()
            PyroxModule._contexts[key] = ctx
            with contextlib.suppress(TypeError):
                weakref.finalize(self, PyroxModule._contexts.pop, key, None)
        return ctx

    def _pyrox_scope_name(self) -> str:
        """Per-instance scope used when building fully-qualified site names.

        Uses an explicit ``pyrox_name`` attribute if the module defines one
        (as a field or class variable); otherwise falls back to the **class
        name**. Both are deterministic, so site names are stable across
        Python runs, checkpoint round-trips, and — critically — Equinox
        pytree reconstruction (``eqx.tree_at``, ``eqx.filter_jit`` with the
        module as an argument, ``jax.tree.unflatten``, deserialization).
        The previous ``{ClassName}_{id}`` fallback changed on every
        reconstruction, silently desynchronizing site names between e.g.
        an MCMC run and a later ``Predictive`` on a rebuilt copy.

        The scope must be **unique among the instances participating in a
        single trace**. With the class-name fallback, two *unnamed*
        instances of the same class collide: under ``handlers.trace`` this
        raises loudly — NumPyro's uniqueness assertion for sample sites,
        and pyrox's duplicate-scope guard in `pyrox_param` for param
        sites (which NumPyro would otherwise silently alias). Under a
        bare ``handlers.seed`` (no trace) there is no uniqueness check,
        so the collision is silent. When stacking several instances of
        one class in a model, give each a distinct ``pyrox_name`` (a
        per-instance field or constructor argument).
        """
        name = getattr(self, "pyrox_name", None)
        if isinstance(name, str) and name:
            return name
        return type(self).__name__

    def _pyrox_fullname(self, name: str) -> str:
        return f"{self._pyrox_scope_name()}.{name}"

    def pyrox_param(
        self,
        name: str,
        init_value: Any,
        *,
        constraint: Any = None,
        event_dim: int | None = None,
    ) -> Any:
        ctx = self._get_context()
        fullname = self._pyrox_fullname(name)
        if ctx.active:
            cached = ctx.get(fullname)
            if cached is not _MISSING:
                return cached
        kwargs: dict[str, Any] = {}
        if constraint is not None:
            kwargs["constraint"] = constraint
        if event_dim is not None:
            kwargs["event_dim"] = event_dim
        # NumPyro's trace asserts uniqueness for sample sites but silently
        # tolerates duplicate param registrations (last write wins). Without
        # a guard, two instances sharing a scope (e.g. unnamed siblings of
        # one class under the class-name fallback) would silently alias
        # their parameters under SVI/MAP. The guard pre-checks the
        # predicted recorded name against every trace this message will
        # actually reach (see `_visible_traces` for the handler-stack
        # semantics: nested traces, scope prefixes incl. hide_types, and
        # block visibility). Same-instance re-registration across calls in
        # one trace (weight sharing) stays allowed via the per-trace
        # ownership sets; ownership does not leak across traces.
        visible = _visible_traces(fullname)
        for tr, recorded in visible:
            if recorded in tr and recorded not in ctx.trace_owned(tr):
                raise ValueError(
                    f"param site {recorded!r} was already registered "
                    "in this trace by a different module instance. Two "
                    f"instances of {type(self).__name__} are sharing "
                    f"the scope {self._pyrox_scope_name()!r} — give "
                    "each a distinct pyrox_name."
                )
        value = numpyro.param(fullname, init_value, **kwargs)
        for tr, recorded in visible:
            ctx.trace_owned(tr).add(recorded)
        return ctx.set(fullname, value)

    def pyrox_sample(self, name: str, prior: Any) -> Any:
        ctx = self._get_context()
        fullname = self._pyrox_fullname(name)
        if ctx.active:
            cached = ctx.get(fullname)
            if cached is not _MISSING:
                return cached
        resolved = (
            prior(self)
            if callable(prior) and not isinstance(prior, dist.Distribution)
            else prior
        )
        if isinstance(resolved, dist.Distribution):
            value = numpyro.sample(fullname, resolved)
        else:
            value = numpyro.deterministic(fullname, resolved)
        return ctx.set(fullname, value)

    def _teardown(self) -> None:
        """Remove this instance's cached context.

        Class-level registries are keyed by ``id(self)``. Equinox modules
        are typically weak-referenceable, so cleanup normally happens via
        `weakref.finalize`. Call this explicitly in environments where
        weak refs are not available or when you need deterministic cleanup.
        """
        PyroxModule._contexts.pop(id(self), None)

pyrox_method

pyrox_method(fn: Callable[..., Any]) -> Callable[..., Any]

Wrap a method so its body runs inside the module's per-call context.

Apply to __call__ (and any other method that registers pyrox sites) so the _Context cache is active for the duration of the call. The cache is cleared when the outermost decorated call returns.

Source code in packages/pyrox/src/pyrox/_core/pyrox_module.py
def pyrox_method(fn: Callable[..., Any]) -> Callable[..., Any]:
    """Wrap a method so its body runs inside the module's per-call context.

    Apply to ``__call__`` (and any other method that registers pyrox sites)
    so the ``_Context`` cache is active for the duration of the call. The
    cache is cleared when the outermost decorated call returns.
    """

    @functools.wraps(fn)
    def wrapper(self: PyroxModule, *args: Any, **kwargs: Any) -> Any:
        with self._get_context():
            return fn(self, *args, **kwargs)

    return wrapper

PyroxParam

PyroxParam

Bases: NamedTuple

Lightweight metadata container for a parameter site.

Bundles init value, constraint, and optional event dimension as a single descriptor. This type is a plain value object — higher-level APIs that consume it (for example a future declarative registration helper) live elsewhere; PyroxModule.pyrox_param takes the fields individually as keyword arguments.

Attributes:

Name Type Description
init_value Any

Initial value, lazy callable, or None to look up an existing param site by name.

constraint Any

NumPyro constraint on the parameter domain; None means unconstrained real.

event_dim int | None

Number of rightmost event dimensions, or None.

Source code in packages/pyrox/src/pyrox/_core/descriptors.py
class PyroxParam(NamedTuple):
    """Lightweight metadata container for a parameter site.

    Bundles init value, constraint, and optional event dimension as a
    single descriptor. This type is a plain value object — higher-level
    APIs that consume it (for example a future declarative registration
    helper) live elsewhere; `PyroxModule.pyrox_param` takes the
    fields individually as keyword arguments.

    Attributes:
        init_value: Initial value, lazy callable, or ``None`` to look up
            an existing param site by name.
        constraint: NumPyro constraint on the parameter domain; ``None``
            means unconstrained real.
        event_dim: Number of rightmost event dimensions, or ``None``.
    """

    init_value: Any = None
    constraint: Any = None
    event_dim: int | None = None

PyroxSample

PyroxSample dataclass

Lightweight metadata container for a random sample site.

Wraps the prior — either a numpyro.distributions.Distribution or a callable (self) -> Distribution for dependent priors that reference other sampled values on the same module. Like PyroxParam, this is a plain value object; call PyroxModule.pyrox_sample with the underlying prior directly.

Source code in packages/pyrox/src/pyrox/_core/descriptors.py
@dataclass(frozen=True)
class PyroxSample:
    """Lightweight metadata container for a random sample site.

    Wraps the prior — either a `numpyro.distributions.Distribution`
    or a callable ``(self) -> Distribution`` for dependent priors that
    reference other sampled values on the same module. Like
    `PyroxParam`, this is a plain value object; call
    `PyroxModule.pyrox_sample` with the underlying prior directly.
    """

    prior: Any | Callable[[Any], Any]

Parameterized

Parameterized

Bases: PyroxModule

Shared base for modules with priors, constraints, and mode switching.

Subclasses typically declare parameters inside setup, which is invoked automatically after __init__ completes. Use register_param to declare a parameter, set_prior to attach a prior, autoguide to pick a guide type, and set_mode to switch between sampling from the prior and sampling from the guide.

Per-instance state (params, priors, guides, mode) lives in a class-level registry keyed by id(self). Cleanup happens via weakref.finalize when the instance is collected; call _teardown for explicit cleanup.

Source code in packages/pyrox/src/pyrox/_core/parameterized.py
class Parameterized(PyroxModule):
    """Shared base for modules with priors, constraints, and mode switching.

    Subclasses typically declare parameters inside `setup`, which is
    invoked automatically after ``__init__`` completes. Use
    `register_param` to declare a parameter, `set_prior` to
    attach a prior, `autoguide` to pick a guide type, and
    `set_mode` to switch between sampling from the prior and
    sampling from the guide.

    Per-instance state (params, priors, guides, mode) lives in a
    class-level registry keyed by ``id(self)``. Cleanup happens via
    `weakref.finalize` when the instance is collected; call
    `_teardown` for explicit cleanup.
    """

    _registry: ClassVar[dict[int, _State]] = {}

    def __post_init__(self) -> None:
        setup = getattr(self, "setup", None)
        if callable(setup):
            setup()

    def _state(self) -> _State:
        key = id(self)
        state = Parameterized._registry.get(key)
        if state is None:
            state = _State()
            Parameterized._registry[key] = state
            with contextlib.suppress(TypeError):
                weakref.finalize(self, Parameterized._registry.pop, key, None)
        return state

    def _entry(self, name: str) -> _Entry:
        entry = self._state().params.get(name)
        if entry is None:
            raise KeyError(
                f"parameter {name!r} not registered; call register_param "
                "first. If this module previously worked and was then "
                "rebuilt by a functional pytree update (eqx.tree_at, "
                "eqx.apply_updates, jit/filter_jit with the module as an "
                "argument, flatten/unflatten, or checkpoint load), note "
                "that reconstruction skips __init__/setup(), so the "
                "rebuilt copy has an empty registry — keep the original "
                "instance for probabilistic calls, or reconstruct via "
                "__init__ so setup() re-registers its parameters."
            )
        return entry

    def register_param(
        self,
        name: str,
        init_value: Any,
        constraint: Any = None,
    ) -> None:
        self._state().params[name] = _Entry(
            init_value=init_value, constraint=constraint
        )

    def set_prior(self, name: str, prior: Any) -> None:
        self._entry(name).prior = prior

    def autoguide(self, name: str, guide_type: GuideType) -> None:
        if guide_type not in _VALID_GUIDES:
            raise ValueError(
                f"guide_type must be one of {sorted(_VALID_GUIDES)!r}, "
                f"got {guide_type!r}"
            )
        self._entry(name).guide_type = guide_type

    def set_mode(self, mode: Mode) -> None:
        if mode not in ("model", "guide"):
            raise ValueError(f"mode must be 'model' or 'guide', got {mode!r}")
        self._state().mode = mode

    def get_param(self, name: str) -> Any:
        entry = self._entry(name)
        state = self._state()
        if state.mode == "model" and entry.prior is not None:
            return self.pyrox_sample(name, entry.prior)
        if state.mode == "guide" and entry.prior is not None:
            return self._guide_param(name, entry)
        return self.pyrox_param(name, entry.init_value, constraint=entry.constraint)

    def load_pyro_samples(self) -> None:
        for name in list(self._state().params):
            self.get_param(name)

    def _teardown(self) -> None:
        Parameterized._registry.pop(id(self), None)
        super()._teardown()

    def _guide_param(self, name: str, entry: _Entry) -> Any:
        guide = entry.guide_type
        if guide == "delta":
            self._reserve_aux(name, ("_loc",))
            return self._guide_delta(name, entry)
        if guide == "normal":
            self._reserve_aux(name, ("_loc", "_scale"))
            return self._guide_normal(name, entry)
        raise NotImplementedError(
            f"guide_type {guide!r} is not yet supported at the "
            "get_param level; materialize via a dedicated guide layer."
        )

    def _reserve_aux(self, name: str, suffixes: tuple[str, ...]) -> None:
        """Fail loudly if a guide's auxiliary site name collides with a
        user-registered parameter.

        The ``normal`` / ``delta`` guides materialize backing sites named
        ``{name}_loc`` (and ``{name}_scale``). If the user also registered a
        parameter with that exact name, the two would share a fully-qualified
        site name and silently clobber each other via the per-call cache (or
        trip NumPyro's duplicate-site check). Reject it at guide time with an
        actionable message rather than let it corrupt inference.
        """
        params = self._state().params
        for suffix in suffixes:
            aux = f"{name}{suffix}"
            if aux in params:
                raise ValueError(
                    f"the guide for parameter {name!r} needs an auxiliary "
                    f"site named {aux!r}, but a parameter with that name is "
                    "already registered; rename one of them to avoid the "
                    "collision."
                )

    def _guide_delta(self, name: str, entry: _Entry) -> Any:
        """Point-estimate (MAP) guide — a constrained param replayed as a Delta.

        Mirrors `numpyro.infer.autoguide.AutoDelta`: the value is a
        `numpyro.param` in the **prior's** support, wrapped in a
        `numpyro.distributions.Delta` **sample** site. NumPyro's ``replay``
        handler conditions the model only on guide *sample* sites, so a bare
        param under the latent's name is invisible to SVI and raises
        ``RuntimeError: Site ... must be sampled in trace``.

        Support and event rank are taken from the resolved prior — matching
        the model site — rather than from the registered ``constraint``:

        * The point estimate must live in the model prior's support, or
          ``replay`` evaluates the prior's log-density outside its domain and
          SVI diverges (a positive prior on an unconstrained-registered param
          would otherwise let ``{name}_loc`` slide to <= 0).
        * The Delta's ``event_dim`` is the prior's, so batched priors keep
          their plate dimensions instead of collapsing the whole value into a
          single event (which would mis-expand under ``numpyro.plate`` /
          block subsampling of the backing param).
        """
        prior = entry.prior
        if callable(prior) and not isinstance(prior, dist.Distribution):
            prior = prior(self)
        if isinstance(prior, dist.Distribution):
            support = prior.support
            event_dim = prior.event_dim
        else:  # non-distribution prior (deterministic value): fall back.
            support = entry.constraint
            event_dim = 0
        loc = self.pyrox_param(
            f"{name}_loc",
            entry.init_value,
            constraint=support,
            event_dim=event_dim,
        )
        loc = jnp.asarray(loc)
        return self.pyrox_sample(name, dist.Delta(loc, event_dim=event_dim))

    def _guide_normal(self, name: str, entry: _Entry) -> Any:
        """Mean-field normal guide in unconstrained space.

        When ``entry.constraint`` is non-trivial, the latent site is a
        ``TransformedDistribution`` wrapping ``Normal(loc, scale)`` with
        the constraint's bijection, so guide draws always land in the
        prior's support. ``loc`` is initialized by the inverse transform
        of ``init_value`` so guide and prior agree at step zero.
        """
        init = jnp.asarray(entry.init_value)
        # A vector parameter usually carries a joint prior (e.g.
        # ``LogNormal(...).expand([D]).to_event(1)``). The guide site has to
        # declare the same event rank or Trace_ELBO rejects the model/guide
        # pair, so read it off the prior rather than assuming element-wise.
        # Callable (dependent) priors must be resolved first, exactly as
        # `_guide_delta` and `pyrox_sample` do — reading `event_dim` off the
        # callable itself would silently report rank 0.
        prior = entry.prior
        if callable(prior) and not isinstance(prior, dist.Distribution):
            prior = prior(self)
        prior_event_dim = prior.event_dim if isinstance(prior, dist.Distribution) else 0
        if _is_real_support(entry.constraint):
            loc = self.pyrox_param(f"{name}_loc", init)
            scale = self.pyrox_param(
                f"{name}_scale",
                jnp.ones_like(init) * 0.1,
                constraint=dist.constraints.positive,
            )
            return self.pyrox_sample(
                name, dist.Normal(loc, scale).to_event(prior_event_dim)
            )
        transform = _biject_to(entry.constraint)
        # loc AND scale live in unconstrained space. For shape-changing
        # transforms (e.g. StickBreaking for a simplex: K -> K-1) the
        # unconstrained shape differs from ``init``; sizing scale from
        # ``init`` would fail to broadcast against ``transform.inv(init)``.
        unconstrained = transform.inv(init)
        loc = self.pyrox_param(f"{name}_loc", unconstrained)
        scale = self.pyrox_param(
            f"{name}_scale",
            jnp.full_like(unconstrained, 0.1),
            constraint=dist.constraints.positive,
        )
        # Promote the base to the transform's domain event rank so the
        # TransformedDistribution's event structure matches the constrained
        # support (a no-op for element-wise transforms like Exp/Sigmoid).
        # The base lives in the transform's *domain*, while the prior's
        # event rank is stated in its codomain, so subtract the rank the
        # transform itself adds. For an element-wise transform the delta is
        # zero and this is just the prior's rank; for a shape-changing one
        # (corr_cholesky: vector domain -> matrix codomain) taking the
        # prior's rank directly would over-promote a base that has fewer
        # dimensions than that, and construction would fail.
        delta = transform.codomain.event_dim - transform.domain.event_dim
        event_dim = max(transform.domain.event_dim, prior_event_dim - delta)
        base = dist.Normal(loc, scale).to_event(event_dim)
        return self.pyrox_sample(name, dist.TransformedDistribution(base, transform))