44from jax import numpy as jnp
55from 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
915from oryx .models .ncde .ncde import MLPNeuralCDE
1016from oryx .models .ncde .term import AbstractNCDETerm , MLPNCDETerm
1117from oryx .models .node .node import MLPNeuralODE
@@ -54,7 +60,7 @@ def __call__(self, state):
5460def 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
6672def 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