Skip to content

Commit f1d768b

Browse files
committed
Add stochastic base models and flatten model
Update MLPModel to MLP
1 parent a72a6a7 commit f1d768b

5 files changed

Lines changed: 98 additions & 16 deletions

File tree

oryx/models/__init__.py

Lines changed: 14 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,15 +4,21 @@
44
Models take inputs and produce outputs, and may have state.
55
"""
66

7-
from oryx.models.base_model import AbstractModel, AbstractStatefulModel
8-
from oryx.models.mlp import MLPModel
9-
from oryx.models.ncde import (
7+
from .base_model import (
8+
AbstractModel,
9+
AbstractStatefulModel,
10+
AbstractStochasticStatefulModel,
11+
AbstractStochaticModel,
12+
)
13+
from .flatten import Flatten
14+
from .mlp import MLP
15+
from .ncde import (
1016
AbstractNCDETerm,
1117
AbstractNeuralCDE,
1218
MLPNCDETerm,
1319
MLPNeuralCDE,
1420
)
15-
from oryx.models.node import (
21+
from .node import (
1622
AbstractNeuralODE,
1723
AbstractNODETerm,
1824
MLPNeuralODE,
@@ -22,6 +28,9 @@
2228
__all__ = [
2329
"AbstractModel",
2430
"AbstractStatefulModel",
31+
"AbstractStochasticStatefulModel",
32+
"AbstractStochaticModel",
33+
"Flatten",
2534
"AbstractNeuralODE",
2635
"AbstractNODETerm",
2736
"MLPNeuralODE",
@@ -30,5 +39,5 @@
3039
"AbstractNCDETerm",
3140
"MLPNeuralCDE",
3241
"MLPNCDETerm",
33-
"MLPModel",
42+
"MLP",
3443
]

oryx/models/base_model.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from typing import Concatenate
33

44
import equinox as eqx
5+
from jaxtyping import Key
56

67

78
class AbstractModel[**InType, OutType](eqx.Module, strict=True):
@@ -25,3 +26,41 @@ def __call__(
2526
self, state: eqx.nn.State, *args: InType.args, **kwargs: InType.kwargs
2627
) -> tuple[eqx.nn.State, *OutType]:
2728
"""Return an output given inputs and the state."""
29+
30+
31+
class AbstractStochaticModel[**InType, OutType](
32+
AbstractModel[
33+
Concatenate[Key, InType],
34+
OutType,
35+
],
36+
strict=True,
37+
):
38+
"""Base class for stochastic models that take inputs and produce outputs."""
39+
40+
# TODO: Revisit this if Python ever adds support for concatenating keyword
41+
# parameters, I'd rather key as a keyword to fit the style of the library
42+
@abstractmethod
43+
def __call__(
44+
self, key: Key, *args: InType.args, **kwargs: InType.kwargs
45+
) -> OutType:
46+
"""Return an output given an input."""
47+
48+
49+
class AbstractStochasticStatefulModel[**InType, *OutType](
50+
AbstractStatefulModel[
51+
Concatenate[Key, InType],
52+
*OutType,
53+
],
54+
AbstractStochaticModel[
55+
Concatenate[eqx.nn.State, InType],
56+
tuple[eqx.nn.State, *OutType],
57+
],
58+
strict=True,
59+
):
60+
"""Base class for stochastic models with state."""
61+
62+
@abstractmethod
63+
def __call__(
64+
self, state: eqx.nn.State, key: Key, *args: InType.args, **kwargs: InType.kwargs
65+
) -> tuple[eqx.nn.State, *OutType]:
66+
"""Return an output given inputs and the state."""

oryx/models/flatten.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,14 @@
1+
from jax import numpy as jnp
2+
from jaxtyping import Array, Float
3+
4+
from .base_model import AbstractModel
5+
6+
7+
class Flatten(
8+
AbstractModel[[Float[Array, " ..."]], Float[Array, " out_size"]], strict=True
9+
):
10+
"""Warps the JAX `jnp.ravel` function to flatten an input array."""
11+
12+
def __call__(self, x: Float[Array, " ..."]) -> Float[Array, " out_size"]:
13+
"""Flatten the input array into a 1D array."""
14+
return jnp.ravel(x)

oryx/models/mlp.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from .base_model import AbstractModel
88

99

10-
class MLPModel(
10+
class MLP(
1111
AbstractModel[[Float[Array, " in_size"]], Float[Array, " out_size"]], strict=True
1212
):
1313
"""Wrapper around eqx.nn.MLP."""

tests/test_models.py

Lines changed: 30 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,14 @@
44
from jax import numpy as jnp
55
from jax import random as jr
66

7-
from oryx.models.base_model import AbstractModel, AbstractStatefulModel
8-
from oryx.models.mlp import MLPModel
7+
from oryx.models.base_model import (
8+
AbstractModel,
9+
AbstractStatefulModel,
10+
AbstractStochasticStatefulModel,
11+
AbstractStochaticModel,
12+
)
13+
from oryx.models.flatten import Flatten
14+
from oryx.models.mlp import MLP
915
from oryx.models.ncde.ncde import MLPNeuralCDE
1016
from oryx.models.ncde.term import AbstractNCDETerm, MLPNCDETerm
1117
from oryx.models.node.node import MLPNeuralODE
@@ -54,7 +60,7 @@ def __call__(self, state):
5460
def test_mlpmodel_forward_shape_and_determinism():
5561
key = jr.key(0)
5662
in_size, out_size = 7, 4
57-
m = MLPModel(in_size=in_size, out_size=out_size, width_size=10, depth=2, key=key)
63+
m = MLP(in_size=in_size, out_size=out_size, width_size=10, depth=2, key=key)
5864
x = jnp.zeros(in_size)
5965
y1 = m(x)
6066
y2 = m(x)
@@ -65,14 +71,14 @@ def test_mlpmodel_forward_shape_and_determinism():
6571

6672
def test_mlpmodel_wrong_input_shape_raises():
6773
key = jr.key(1)
68-
m = MLPModel(in_size=5, out_size=3, width_size=8, depth=1, key=key)
74+
m = MLP(in_size=5, out_size=3, width_size=8, depth=1, key=key)
6975
x_bad = jnp.ones((6,))
7076
with pytest.raises(Exception):
7177
_ = m(x_bad)
7278

7379

7480
@pytest.mark.parametrize("time_in_input", [True, False])
75-
def test_mlpneuralode_solve_and_call(time_in_input):
81+
def test_mlpneuralode_solve_and_call(time_in_input: bool):
7682
key = jr.key(4)
7783
in_size, out_size, latent_size = 3, 2, 4
7884
model = MLPNeuralODE(
@@ -138,7 +144,7 @@ class NoCallNCDETerm(AbstractNCDETerm):
138144

139145

140146
@pytest.mark.parametrize("add_time", [True, False])
141-
def test_mlpncdeterm_output_shape(add_time):
147+
def test_mlpncdeterm_output_shape(add_time: bool):
142148
key = jr.key(6)
143149
input_size, data_size = 3, 5
144150
term = MLPNCDETerm(
@@ -157,7 +163,7 @@ def test_mlpncdeterm_output_shape(add_time):
157163

158164

159165
@pytest.mark.parametrize("inference", [True, False])
160-
def test_mlpneuralcde(inference):
166+
def test_mlpneuralcde(inference: bool):
161167
key = jr.key(7)
162168
in_size, out_size, latent_size, state_size = 4, 2, 6, 4
163169
model, state = eqx.nn.make_with_state(MLPNeuralCDE)(
@@ -191,7 +197,7 @@ def test_mlpneuralcde(inference):
191197

192198

193199
@pytest.mark.parametrize("time_in_input", [True, False])
194-
def test_ncde_z0_and_coeffs_shapes(time_in_input):
200+
def test_ncde_z0_and_coeffs_shapes(time_in_input: bool):
195201
key = jr.key(0)
196202
in_size, latent_size = 3, 5
197203
model = MLPNeuralCDE(
@@ -219,7 +225,7 @@ def test_ncde_z0_and_coeffs_shapes(time_in_input):
219225

220226

221227
@pytest.mark.parametrize("inference", [True, False])
222-
def test_ncde_solve_output_shape(inference):
228+
def test_ncde_solve_output_shape(inference: bool):
223229
key = jr.key(1)
224230
in_size, latent_size = 2, 4
225231
model = MLPNeuralCDE(
@@ -242,7 +248,7 @@ def test_ncde_solve_output_shape(inference):
242248

243249

244250
@pytest.mark.parametrize("inference", [True, False])
245-
def test_ncde_state_rollover(inference):
251+
def test_ncde_state_rollover(inference: bool):
246252
key = jr.key(2)
247253
in_size, out_size, latent_size = 2, 1, 3
248254
model, state = eqx.nn.make_with_state(MLPNeuralCDE)(
@@ -266,3 +272,17 @@ def test_ncde_state_rollover(inference):
266272
assert not jnp.isnan(ts).any()
267273
assert jnp.allclose(ts, jnp.array([1.0, 2.0, 3.0]))
268274
assert model.t1(state) == 3.0
275+
276+
277+
@pytest.mark.parametrize("dims", range(4))
278+
def test_flatten_model(dims: int):
279+
model = Flatten()
280+
281+
input_shape = (2,) * dims
282+
input_data = jnp.ones(input_shape)
283+
284+
output = model(input_data)
285+
286+
assert isinstance(output, jnp.ndarray)
287+
assert output.shape == (2**dims,)
288+
assert jnp.array_equal(output, jnp.ones((2**dims,)))

0 commit comments

Comments
 (0)