Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 1 addition & 54 deletions pyrenew/observation/noise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 2 additions & 0 deletions pyrenew/randomvariable/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,12 @@
StaticDistributionalVariable,
)
from pyrenew.randomvariable.transformedvariable import TransformedVariable
from pyrenew.randomvariable.vectorizedrv import VectorizedRV

__all__ = [
"DistributionalVariable",
"StaticDistributionalVariable",
"DynamicDistributionalVariable",
"TransformedVariable",
"VectorizedRV",
]
62 changes: 62 additions & 0 deletions pyrenew/randomvariable/vectorizedrv.py
Original file line number Diff line number Diff line change
@@ -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)