From bca8567075b86cb1434da4d86b6c89c4f8848853 Mon Sep 17 00:00:00 2001 From: Danny Date: Thu, 3 Apr 2025 20:47:23 -0400 Subject: [PATCH 1/2] fix for t0==t1 when initializing ys, as opposed to inside the loop --- diffrax/_integrate.py | 83 +++++++++++++------------------------------ 1 file changed, 25 insertions(+), 58 deletions(-) diff --git a/diffrax/_integrate.py b/diffrax/_integrate.py index 88c014aa..ddd26a7b 100644 --- a/diffrax/_integrate.py +++ b/diffrax/_integrate.py @@ -788,60 +788,10 @@ def _save_t1(subsaveat, save_state): save_state = _save(tfinal, yfinal, args, subsaveat.fn, save_state) return save_state - def _save_if_t0_equals_t1(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) - - def _make_ys(out, old_outs): - outs = jnp.stack([out] * out_size) - if subsaveat.steps: - outs = jnp.concatenate( - [ - outs, - jnp.full( - (max_steps,) + out.shape, jnp.inf, dtype=out.dtype - ), - ] - ) - assert outs.shape == old_outs.shape - return outs - - ts = jnp.full(out_size, t0) - if subsaveat.steps: - ts = jnp.concatenate((ts, jnp.full(max_steps, jnp.inf, dtype=ts.dtype))) - assert ts.shape == save_state.ts.shape - ys = jtu.tree_map(_make_ys, subsaveat.fn(t0, yfinal, args), save_state.ys) - 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_if_t0_equals_t1, 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 ) @@ -1249,6 +1199,27 @@ def _allocate_output(subsaveat: SubSaveAt) -> SaveState: out_size += 1 if subsaveat.ts is not None: out_size += len(subsaveat.ts) + if subsaveat.t1 and not subsaveat.steps: + out_size += 1 + + def ts_ys_inf_subsave(_out_size): + _ts = jnp.full(_out_size, direction * jnp.inf, dtype=time_dtype) + _ys = jtu.tree_map( + lambda y: jnp.full((_out_size,) + y.shape, jnp.inf, dtype=y.dtype), + eqx.filter_eval_shape(subsaveat.fn, t0, y0, args), + ) + return _ts, _ys + + # if t0 == t1 then we don't enter the integration loop. In this case we have to + # manually set the saved states to y0 as opposed to the default inf + ts, ys = ts_ys_inf_subsave(out_size) + if out_size != 0: + y0_subsave = jtu.tree_map( + lambda y: jnp.stack([subsaveat.fn(t0, y, args)] * out_size), y0 + ) + ys = jtu.tree_map( + lambda _y0, _ys: jnp.where(t0 == t1, _y0, _ys), y0_subsave, ys + ) if subsaveat.steps: # We have no way of knowing how many steps we'll actually end up taking, and # XLA doesn't support dynamic shapes. So we just have to allocate the @@ -1257,16 +1228,12 @@ def _allocate_output(subsaveat: SubSaveAt) -> SaveState: raise ValueError( "`max_steps=None` is incompatible with saving at `steps=True`" ) - out_size += max_steps - if subsaveat.t1 and not subsaveat.steps: - out_size += 1 + ts_steps, ys_steps = ts_ys_inf_subsave(max_steps) + ts = jnp.concatenate((ts, ts_steps)) + ys = jtu.tree_map(lambda x, y: jnp.concatenate((x, y)), ys, ys_steps) + saveat_ts_index = 0 save_index = 0 - ts = jnp.full(out_size, direction * jnp.inf, dtype=time_dtype) - struct = eqx.filter_eval_shape(subsaveat.fn, t0, y0, args) - ys = jtu.tree_map( - lambda y: jnp.full((out_size,) + y.shape, jnp.inf, dtype=y.dtype), struct - ) return SaveState( ts=ts, ys=ys, save_index=save_index, saveat_ts_index=saveat_ts_index ) From 36c1d3e4b05456332afa1923340cd924f76da134 Mon Sep 17 00:00:00 2001 From: Danny Date: Fri, 4 Apr 2025 08:18:15 -0400 Subject: [PATCH 2/2] bug fix --- diffrax/_integrate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/diffrax/_integrate.py b/diffrax/_integrate.py index ddd26a7b..02ccbba8 100644 --- a/diffrax/_integrate.py +++ b/diffrax/_integrate.py @@ -1215,7 +1215,7 @@ def ts_ys_inf_subsave(_out_size): ts, ys = ts_ys_inf_subsave(out_size) if out_size != 0: y0_subsave = jtu.tree_map( - lambda y: jnp.stack([subsaveat.fn(t0, y, args)] * out_size), y0 + lambda y: jnp.stack([y] * out_size), subsaveat.fn(t0, y0, args) ) ys = jtu.tree_map( lambda _y0, _ys: jnp.where(t0 == t1, _y0, _ys), y0_subsave, ys