From 650cd422f0485587ba55188a84491fb23acfd5fc Mon Sep 17 00:00:00 2001 From: Evgenii Zheltonozhskii Date: Mon, 22 Apr 2024 23:43:12 +0100 Subject: [PATCH 1/4] Fix complex tests --- diffrax/_brownian/tree.py | 12 +++--- diffrax/_global_interpolation.py | 48 ++++++++++++++--------- diffrax/_local_interpolation.py | 25 +++++++----- diffrax/_root_finder/_verychord.py | 10 +++-- diffrax/_solver/dopri8.py | 3 +- diffrax/_solver/kencarp3.py | 3 +- diffrax/_solver/milstein.py | 11 ++++-- diffrax/_step_size_controller/adaptive.py | 2 +- test/helpers.py | 2 + test/test_global_interpolation.py | 13 +++--- test/test_integrate.py | 25 +++++++----- test/test_interpolation.py | 17 +++++--- test/test_vmap.py | 7 ++-- 13 files changed, 113 insertions(+), 65 deletions(-) diff --git a/diffrax/_brownian/tree.py b/diffrax/_brownian/tree.py index 88c10524..813384c4 100644 --- a/diffrax/_brownian/tree.py +++ b/diffrax/_brownian/tree.py @@ -10,7 +10,8 @@ import jax.random as jr import jax.tree_util as jtu import lineax.internal as lxi -from jaxtyping import Array, Float, PRNGKeyArray, PyTree +from jaxtyping import Array, Inexact, PRNGKeyArray, PyTree +from lineax.internal import complex_to_real_dtype from .._custom_types import ( AbstractBrownianIncrement, @@ -54,9 +55,9 @@ # For the midpoint rule for generating space-time Levy area see Theorem 6.1.6. # For the general interpolation rule for space-time Levy area see Theorem 6.1.4. -FloatDouble: TypeAlias = tuple[Float[Array, " *shape"], Float[Array, " *shape"]] +FloatDouble: TypeAlias = tuple[Inexact[Array, " *shape"], Inexact[Array, " *shape"]] FloatTriple: TypeAlias = tuple[ - Float[Array, " *shape"], Float[Array, " *shape"], Float[Array, " *shape"] + Inexact[Array, " *shape"], Inexact[Array, " *shape"], Inexact[Array, " *shape"] ] _Spline: TypeAlias = Literal["sqrt", "quad", "zero"] _BrownianReturn = TypeVar("_BrownianReturn", bound=AbstractBrownianIncrement) @@ -283,9 +284,10 @@ def _evaluate_leaf( tuple[RealScalarLike, Array], tuple[RealScalarLike, Array, Array, Array] ]: shape, dtype = struct.shape, struct.dtype + tdtype = complex_to_real_dtype(dtype) - t0 = jnp.zeros((), dtype) - r = jnp.asarray(r, dtype) + t0 = jnp.zeros((), tdtype) + r = jnp.asarray(r, tdtype) if self.levy_area is SpaceTimeLevyArea: state_key, init_key_w, init_key_la = jr.split(key, 3) diff --git a/diffrax/_global_interpolation.py b/diffrax/_global_interpolation.py index 31d8b038..1964a1f3 100644 --- a/diffrax/_global_interpolation.py +++ b/diffrax/_global_interpolation.py @@ -143,7 +143,10 @@ def _index(_ys): next_t = self.ts[index + 1] diff_t = next_t - prev_t - return (prev_ys**ω + (next_ys**ω - prev_ys**ω) * (fractional_part / diff_t)).ω + with jax.numpy_dtype_promotion("standard"): + return ( + prev_ys**ω + (next_ys**ω - prev_ys**ω) * (fractional_part / diff_t) + ).ω @eqx.filter_jit def derivative(self, t: RealScalarLike, left: bool = True) -> PyTree[Array]: @@ -165,10 +168,11 @@ def derivative(self, t: RealScalarLike, left: bool = True) -> PyTree[Array]: index, _ = self._interpret_t(t, left) - return ( - (ω(self.ys)[index + 1] - ω(self.ys)[index]) - / (self.ts[index + 1] - self.ts[index]) - ).ω + with jax.numpy_dtype_promotion("standard"): + return ( + (ω(self.ys)[index + 1] - ω(self.ys)[index]) + / (self.ts[index + 1] - self.ts[index]) + ).ω LinearInterpolation.__init__.__doc__ = """**Arguments:** @@ -254,10 +258,11 @@ def evaluate( d, c, b, a = self.coeffs - return ( - ω(a)[index] - + frac * (ω(b)[index] + frac * (ω(c)[index] + frac * ω(d)[index])) - ).ω + with jax.numpy_dtype_promotion("standard"): + return ( + ω(a)[index] + + frac * (ω(b)[index] + frac * (ω(c)[index] + frac * ω(d)[index])) + ).ω @eqx.filter_jit def derivative( @@ -283,7 +288,8 @@ def derivative( d, c, b, _ = self.coeffs - return (ω(b)[index] + frac * (2 * ω(c)[index] + frac * 3 * ω(d)[index])).ω + with jax.numpy_dtype_promotion("standard"): + return (ω(b)[index] + frac * (2 * ω(c)[index] + frac * 3 * ω(d)[index])).ω CubicInterpolation.__init__.__doc__ = """**Arguments:** @@ -622,8 +628,9 @@ def _hermite_forward( ]: prev_ti, prev_yi, prev_deriv_i = carry ti, yi, next_ti, next_yi = value - first_deriv_i = (next_yi - yi) / (next_ti - ti) - later_deriv_i = (yi - prev_yi) / (ti - prev_ti) + with jax.numpy_dtype_promotion("standard"): + first_deriv_i = (next_yi - yi) / (next_ti - ti) + later_deriv_i = (yi - prev_yi) / (ti - prev_ti) deriv_i = jnp.where(jnp.isnan(prev_yi), first_deriv_i, later_deriv_i) cond = jnp.isnan(yi) carry_ti = jnp.where(cond, prev_ti, ti) @@ -635,13 +642,15 @@ def _hermite_forward( def _hermite_coeffs(t0, y0, deriv0, t1, y1): t_diff = t1 - t0 - deriv1 = (y1 - y0) / t_diff - d_deriv = deriv1 - deriv0 - a = y0 - b = deriv0 - c = 2 * d_deriv / t_diff - d = -d_deriv / t_diff**2 + with jax.numpy_dtype_promotion("standard"): + deriv1 = (y1 - y0) / t_diff + d_deriv = deriv1 - deriv0 + + a = y0 + b = deriv0 + c = 2 * d_deriv / t_diff + d = -d_deriv / (t_diff**2) return d, c, b, a @@ -684,7 +693,8 @@ def _backward_hermite_coefficients( else: y0 = jnp.broadcast_to(replace_nans_at_start, ys[0].shape) if deriv0 is None: - deriv0 = (next_ys[0] - y0) / (next_ts[0] - t0) + with jax.numpy_dtype_promotion("standard"): + deriv0 = (next_ys[0] - y0) / (next_ts[0] - t0) else: deriv0 = jnp.broadcast_to(deriv0, ys[0].shape) ts = ts[:-1] diff --git a/diffrax/_local_interpolation.py b/diffrax/_local_interpolation.py index 97ccdb15..0098a059 100644 --- a/diffrax/_local_interpolation.py +++ b/diffrax/_local_interpolation.py @@ -1,6 +1,7 @@ from collections.abc import Callable from typing import cast, Optional, TYPE_CHECKING +import jax import jax.numpy as jnp import jax.tree_util as jtu import numpy as np @@ -35,12 +36,15 @@ def evaluate( self, t0: RealScalarLike, t1: Optional[RealScalarLike] = None, left: bool = True ) -> PyTree[Array]: del left - if t1 is None: - coeff = linear_rescale(self.t0, t0, self.t1) - return (self.y0**ω + coeff * (self.y1**ω - self.y0**ω)).call(jnp.asarray).ω - else: - coeff = (t1 - t0) / (self.t1 - self.t0) - return (coeff * (self.y1**ω - self.y0**ω)).call(jnp.asarray).ω + with jax.numpy_dtype_promotion("standard"): + if t1 is None: + coeff = linear_rescale(self.t0, t0, self.t1) + return ( + (self.y0**ω + coeff * (self.y1**ω - self.y0**ω)).call(jnp.asarray).ω + ) + else: + coeff = (t1 - t0) / (self.t1 - self.t0) + return (coeff * (self.y1**ω - self.y0**ω)).call(jnp.asarray).ω class ThirdOrderHermitePolynomialInterpolation(AbstractLocalInterpolation): @@ -82,7 +86,8 @@ def evaluate( t = linear_rescale(self.t0, t0, self.t1) def _eval(_coeffs): - return jnp.polyval(_coeffs, t) + with jax.numpy_dtype_promotion("standard"): + return jnp.polyval(_coeffs, t) return jtu.tree_map(_eval, self.coeffs) @@ -104,7 +109,8 @@ def __init__( k: PyTree[Shaped[Array, "order ?*y"], "Y"], ): def _calculate(_y0, _y1, _k): - _ymid = _y0 + jnp.tensordot(self.c_mid, _k, axes=1) + with jax.numpy_dtype_promotion("standard"): + _ymid = _y0 + jnp.tensordot(self.c_mid, _k, axes=1) _f0 = _k[0] _f1 = _k[-1] # TODO: rewrite as matrix-vector product? @@ -127,6 +133,7 @@ def evaluate( t = linear_rescale(self.t0, t0, self.t1) def _eval(_coeffs): - return jnp.polyval(_coeffs, t) + with jax.numpy_dtype_promotion("standard"): + return jnp.polyval(_coeffs, t) return jtu.tree_map(_eval, self.coeffs) diff --git a/diffrax/_root_finder/_verychord.py b/diffrax/_root_finder/_verychord.py index 89300108..be1b28b9 100644 --- a/diffrax/_root_finder/_verychord.py +++ b/diffrax/_root_finder/_verychord.py @@ -11,6 +11,7 @@ import optimistix as optx from equinox.internal import ω from jaxtyping import Array, Bool, PyTree, Scalar +from lineax.internal import complex_to_real_dtype from .._custom_types import Y @@ -97,11 +98,12 @@ def init( y_dtype = lxi.default_floating_dtype() else: y_dtype = jnp.result_type(*y_leaves) + diff_dtype = complex_to_real_dtype(y_dtype) init_state = _VeryChordState( linear_state=linear_state, diff=jtu.tree_map(lambda x: jnp.full(x.shape, jnp.inf, x.dtype), y), - diffsize=jnp.array(jnp.inf, dtype=y_dtype), - diffsize_prev=jnp.array(1.0, dtype=y_dtype), + diffsize=jnp.array(jnp.inf, dtype=diff_dtype), + diffsize_prev=jnp.array(1.0, dtype=diff_dtype), result=optx.RESULTS.successful, step=jnp.array(0), ) @@ -127,7 +129,9 @@ def step( ) diff = sol.value new_y = (y**ω - diff**ω).ω - scale = (self.atol + self.rtol * ω(new_y).call(jnp.abs)).ω + + with jax.numpy_dtype_promotion("standard"): + scale = (self.atol + self.rtol * ω(new_y).call(jnp.abs)).ω diffsize = self.norm((diff**ω / scale**ω).ω) new_state = _VeryChordState( linear_state=state.linear_state, diff --git a/diffrax/_solver/dopri8.py b/diffrax/_solver/dopri8.py index ab9ef64a..428320c7 100644 --- a/diffrax/_solver/dopri8.py +++ b/diffrax/_solver/dopri8.py @@ -298,7 +298,8 @@ def evaluate( return self.evaluate(t1) - self.evaluate(t0) t = linear_rescale(self.t0, t0, self.t1) coeffs = _vmap_polyval(jnp.asarray(self.eval_coeffs, dtype=t.dtype), t) * t - return (self.y0**ω + vector_tree_dot(coeffs, self.k) ** ω).ω + with jax.numpy_dtype_promotion("standard"): + return (self.y0**ω + vector_tree_dot(coeffs, self.k) ** ω).ω class Dopri8(AbstractERK): diff --git a/diffrax/_solver/kencarp3.py b/diffrax/_solver/kencarp3.py index d24546fd..f7538f4e 100644 --- a/diffrax/_solver/kencarp3.py +++ b/diffrax/_solver/kencarp3.py @@ -117,7 +117,8 @@ def evaluate( explicit_k, implicit_k = self.k k = (explicit_k**ω + implicit_k**ω).ω coeffs = t * jax.vmap(lambda row: jnp.polyval(row, t))(self.coeffs) - return (self.y0**ω + vector_tree_dot(coeffs, k) ** ω).ω + with jax.numpy_dtype_promotion("standard"): + return (self.y0**ω + vector_tree_dot(coeffs, k) ** ω).ω class _KenCarp3Interpolation(KenCarpInterpolation): diff --git a/diffrax/_solver/milstein.py b/diffrax/_solver/milstein.py index 8bc1ea5b..3d14e343 100644 --- a/diffrax/_solver/milstein.py +++ b/diffrax/_solver/milstein.py @@ -214,7 +214,8 @@ def step( leaf = jnp.tensordot(l1[..., None], l2[None, ...], axes=1) if i1 == i2: eye = jnp.eye(l1.size).reshape(l1.shape + l1.shape) - leaf = leaf - Δt * eye + with jax.numpy_dtype_promotion("standard"): + leaf = leaf - Δt * eye leaves_ΔwΔw.append(leaf) tree_ΔwΔw = tree_Δw.compose(tree_Δw) ΔwΔw = jtu.tree_unflatten(tree_ΔwΔw, leaves_ΔwΔw) @@ -236,7 +237,9 @@ def _to_vmap(_g0): # _g0 has structure (tree(y0), leaf(y0)) _, _jvp = jax.jvp(_to_vjp, (y0,), (_g0,)) # jvp has structure (tree(g0), leaf(g0)) - _jvp_matrix = jax.jacfwd(lambda _Δw: diffusion.prod(_jvp, _Δw))(Δw) + _jvp_matrix = jax.jacfwd( + lambda _Δw: diffusion.prod(_jvp, _Δw), holomorphic=jnp.iscomplexobj(Δw) + )(Δw) # _jvp_matrix has structure (tree(y0), tree(Δw), leaf(y0), leaf(Δw)) return _jvp_matrix @@ -282,7 +285,9 @@ def _to_treemap(_Δw, _g0): Δw_treedef = jtu.tree_structure(Δw) # g0 has structure (tree(g0), leaf(g0)) # Which we now transform into its isomorphic matrix form, as above. - g0_matrix = jax.jacfwd(lambda _Δw: diffusion.prod(g0, _Δw))(Δw) + g0_matrix = jax.jacfwd( + lambda _Δw: diffusion.prod(g0, _Δw), holomorphic=jnp.iscomplexobj(Δw) + )(Δw) # g0_matrix has structure (tree(y0), tree(Δw), leaf(y0), leaf(Δw)) g0_matrix = jtu.tree_transpose(y_treedef, Δw_treedef, g0_matrix) # g0_matrix has structure (tree(Δw), tree(y0), leaf(y0), leaf(Δw)) diff --git a/diffrax/_step_size_controller/adaptive.py b/diffrax/_step_size_controller/adaptive.py index 343b5920..20eadf92 100644 --- a/diffrax/_step_size_controller/adaptive.py +++ b/diffrax/_step_size_controller/adaptive.py @@ -611,7 +611,7 @@ def _scale(_y0, _y1_candidate, _y_error): factor = lax.stop_gradient(factor) factor = eqxi.nondifferentiable(factor) with jax.numpy_dtype_promotion("standard"): - dt = prev_dt * factor.astype(prev_dt) + dt = prev_dt * factor # E.g. we failed an implicit step, so y_error=inf, so inv_scaled_error=0, # so factor=factormin, and we shrunk our step. diff --git a/test/helpers.py b/test/helpers.py index ed20ac46..ab9e37af 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -18,6 +18,7 @@ ) from jax import Array from jaxtyping import PRNGKeyArray, PyTree, Shaped +from lineax.internal import complex_to_real_dtype all_ode_solvers = ( @@ -251,6 +252,7 @@ def sde_solver_strong_order( bm_tol, saveat, ) + dts = 2.0 ** jnp.arange(-3, -3 - num_levels, -1, dtype=dtype) errs_list, steps_list = [], [] for level in range(level_coarse, level_fine + 1): diff --git a/test/test_global_interpolation.py b/test/test_global_interpolation.py index 5103000f..97f1938b 100644 --- a/test/test_global_interpolation.py +++ b/test/test_global_interpolation.py @@ -250,14 +250,15 @@ def test_cubic_interpolation_deriv0(unsqueeze): @pytest.mark.parametrize("mode", ["linear", "cubic"]) -def test_interpolation_classes(mode, getkey): +@pytest.mark.parametrize("dtype", [jnp.float64, jnp.complex128]) +def test_interpolation_classes(mode, dtype, getkey): length = 8 num_channels = 3 ts_ = [ jnp.linspace(0, 10, length), jnp.array([0.0, 2.0, 3.0, 3.1, 4.0, 4.1, 5.0, 5.1]), ] - _make = lambda: jr.normal(getkey(), (length, num_channels)) + _make = lambda: jr.normal(getkey(), (length, num_channels), dtype=dtype) ys_ = [_make(), [_make(), {"a": _make(), "b": _make()}], {}, []] for ts in ts_: assert len(ts) == length @@ -288,9 +289,9 @@ def test_interpolation_classes(mode, getkey): def _test(firstval, vals, y0, y1): vals = jnp.concatenate([firstval[None], vals]) with jax.numpy_rank_promotion("allow"): - true_vals = y0 + ((points - t0) / (t1 - t0))[:, None] * ( - y1 - y0 - ) + true_vals = y0 + ((points - t0) / (t1 - t0)).astype( + y0.dtype + )[:, None] * (y1 - y0) assert tree_allclose(vals, true_vals) jtu.tree_map(_test, firstval, vals, y0, y1) @@ -299,7 +300,7 @@ def _test(firstval, vals, y0, y1): def _test2(firstderiv, derivs, y0, y1): derivs = jnp.concatenate([firstderiv[None], derivs]) - true_derivs = (y1 - y0) / (t1 - t0) + true_derivs = (y1 - y0) / (t1 - t0).astype(y0.dtype) true_derivs = jnp.broadcast_to(true_derivs, derivs.shape) assert tree_allclose(derivs, true_derivs) diff --git a/test/test_integrate.py b/test/test_integrate.py index b2611867..cdbd3fab 100644 --- a/test/test_integrate.py +++ b/test/test_integrate.py @@ -144,13 +144,16 @@ def f(t, y, args): @pytest.mark.parametrize("solver", all_ode_solvers + all_split_solvers) -def test_ode_order(solver): +@pytest.mark.parametrize("dtype", [jnp.float64, jnp.complex128]) +def test_ode_order(solver, dtype): solver = implicit_tol(solver) key = jr.PRNGKey(5678) akey, ykey = jr.split(key, 2) - A = jr.normal(akey, (10, 10), dtype=jnp.float64) * 0.5 + A = jr.normal(akey, (10, 10), dtype=dtype) * 0.5 + if jnp.iscomplexobj(A) and isinstance(solver, diffrax.AbstractImplicitSolver): + return if ( solver.term_structure == diffrax.MultiTerm[tuple[diffrax.AbstractTerm, diffrax.AbstractTerm]] @@ -171,7 +174,7 @@ def f(t, y, args): term = diffrax.ODETerm(f) t0 = 0 t1 = 4 - y0 = jr.normal(ykey, (10,), dtype=jnp.float64) + y0 = jr.normal(ykey, (10,), dtype=dtype) true_yT = jax.scipy.linalg.expm((t1 - t0) * A) @ y0 exponents = [] @@ -223,16 +226,19 @@ def _squareplus(x): def _drift(t, y, args): drift_mlp, _, _ = args - return 0.5 * drift_mlp(y) + with jax.numpy_dtype_promotion("standard"): + return 0.5 * drift_mlp(y) def _diffusion(t, y, args): _, diffusion_mlp, noise_dim = args - return 0.25 * diffusion_mlp(y).reshape(3, noise_dim) + with jax.numpy_dtype_promotion("standard"): + return 0.25 * diffusion_mlp(y).reshape(3, noise_dim) @pytest.mark.parametrize("solver_ctr,noise,theoretical_order", _solvers_and_orders()) -def test_sde_strong_order(solver_ctr, noise, theoretical_order): +@pytest.mark.parametrize("dtype", [jnp.float64]) +def test_sde_strong_order(solver_ctr, noise, theoretical_order, dtype): key = jr.PRNGKey(5678) driftkey, diffusionkey, ykey, bmkey = jr.split(key, 4) num_samples = 20 @@ -270,7 +276,7 @@ def test_sde_strong_order(solver_ctr, noise, theoretical_order): t0 = 0.0 t1 = 2.0 - y0 = jr.normal(ykey, (3,), dtype=jnp.float64) + y0 = jr.normal(ykey, (3,), dtype=dtype) def get_terms(bm): return MultiTerm(ODETerm(_drift), ControlTerm(_diffusion, bm)) @@ -333,9 +339,10 @@ def get_dt_and_controller(level): diffrax.SaveAt(dense=True), ), ) -def test_reverse_time(solver_ctr, dt0, saveat, getkey): +@pytest.mark.parametrize("dtype", [jnp.float64, jnp.complex128]) +def test_reverse_time(solver_ctr, dt0, saveat, dtype, getkey): key = getkey() - y0 = jr.normal(key, (2, 2)) + y0 = jr.normal(key, (2, 2), dtype=dtype) stepsize_controller = ( diffrax.PIDController(rtol=1e-3, atol=1e-6) if dt0 is None diff --git a/test/test_interpolation.py b/test/test_interpolation.py index 1a677f27..9a085a9f 100644 --- a/test/test_interpolation.py +++ b/test/test_interpolation.py @@ -2,6 +2,7 @@ import jax import jax.numpy as jnp import jax.random as jr +import pytest from .helpers import all_ode_solvers, all_split_solvers, implicit_tol, tree_allclose @@ -19,9 +20,10 @@ def _test_path_endpoints(path, name, y0, y1): assert tree_allclose(y1, path.evaluate(path.t1)) -def test_derivative(getkey): +@pytest.mark.parametrize("dtype", (jnp.float64, jnp.complex128)) +def test_derivative(dtype, getkey): ts = jnp.linspace(0, 5, 8) - ys = jr.normal(getkey(), (8, 4)) + ys = jr.normal(getkey(), (8, 4), dtype=dtype) paths = [] @@ -34,7 +36,7 @@ def test_derivative(getkey): cubic_interp = diffrax.CubicInterpolation(ts=ts, coeffs=cubic_coeffs) paths.append((cubic_interp, "cubic", ys[0], ys[-1])) - y0 = jr.normal(getkey(), (3,)) + y0 = jr.normal(getkey(), (3,), dtype=dtype) dense_interp = diffrax.diffeqsolve( diffrax.ODETerm(lambda t, y, p: -y), diffrax.Euler(), @@ -56,7 +58,10 @@ def test_derivative(getkey): for solver in all_ode_solvers: solver = implicit_tol(solver) - y0 = jr.normal(getkey(), (3,)) + y0 = jr.normal(getkey(), (3,), dtype=dtype) + + if jnp.iscomplexobj(y0) and isinstance(solver, diffrax.AbstractImplicitSolver): + continue solution = diffrax.diffeqsolve( diffrax.ODETerm(lambda t, y, p: -y), solver, @@ -71,7 +76,9 @@ def test_derivative(getkey): for solver in all_split_solvers: solver = implicit_tol(solver) - y0 = jr.normal(getkey(), (3,)) + y0 = jr.normal(getkey(), (3,), dtype=dtype) + if jnp.iscomplexobj(y0) and isinstance(solver, diffrax.AbstractImplicitSolver): + continue solution = diffrax.diffeqsolve( diffrax.MultiTerm( diffrax.ODETerm(lambda t, y, p: -0.7 * y), diff --git a/test/test_vmap.py b/test/test_vmap.py index 7c5ff90c..b6379df9 100644 --- a/test/test_vmap.py +++ b/test/test_vmap.py @@ -10,15 +10,16 @@ "stepsize_controller", (diffrax.ConstantStepSize(), diffrax.PIDController(rtol=1e-3, atol=1e-6)), ) -def test_vmap_y0(stepsize_controller): +@pytest.mark.parametrize("dtype", (jnp.float64, jnp.complex128)) +def test_vmap_y0(stepsize_controller, dtype): t0 = 0 t1 = 1 dt0 = 0.1 key = jr.PRNGKey(5678) - y0 = jr.normal(key, (10, 2)) - a = jnp.array([[-0.2, 1], [1, -0.2]]) + y0 = jr.normal(key, (10, 2), dtype=dtype) + a = jnp.array([[-0.2, 1], [1, -0.2]], dtype=dtype) def f(t, y, args): return a @ y From 433cc1b95a4e32897a9778306b7fe2ffdf605b43 Mon Sep 17 00:00:00 2001 From: Evgenii Zheltonozhskii Date: Tue, 23 Apr 2024 00:19:57 +0100 Subject: [PATCH 2/4] Fix more complex tests --- diffrax/_root_finder/_verychord.py | 2 +- diffrax/_solver/implicit_euler.py | 18 ++++++++++-------- diffrax/_step_size_controller/adaptive.py | 3 +-- test/test_solver.py | 23 +++++++++++++++-------- 4 files changed, 27 insertions(+), 19 deletions(-) diff --git a/diffrax/_root_finder/_verychord.py b/diffrax/_root_finder/_verychord.py index be1b28b9..921a6dd4 100644 --- a/diffrax/_root_finder/_verychord.py +++ b/diffrax/_root_finder/_verychord.py @@ -132,7 +132,7 @@ def step( with jax.numpy_dtype_promotion("standard"): scale = (self.atol + self.rtol * ω(new_y).call(jnp.abs)).ω - diffsize = self.norm((diff**ω / scale**ω).ω) + diffsize = self.norm((diff**ω / scale**ω).ω) new_state = _VeryChordState( linear_state=state.linear_state, diff=diff, diff --git a/diffrax/_solver/implicit_euler.py b/diffrax/_solver/implicit_euler.py index eb3bdb00..6d9c79bd 100644 --- a/diffrax/_solver/implicit_euler.py +++ b/diffrax/_solver/implicit_euler.py @@ -2,6 +2,7 @@ from typing import ClassVar from typing_extensions import TypeAlias +import jax import optimistix as optx from equinox.internal import ω @@ -81,14 +82,15 @@ def step( # write out a `ButcherTableau` and use `AbstractSDIRK`. k0 = terms.vf_prod(t0, y0, args, control) args = (terms.vf_prod, t1, y0, args, control) - nonlinear_sol = optx.root_find( - _implicit_relation, - self.root_finder, - k0, - args, - throw=False, - max_steps=self.root_find_max_steps, - ) + with jax.numpy_dtype_promotion("standard"): + nonlinear_sol = optx.root_find( + _implicit_relation, + self.root_finder, + k0, + args, + throw=False, + max_steps=self.root_find_max_steps, + ) k1 = nonlinear_sol.value y1 = (y0**ω + k1**ω).ω # Use the trapezoidal rule for adaptive step sizing. diff --git a/diffrax/_step_size_controller/adaptive.py b/diffrax/_step_size_controller/adaptive.py index 20eadf92..c3458b7a 100644 --- a/diffrax/_step_size_controller/adaptive.py +++ b/diffrax/_step_size_controller/adaptive.py @@ -610,8 +610,7 @@ def _scale(_y0, _y1_candidate, _y_error): # a grad API boundary as part of a larger model.) factor = lax.stop_gradient(factor) factor = eqxi.nondifferentiable(factor) - with jax.numpy_dtype_promotion("standard"): - dt = prev_dt * factor + dt = prev_dt * factor.astype(prev_dt) # E.g. we failed an implicit step, so y_error=inf, so inv_scaled_error=0, # so factor=factormin, and we shrunk our step. diff --git a/test/test_solver.py b/test/test_solver.py index f10b8b39..28f5064c 100644 --- a/test/test_solver.py +++ b/test/test_solver.py @@ -274,7 +274,8 @@ def order(self, terms): # Essentially used as a check that our general IMEX implementation is correct. -def test_sil3(): +@pytest.mark.parametrize("dtype", (jnp.float64, jnp.complex128)) +def test_sil3(dtype): class ReferenceSil3(diffrax.AbstractImplicitSolver): term_structure = diffrax.MultiTerm[ tuple[diffrax.AbstractTerm, diffrax.AbstractTerm] @@ -313,7 +314,8 @@ def _second_stage(ya, _): return ya - (y0 + (1 / 3) * f0 + (1 / 6) * g0 + (1 / 6) * g1) ta = t0 + (1 / 3) * dt - ya = optx.root_find(_second_stage, self.root_finder, y0).value + with jax.numpy_dtype_promotion("standard"): + ya = optx.root_find(_second_stage, self.root_finder, y0).value fs.append(ex_vf_prod(ta, ya)) gs.append(im_vf_prod(ta, ya)) @@ -326,7 +328,9 @@ def _third_stage(yb, _): ) tb = t0 + (2 / 3) * dt - yb = optx.root_find(_third_stage, self.root_finder, ya).value + + with jax.numpy_dtype_promotion("standard"): + yb = optx.root_find(_third_stage, self.root_finder, ya).value fs.append(ex_vf_prod(tb, yb)) gs.append(im_vf_prod(tb, yb)) @@ -345,7 +349,8 @@ def _fourth_stage(yc, _): ) tc = t1 - yc = optx.root_find(_fourth_stage, self.root_finder, yb).value + with jax.numpy_dtype_promotion("standard"): + yc = optx.root_find(_fourth_stage, self.root_finder, yb).value fs.append(ex_vf_prod(tc, yc)) gs.append(im_vf_prod(tc, yc)) @@ -379,17 +384,19 @@ def _fourth_stage(yc, _): mlp2 = eqx.nn.MLP(3, 2, 8, 1, key=mlpkey2) def f1(t, y, args): - y = jnp.concatenate([t[None], y]) - return mlp1(y) + with jax.numpy_dtype_promotion("standard"): + y = jnp.concatenate([t[None], y]) + return mlp1(y) def f2(t, y, args): y = jnp.concatenate([t[None], y]) - return mlp2(y) + with jax.numpy_dtype_promotion("standard"): + return mlp2(y) terms = diffrax.MultiTerm(diffrax.ODETerm(f1), diffrax.ODETerm(f2)) t0 = jnp.array(0.3) t1 = jnp.array(1.5) - y0 = jr.normal(ykey, (2,), dtype=jnp.float64) + y0 = jr.normal(ykey, (2,), dtype=dtype) args = None state = solver.init(terms, t0, t1, y0, args) From ced7b6b84514b896a8aeb0eacfcb107bc00b1fc2 Mon Sep 17 00:00:00 2001 From: Evgenii Zheltonozhskii Date: Tue, 30 Apr 2024 01:25:18 +0200 Subject: [PATCH 3/4] New sde related fixes --- diffrax/_brownian/tree.py | 55 +++++++++++++++++-------------- diffrax/_integrate.py | 3 +- diffrax/_solver/implicit_euler.py | 18 +++++----- diffrax/_solver/srk.py | 5 ++- test/helpers.py | 16 +++++---- test/test_sde.py | 10 ++++-- test/test_solver.py | 11 +++---- 7 files changed, 64 insertions(+), 54 deletions(-) diff --git a/diffrax/_brownian/tree.py b/diffrax/_brownian/tree.py index 813384c4..3d40e186 100644 --- a/diffrax/_brownian/tree.py +++ b/diffrax/_brownian/tree.py @@ -91,7 +91,7 @@ def _levy_diff(_, x0: tuple, x1: tuple) -> AbstractBrownianIncrement: assert len(x1) == 2 dt0, w0 = x0 dt1, w1 = x1 - su = jnp.asarray(dt1 - dt0, dtype=w0.dtype) + su = jnp.asarray(dt1 - dt0, dtype=complex_to_real_dtype(w0.dtype)) return BrownianIncrement(dt=su, W=w1 - w0) elif len(x0) == 4: # space-time levy area case @@ -100,12 +100,13 @@ def _levy_diff(_, x0: tuple, x1: tuple) -> AbstractBrownianIncrement: dt1, w1, hh1, bhh1 = x1 w_su = w1 - w0 - su = jnp.asarray(dt1 - dt0, dtype=w0.dtype) + su = jnp.asarray(dt1 - dt0, dtype=complex_to_real_dtype(w0.dtype)) _su = jnp.where(jnp.abs(su) < jnp.finfo(su).eps, jnp.inf, su) inverse_su = 1 / _su - u_bb_s = dt1 * w0 - dt0 * w1 - bhh_su = bhh1 - bhh0 - 0.5 * u_bb_s # bhh_su = H_{s,u} * (u-s) - hh_su = inverse_su * bhh_su + with jax.numpy_dtype_promotion("standard"): + u_bb_s = dt1 * w0 - dt0 * w1 + bhh_su = bhh1 - bhh0 - 0.5 * u_bb_s # bhh_su = H_{s,u} * (u-s) + hh_su = inverse_su * bhh_su return SpaceTimeLevyArea(dt=su, W=w_su, H=hh_su) else: assert False @@ -396,27 +397,31 @@ def _body_fun(_state: _State): a = d_prime * sr3 * sr_ru_half b = d_prime * ru3 * sr_ru_half - w_sr = sr / su * w_su + 6 * sr * ru / su3 * bhh_su + 2 * (a + b) / su * x1 - w_r = w_s + w_sr - c = jnp.sqrt(3 * sr3 * ru3) / (6 * d) - bhh_sr = sr3 / su3 * bhh_su - a * x1 + c * x2 - bhh_r = bhh_s + bhh_sr + 0.5 * (r * w_s - s * w_r) + with jax.numpy_dtype_promotion("standard"): + w_sr = ( + sr / su * w_su + 6 * sr * ru / su3 * bhh_su + 2 * (a + b) / su * x1 + ) + w_r = w_s + w_sr + c = jnp.sqrt(3 * sr3 * ru3) / (6 * d) + bhh_sr = sr3 / su3 * bhh_su - a * x1 + c * x2 + bhh_r = bhh_s + bhh_sr + 0.5 * (r * w_s - s * w_r) - inverse_r = 1 / jnp.where(jnp.abs(r) < jnp.finfo(r).eps, jnp.inf, r) - hh_r = inverse_r * bhh_r + inverse_r = 1 / jnp.where(jnp.abs(r) < jnp.finfo(r).eps, jnp.inf, r) + hh_r = inverse_r * bhh_r elif self.levy_area is BrownianIncrement: - w_mean = w_s + sr / su * w_su - if self._spline == "sqrt": - z = jr.normal(final_state.key, shape, dtype) - bb = jnp.sqrt(sr * ru / su) * z - elif self._spline == "quad": - z = jr.normal(final_state.key, shape, dtype) - bb = (sr * ru / su) * z - elif self._spline == "zero": - bb = jnp.zeros(shape, dtype) - else: - assert False + with jax.numpy_dtype_promotion("standard"): + w_mean = w_s + sr / su * w_su + if self._spline == "sqrt": + z = jr.normal(final_state.key, shape, dtype) + bb = jnp.sqrt(sr * ru / su) * z + elif self._spline == "quad": + z = jr.normal(final_state.key, shape, dtype) + bb = (sr * ru / su) * z + elif self._spline == "zero": + bb = jnp.zeros(shape, dtype) + else: + assert False w_r = w_mean + bb return r, w_r @@ -499,8 +504,8 @@ def _brownian_arch( w_t = w_s + w_st w_stu = (w_s, w_t, w_u) - - bhh_t = bhh_s + bhh_st + 0.5 * (t * w_s - s * w_t) + with jax.numpy_dtype_promotion("standard"): + bhh_t = bhh_s + bhh_st + 0.5 * (t * w_s - s * w_t) bhh_stu = (bhh_s, bhh_t, bhh_u) bkk_stu = None bkk_st_tu = None diff --git a/diffrax/_integrate.py b/diffrax/_integrate.py index 5a18ff2a..f052dcf0 100644 --- a/diffrax/_integrate.py +++ b/diffrax/_integrate.py @@ -155,7 +155,8 @@ def _check(term_cls, term, term_contr_kwargs, yi): # If we've got to this point then the term is compatible try: - jtu.tree_map(_check, term_structure, terms, contr_kwargs, y) + with jax.numpy_dtype_promotion("standard"): + jtu.tree_map(_check, term_structure, terms, contr_kwargs, y) except ValueError: # ValueError may also arise from mismatched tree structures return False diff --git a/diffrax/_solver/implicit_euler.py b/diffrax/_solver/implicit_euler.py index 6d9c79bd..eb3bdb00 100644 --- a/diffrax/_solver/implicit_euler.py +++ b/diffrax/_solver/implicit_euler.py @@ -2,7 +2,6 @@ from typing import ClassVar from typing_extensions import TypeAlias -import jax import optimistix as optx from equinox.internal import ω @@ -82,15 +81,14 @@ def step( # write out a `ButcherTableau` and use `AbstractSDIRK`. k0 = terms.vf_prod(t0, y0, args, control) args = (terms.vf_prod, t1, y0, args, control) - with jax.numpy_dtype_promotion("standard"): - nonlinear_sol = optx.root_find( - _implicit_relation, - self.root_finder, - k0, - args, - throw=False, - max_steps=self.root_find_max_steps, - ) + nonlinear_sol = optx.root_find( + _implicit_relation, + self.root_finder, + k0, + args, + throw=False, + max_steps=self.root_find_max_steps, + ) k1 = nonlinear_sol.value y1 = (y0**ω + k1**ω).ω # Use the trapezoidal rule for adaptive step sizing. diff --git a/diffrax/_solver/srk.py b/diffrax/_solver/srk.py index c40019df..c926ae68 100644 --- a/diffrax/_solver/srk.py +++ b/diffrax/_solver/srk.py @@ -12,6 +12,7 @@ import numpy as np from equinox.internal import ω from jaxtyping import Array, Float, PyTree +from lineax.internal import complex_to_real_dtype from .._custom_types import ( AbstractBrownianIncrement, @@ -340,7 +341,9 @@ def step( # First the drift related stuff a = self._embed_a_lower(self.tableau.a, dtype) - c = jnp.asarray(np.insert(self.tableau.c, 0, 0.0), dtype=dtype) + c = jnp.asarray( + np.insert(self.tableau.c, 0, 0.0), dtype=complex_to_real_dtype(dtype) + ) b_sol = jnp.asarray(self.tableau.b_sol, dtype=dtype) def make_zeros(): diff --git a/test/helpers.py b/test/helpers.py index ab9e37af..6eb10c61 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -18,7 +18,6 @@ ) from jax import Array from jaxtyping import PRNGKeyArray, PyTree, Shaped -from lineax.internal import complex_to_real_dtype all_ode_solvers = ( @@ -252,7 +251,6 @@ def sde_solver_strong_order( bm_tol, saveat, ) - dts = 2.0 ** jnp.arange(-3, -3 - num_levels, -1, dtype=dtype) errs_list, steps_list = [], [] for level in range(level_coarse, level_fine + 1): @@ -277,7 +275,8 @@ def sde_solver_strong_order( steps_list.append(jnp.average(steps)) errs_arr = jnp.array(errs_list) steps_arr = jnp.array(steps_list) - order, _ = jnp.polyfit(jnp.log(1 / steps_arr), jnp.log(errs_arr), 1) + with jax.numpy_dtype_promotion("standard"): + order, _ = jnp.polyfit(jnp.log(1 / steps_arr), jnp.log(errs_arr), 1) return steps_arr, errs_arr, order @@ -360,12 +359,14 @@ def _squareplus(x): def drift(t, y, args): mlp, _, _ = args - return 0.25 * mlp(y) + with jax.numpy_dtype_promotion("standard"): + return 0.25 * mlp(y) def diffusion(t, y, args): _, mlp, noise_dim = args - return 1.0 * mlp(y).reshape(3, noise_dim) + with jax.numpy_dtype_promotion("standard"): + return 1.0 * mlp(y).reshape(3, noise_dim) def get_mlp_sde(t0, t1, dtype, key, noise_dim): @@ -447,8 +448,9 @@ def ft(t): drift_mlp = init_linear_weight(drift_mlp, lap_init, driftkey) def _drift(t, y, _): - mlp_out = drift_mlp(jnp.concatenate([y, ft(t)])) - return (0.01 * mlp_out - 0.5 * y**3) / (jnp.sum(y**2) + 1) + with jax.numpy_dtype_promotion("standard"): + mlp_out = drift_mlp(jnp.concatenate([y, ft(t)])) + return (0.01 * mlp_out - 0.5 * y**3) / (jnp.sum(y**2) + 1) diffusion_mx = jr.normal(diffusionkey, (4, y_dim, noise_dim), dtype=dtype) diff --git a/test/test_sde.py b/test/test_sde.py index 5ffb3a6d..42493c76 100644 --- a/test/test_sde.py +++ b/test/test_sde.py @@ -43,8 +43,12 @@ def _solvers_and_orders(): # converges to its own limit (i.e. using itself as reference), and then in a # different test check whether that limit is the same as the Euler/Heun limit. @pytest.mark.parametrize("solver_ctr,noise,theoretical_order", _solvers_and_orders()) +@pytest.mark.parametrize( + "dtype", + (jnp.float64,), +) def test_sde_strong_order_new( - solver_ctr, noise: Literal["any", "com", "add"], theoretical_order + solver_ctr, noise: Literal["any", "com", "add"], theoretical_order, dtype ): bmkey = jr.PRNGKey(5678) sde_key = jr.PRNGKey(11) @@ -54,7 +58,7 @@ def test_sde_strong_order_new( t1 = 5.3 if noise == "add": - sde = get_time_sde(t0, t1, jnp.float64, sde_key, noise_dim=7) + sde = get_time_sde(t0, t1, dtype, sde_key, noise_dim=7) else: if noise == "com": noise_dim = 1 @@ -62,7 +66,7 @@ def test_sde_strong_order_new( noise_dim = 5 else: assert False - sde = get_mlp_sde(t0, t1, jnp.float64, sde_key, noise_dim=noise_dim) + sde = get_mlp_sde(t0, t1, dtype, sde_key, noise_dim=noise_dim) ref_solver = solver_ctr() level_coarse, level_fine = 1, 7 diff --git a/test/test_solver.py b/test/test_solver.py index 28f5064c..aa618712 100644 --- a/test/test_solver.py +++ b/test/test_solver.py @@ -274,7 +274,7 @@ def order(self, terms): # Essentially used as a check that our general IMEX implementation is correct. -@pytest.mark.parametrize("dtype", (jnp.float64, jnp.complex128)) +@pytest.mark.parametrize("dtype", (jnp.float64,)) def test_sil3(dtype): class ReferenceSil3(diffrax.AbstractImplicitSolver): term_structure = diffrax.MultiTerm[ @@ -314,8 +314,7 @@ def _second_stage(ya, _): return ya - (y0 + (1 / 3) * f0 + (1 / 6) * g0 + (1 / 6) * g1) ta = t0 + (1 / 3) * dt - with jax.numpy_dtype_promotion("standard"): - ya = optx.root_find(_second_stage, self.root_finder, y0).value + ya = optx.root_find(_second_stage, self.root_finder, y0).value fs.append(ex_vf_prod(ta, ya)) gs.append(im_vf_prod(ta, ya)) @@ -329,8 +328,7 @@ def _third_stage(yb, _): tb = t0 + (2 / 3) * dt - with jax.numpy_dtype_promotion("standard"): - yb = optx.root_find(_third_stage, self.root_finder, ya).value + yb = optx.root_find(_third_stage, self.root_finder, ya).value fs.append(ex_vf_prod(tb, yb)) gs.append(im_vf_prod(tb, yb)) @@ -349,8 +347,7 @@ def _fourth_stage(yc, _): ) tc = t1 - with jax.numpy_dtype_promotion("standard"): - yc = optx.root_find(_fourth_stage, self.root_finder, yb).value + yc = optx.root_find(_fourth_stage, self.root_finder, yb).value fs.append(ex_vf_prod(tc, yc)) gs.append(im_vf_prod(tc, yc)) From 21be39942fa57b8da09c8fa44552d3d33b5f646b Mon Sep 17 00:00:00 2001 From: Evgenii Zheltonozhskii Date: Tue, 30 Apr 2024 02:24:29 +0200 Subject: [PATCH 4/4] New sde related fixes --- diffrax/_brownian/path.py | 3 ++- diffrax/_solver/srk.py | 14 ++++++++------ diffrax/_term.py | 6 ++++-- test/helpers.py | 3 ++- test/test_brownian.py | 7 +++++++ test/test_sde.py | 36 +++++++++++++++++++++++++----------- 6 files changed, 48 insertions(+), 21 deletions(-) diff --git a/diffrax/_brownian/path.py b/diffrax/_brownian/path.py index 66075069..3fa481aa 100644 --- a/diffrax/_brownian/path.py +++ b/diffrax/_brownian/path.py @@ -9,6 +9,7 @@ import jax.tree_util as jtu import lineax.internal as lxi from jaxtyping import Array, PRNGKeyArray, PyTree +from lineax.internal import complex_to_real_dtype from .._custom_types import ( AbstractBrownianIncrement, @@ -130,7 +131,7 @@ def _evaluate_leaf( ): w_std = jnp.sqrt(t1 - t0).astype(shape.dtype) w = jr.normal(key, shape.shape, shape.dtype) * w_std - dt = jnp.asarray(t1 - t0, dtype=shape.dtype) + dt = jnp.asarray(t1 - t0, dtype=complex_to_real_dtype(shape.dtype)) if levy_area is SpaceTimeLevyArea: key, key_hh = jr.split(key, 2) diff --git a/diffrax/_solver/srk.py b/diffrax/_solver/srk.py index c926ae68..494ceacd 100644 --- a/diffrax/_solver/srk.py +++ b/diffrax/_solver/srk.py @@ -406,7 +406,7 @@ def aux_add_levy(w_leaf, *levy_leaves): def _comp_g(_t): return diffusion.vf(_t, y0, args) - g0_g1 = _comp_g(jnp.array([t0, t1], dtype=dtype)) + g0_g1 = _comp_g(jnp.array([t0, t1], dtype=complex_to_real_dtype(dtype))) g0 = jtu.tree_map(lambda g_leaf: g_leaf[0], g0_g1) # g_delta = 0.5 * g1 - g0 g_delta = jtu.tree_map(lambda g_leaf: 0.5 * (g_leaf[1] - g_leaf[0]), g0_g1) @@ -534,13 +534,15 @@ def compute_and_insert_kf_j(_h_kfs_in): return (_h_kfs, None, None), None def compute_and_insert_kg_j(_w_kgs_in, _levylist_kgs_in): - _w_kg_j = diffusion.vf_prod(t0 + c_j * h, z_j, args, w) + with jax.numpy_dtype_promotion("standard"): + _w_kg_j = diffusion.vf_prod(t0 + c_j * h, z_j, args, w) new_w_kgs = insert_jth_stage(_w_kgs_in, _w_kg_j, j) - _levylist_kg_j = [ - diffusion.vf_prod(t0 + c_j * h, z_j, args, levy) - for levy in levy_areas - ] + with jax.numpy_dtype_promotion("standard"): + _levylist_kg_j = [ + diffusion.vf_prod(t0 + c_j * h, z_j, args, levy) + for levy in levy_areas + ] new_levylist_kgs = insert_jth_stage(_levylist_kgs_in, _levylist_kg_j, j) return new_w_kgs, new_levylist_kgs diff --git a/diffrax/_term.py b/diffrax/_term.py index 2f7eca30..e61ca751 100644 --- a/diffrax/_term.py +++ b/diffrax/_term.py @@ -361,7 +361,8 @@ class WeaklyDiagonalControlTerm(_AbstractControlTerm[_VF, _Control]): """ def prod(self, vf: _VF, control: _Control) -> Y: - return jtu.tree_map(operator.mul, vf, control) + with jax.numpy_dtype_promotion("standard"): + return jtu.tree_map(operator.mul, vf, control) class _ControlToODE(eqx.Module): @@ -461,7 +462,8 @@ def contr(self, t0: RealScalarLike, t1: RealScalarLike, **kwargs) -> _Control: return (self.direction * self.term.contr(_t0, _t1, **kwargs) ** ω).ω def prod(self, vf: _VF, control: _Control) -> Y: - return self.term.prod(vf, control) + with jax.numpy_dtype_promotion("standard"): + return self.term.prod(vf, control) def vf_prod(self, t: RealScalarLike, y: Y, args: Args, control: _Control) -> Y: t = t * self.direction diff --git a/test/helpers.py b/test/helpers.py index 6eb10c61..b7e33b54 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -99,7 +99,8 @@ def path_l2_dist( # all but the first two axes (which represent the number of samples # and the length of saveat). Also sum all the PyTree leaves. def sum_square_diff(y1, y2): - square_diff = jnp.square(y1 - y2) + with jax.numpy_dtype_promotion("standard"): + square_diff = jnp.square(y1 - y2) # sum all but the first two axes axes = range(2, square_diff.ndim) out = jnp.sum(square_diff, axis=axes) diff --git a/test/test_brownian.py b/test/test_brownian.py index ac227859..70eeb91b 100644 --- a/test/test_brownian.py +++ b/test/test_brownian.py @@ -46,6 +46,11 @@ def test_shape_and_dtype(ctr, levy_area, use_levy, getkey): (2,), (3, 4), (1, 2, 3, 4), + (1, 2, 3, 4), + { + "a": (1,), + "b": (2, 3), + }, { "a": (1,), "b": (2, 3), @@ -66,7 +71,9 @@ def test_shape_and_dtype(ctr, levy_area, use_levy, getkey): jnp.float16, jnp.float32, jnp.float64, + jnp.complex128, {"a": None, "b": jnp.float64}, + {"a": jnp.float64, "b": jnp.complex128}, (jnp.float16, (jnp.float32, jnp.float64)), ) diff --git a/test/test_sde.py b/test/test_sde.py index 42493c76..b2f62245 100644 --- a/test/test_sde.py +++ b/test/test_sde.py @@ -116,8 +116,12 @@ def get_dt_and_controller(level): # using a single reference solution. We use Euler if the solver is Ito # and Heun if the solver is Stratonovich. @pytest.mark.parametrize("solver_ctr,noise,theoretical_order", _solvers_and_orders()) +@pytest.mark.parametrize( + "dtype", + (jnp.float64,), +) def test_sde_strong_limit( - solver_ctr, noise: Literal["any", "com", "add"], theoretical_order + solver_ctr, noise: Literal["any", "com", "add"], theoretical_order, dtype ): bmkey = jr.PRNGKey(5678) sde_key = jr.PRNGKey(11) @@ -127,7 +131,7 @@ def test_sde_strong_limit( t1 = 5.3 if noise == "add": - sde = get_time_sde(t0, t1, jnp.float64, sde_key, noise_dim=3) + sde = get_time_sde(t0, t1, dtype, sde_key, noise_dim=3) level_fine = 12 if theoretical_order <= 1.0: level_coarse = 11 @@ -141,7 +145,7 @@ def test_sde_strong_limit( noise_dim = 5 else: assert False - sde = get_mlp_sde(t0, t1, jnp.float64, sde_key, noise_dim=noise_dim) + sde = get_mlp_sde(t0, t1, dtype, sde_key, noise_dim=noise_dim) # Reference solver is always an ODE-viable solver, so its implementation has been # verified by the ODE tests like test_ode_order. @@ -210,9 +214,12 @@ def get_matrix(y_leaf): @pytest.mark.parametrize("shape", [(), (5, 2)]) @pytest.mark.parametrize("solver_ctr", _solvers()) -def test_sde_solver_shape(shape, solver_ctr): +@pytest.mark.parametrize( + "dtype", + (jnp.float64, jnp.complex128), +) +def test_sde_solver_shape(shape, solver_ctr, dtype): pytree = ({"a": 0, "b": [0, 0]}, 0, 0) - dtype = jnp.float64 key = jr.PRNGKey(0) y0 = jtu.tree_map(lambda _: jr.normal(key, shape, dtype=dtype), pytree) t0, t1, dt0 = 0.0, 1.0, 0.3 @@ -236,8 +243,7 @@ def test_sde_solver_shape(shape, solver_ctr): assert leaf[0].shape == shape -def _weakly_diagonal_noise_helper(solver): - dtype = jnp.float64 +def _weakly_diagonal_noise_helper(solver, dtype): w_shape = (3,) args = (0.5, 1.2) @@ -265,9 +271,17 @@ def _drift(t, y, args): @pytest.mark.parametrize("solver_ctr", _solvers()) -def test_weakly_diagonal_noise(solver_ctr): - _weakly_diagonal_noise_helper(solver_ctr()) +@pytest.mark.parametrize( + "dtype", + (jnp.float64, jnp.complex128), +) +def test_weakly_diagonal_noise(solver_ctr, dtype): + _weakly_diagonal_noise_helper(solver_ctr(), dtype) -def test_halfsolver_term_compatible(): - _weakly_diagonal_noise_helper(diffrax.HalfSolver(diffrax.SPaRK())) +@pytest.mark.parametrize( + "dtype", + (jnp.float64, jnp.complex128), +) +def test_halfsolver_term_compatible(dtype): + _weakly_diagonal_noise_helper(diffrax.HalfSolver(diffrax.SPaRK()), dtype)