diff --git a/diffrax/_integrate.py b/diffrax/_integrate.py index 3a1c59b7..1d3e7209 100644 --- a/diffrax/_integrate.py +++ b/diffrax/_integrate.py @@ -775,9 +775,56 @@ def _save_t1(subsaveat, save_state): save_state = _save(tfinal, yfinal, args, subsaveat.fn, save_state) return save_state + def _save_ts(subsaveat: SubSaveAt, save_state: SaveState) -> SaveState: + if subsaveat.ts is not None: + out_size = 1 if subsaveat.t0 else 0 + out_size += 1 if subsaveat.t1 and not subsaveat.steps else 0 + out_size += len(subsaveat.ts) + ys = jtu.tree_map( + lambda y: jnp.stack([y] * out_size), + subsaveat.fn(t0, yfinal, args), + ) + ts = jnp.full(out_size, t0) + if subsaveat.steps: + ysteps = jtu.tree_map( + lambda y: jnp.stack([y] * max_steps), + subsaveat.fn(t0, jnp.full_like(yfinal, jnp.inf), args), + ) + ys = jtu.tree_map( + lambda _ys, _ysteps: jnp.concatenate([_ys, _ysteps], axis=0), + ys, + ysteps, + ) + ts = jnp.concatenate((ts, jnp.full(max_steps, jnp.inf))) + save_state = SaveState( + saveat_ts_index=out_size, + ts=ts, + ys=ys, + save_index=out_size, + ) + return save_state + save_state = jtu.tree_map( _save_t1, saveat.subs, final_state.save_state, is_leaf=_is_subsaveat ) + + # if t0 == t1 then we don't enter the integration loop. In this case we have to + # manually update the saved ts and ys if we want to save at "intermediate" + # times specified by saveat.subs.ts + save_state = jax.lax.cond( + eqxi.unvmap_any(t0 == t1), + lambda __save_state: jax.lax.cond( + t0 == t1, + lambda _save_state: jtu.tree_map( + _save_ts, saveat.subs, _save_state, is_leaf=_is_subsaveat + ), + lambda _save_state: _save_state, + __save_state, + ), + lambda __save_state: __save_state, + save_state, + ) + final_state = eqx.tree_at( lambda s: s.save_state, final_state, save_state, is_leaf=_is_none ) diff --git a/test/test_saveat_solution.py b/test/test_saveat_solution.py index cc73bf99..ba536b04 100644 --- a/test/test_saveat_solution.py +++ b/test/test_saveat_solution.py @@ -147,6 +147,91 @@ def test_saveat_solution(): assert sol.result == diffrax.RESULTS.successful +@pytest.mark.parametrize("subs", [True, False]) +def test_t0_eq_t1(subs): + y0 = jnp.array([2.0]) + ts = jnp.linspace(1.0, 1.0, 3) + max_steps = 256 + if subs: + get0 = diffrax.SubSaveAt( + ts=ts, + t1=True, + ) + get1 = diffrax.SubSaveAt( + t0=True, + ts=ts, + ) + get2 = diffrax.SubSaveAt( + t0=True, + ts=ts, + steps=True, + ) + subs = (get0, get1, get2) + saveat = diffrax.SaveAt(subs=subs) + else: + saveat = diffrax.SaveAt(t0=True, t1=True, ts=ts) + term = diffrax.ODETerm(lambda t, y, args: y) + sol = diffrax.diffeqsolve( + term, + t0=ts[0], + t1=ts[-1], + y0=y0, + dt0=0.1, + solver=diffrax.Dopri5(), + saveat=saveat, + max_steps=max_steps, + ) + if subs: + compare = jnp.full((len(ts) + 1, *y0.shape), y0) + compare_2 = jnp.concatenate( + (compare, jnp.full((max_steps, *y0.shape), jnp.inf)) + ) + assert tree_allclose(sol.ys[0], compare) # pyright: ignore + assert tree_allclose(sol.ys[1], compare) # pyright: ignore + assert tree_allclose(sol.ys[2], compare_2) # pyright: ignore + else: + compare = jnp.full((len(ts) + 2, *y0.shape), y0) + assert tree_allclose(sol.ys, compare) + + +@pytest.mark.parametrize("subs", [True, False]) +def test_vmap_t0_eq_t1(subs): + ntsave = 4 + y0 = jnp.array([2.0]) + term = diffrax.ODETerm(lambda t, y, args: y) + + def _solve(tf): + ts = jnp.linspace(0.0, tf, ntsave) + get0 = diffrax.SubSaveAt( + ts=ts, + t1=True, + ) + get1 = diffrax.SubSaveAt( + t0=True, + ts=ts, + ) + subs = (get0, get1) + saveat = diffrax.SaveAt(subs=subs) + return diffrax.diffeqsolve( + term, + t0=ts[0], + t1=ts[-1], + y0=y0, + dt0=0.1, + solver=diffrax.Dopri5(), + saveat=saveat, + ) + + compare = jnp.full((ntsave + 1, *y0.shape), y0) + sol = jax.vmap(_solve)(jnp.array([0.0, 1.0])) + assert tree_allclose(sol.ys[0][0], compare) # pyright: ignore + assert tree_allclose(sol.ys[1][0], compare) # pyright: ignore + + regular_solve = _solve(1.0) + assert tree_allclose(sol.ys[0][1], regular_solve.ys[0]) # pyright: ignore + assert tree_allclose(sol.ys[1][1], regular_solve.ys[1]) # pyright: ignore + + def test_trivial_dense(): term = diffrax.ODETerm(lambda t, y, args: -0.5 * y) y0 = jnp.array([2.1])