diff --git a/pyrenew/observation/noise.py b/pyrenew/observation/noise.py index 0733078d..f9352741 100644 --- a/pyrenew/observation/noise.py +++ b/pyrenew/observation/noise.py @@ -32,64 +32,11 @@ from jax.typing import ArrayLike from pyrenew.metaclass import RandomVariable +from pyrenew.randomvariable import VectorizedRV _EPSILON = 1e-10 -class VectorizedRV(RandomVariable): - """ - Wrapper that adds n_groups support to simple RandomVariables. - - Uses numpyro.plate to vectorize sampling, enabling simple RVs - to work with noise models expecting the group-level interface. - - Parameters - ---------- - name - A name for this random variable. - The numpyro plate is named ``f"{name}_plate"``. - rv - The underlying RandomVariable to wrap. - """ - - def __init__(self, name: str, rv: RandomVariable) -> None: - """ - Initialize VectorizedRV wrapper. - - Parameters - ---------- - name - A name for this random variable. - The numpyro plate is named ``f"{name}_plate"``. - rv - The underlying RandomVariable to wrap. - """ - super().__init__(name=name) - self.rv = rv - self.plate_name = f"{name}_plate" - - def validate(self) -> None: # pragma: no cover - """Validate the underlying RV.""" - self.rv.validate() - - def sample(self, n_groups: int, **kwargs: object) -> ArrayLike: - """ - Sample n_groups values using numpyro.plate. - - Parameters - ---------- - n_groups - Number of group-level values to sample. - - Returns - ------- - ArrayLike - Array of shape (n_groups,). - """ - with numpyro.plate(self.plate_name, n_groups): - return self.rv(**kwargs) - - class CountNoise(ABC): """ Abstract base for count observation noise models. diff --git a/pyrenew/randomvariable/__init__.py b/pyrenew/randomvariable/__init__.py index c599d101..54c1ced1 100644 --- a/pyrenew/randomvariable/__init__.py +++ b/pyrenew/randomvariable/__init__.py @@ -6,10 +6,12 @@ StaticDistributionalVariable, ) from pyrenew.randomvariable.transformedvariable import TransformedVariable +from pyrenew.randomvariable.vectorizedrv import VectorizedRV __all__ = [ "DistributionalVariable", "StaticDistributionalVariable", "DynamicDistributionalVariable", "TransformedVariable", + "VectorizedRV", ] diff --git a/pyrenew/randomvariable/vectorizedrv.py b/pyrenew/randomvariable/vectorizedrv.py new file mode 100644 index 00000000..c77bf40e --- /dev/null +++ b/pyrenew/randomvariable/vectorizedrv.py @@ -0,0 +1,62 @@ +# numpydoc ignore=GL08 + +from __future__ import annotations + +import numpyro +from jax.typing import ArrayLike + +from pyrenew.metaclass import RandomVariable + + +class VectorizedRV(RandomVariable): + """ + Wrapper that adds n_groups support to simple RandomVariables. + + Uses numpyro.plate to vectorize sampling, enabling simple RVs + to work with noise models expecting the group-level interface. + + Parameters + ---------- + name + A name for this random variable. + The numpyro plate is named ``f"{name}_plate"``. + rv + The underlying RandomVariable to wrap. + """ + + def __init__(self, name: str, rv: RandomVariable) -> None: + """ + Initialize VectorizedRV wrapper. + + Parameters + ---------- + name + A name for this random variable. + The numpyro plate is named ``f"{name}_plate"``. + rv + The underlying RandomVariable to wrap. + """ + super().__init__(name=name) + self.rv = rv + self.plate_name = f"{name}_plate" + + def validate(self) -> None: # pragma: no cover + """Validate the underlying RV.""" + self.rv.validate() + + def sample(self, n_groups: int, **kwargs: object) -> ArrayLike: + """ + Sample n_groups values using numpyro.plate. + + Parameters + ---------- + n_groups + Number of group-level values to sample. + + Returns + ------- + ArrayLike + Array of shape (n_groups,). + """ + with numpyro.plate(self.plate_name, n_groups): + return self.rv(**kwargs)