diff --git a/diffrax/_step_size_controller/clip.py b/diffrax/_step_size_controller/clip.py index 4899da4b..14b52cc4 100644 --- a/diffrax/_step_size_controller/clip.py +++ b/diffrax/_step_size_controller/clip.py @@ -98,6 +98,17 @@ def _bump_next_t0(next_t0, ts): return next_t0, made_jump1 | made_jump2 +def _drop_step_tangent(next_t0, next_t1, landed): + # After a step that landed on a clip time, the proposed step keeps its value but + # not its tangent: see the comment at its use in `adapt_step_size`. + # (Written so that the value is exactly `next_t1`, and the tangent that of + # `next_t0`.) + rigid_t1 = jax.lax.stop_gradient(next_t1) + ( + next_t0 - jax.lax.stop_gradient(next_t0) + ) + return jnp.where(landed, rigid_t1, next_t1) + + def _find_idx_with_hint(t: RealScalarLike, ts: Array | None, hint: IntScalarLike): # Find index of first element of `ts` strictly greater than `t`. # Uses a linear search starting from `hint`. The value `hint` is assumed to be in @@ -367,7 +378,8 @@ def callback(_keep_step, _t1): # propose a step over the interval [something, prevbefore(x)], then on the # next step the inner controller will propose a step over [prevbefore(x), x] # which definitely isn't desired! - _next_t0, _ = _bump_next_t0(next_t0, step_ts) + _next_t0, landed_on_step = _bump_next_t0(next_t0, step_ts) + next_t1 = _drop_step_tangent(next_t0, next_t1, landed_on_step) step_index = _find_idx_with_hint(_next_t0, step_ts, step_index) next_t1 = _clip_t(next_t1, step_index, step_ts, False) step_info = step_index, step_ts @@ -376,6 +388,15 @@ def callback(_keep_step, _t1): else: jump_index, jump_ts = controller_state.jump_info next_t0, made_jump2 = _bump_next_t0(next_t0, jump_ts) + # The step we just made was clipped to (just before) this jump time, so it + # is often very short, whilst the tangent of its length, d(jump) - d(t0), + # is O(1). The inner controller proposes the next step as a multiple of + # that length, and keeping its tangent would scale the tangents of all the + # following step times by (step size / clipped step size): with several + # close traced `jump_ts`, the derivative with respect to them blows up. + # Instead, the proposed step moves rigidly with the jump time. (Likewise + # after landing on a `step_ts` time, above.) + next_t1 = _drop_step_tangent(next_t0, next_t1, made_jump2) # This next line is to fix # https://github.com/patrick-kidger/diffrax/issues/713 # TODO: should we add this to the `step_ts` branch as well? diff --git a/test/test_adaptive_stepsize_controller.py b/test/test_adaptive_stepsize_controller.py index 4a51336f..5fa2500c 100644 --- a/test/test_adaptive_stepsize_controller.py +++ b/test/test_adaptive_stepsize_controller.py @@ -239,6 +239,63 @@ def forcing(s): assert tree_allclose(finite_diff, autodiff) +# The step clipped to just before a jump is often tiny, whilst the tangent of its +# length, d(jump) - d(t0), is O(1). This used to be carried into all later steps by +# `PIDController` (whose `dt = prev_dt * factor` only stopped the gradient of `factor`) +# and grew jump after jump, so that the gradient with respect to several close jump +# times was wrong by many orders of magnitude. With `step_ts` too, the steps that land +# on a `step_ts` time did the same. +@pytest.mark.parametrize("with_step_ts", [False, True]) +@pytest.mark.parametrize( + "adjoint", + [ + diffrax.RecursiveCheckpointAdjoint(), + diffrax.DirectAdjoint(), + diffrax.ForwardMode(), + ], +) +def test_grad_wrt_close_jump_ts(adjoint, with_step_ts): + def run(shift): + # A train of 7 pulses 0.04 apart, whose 14 edges all move with `shift`. + length = (1.78 - 6 * 0.04) / 7 + starts = 0.1 + jnp.arange(7.0) * (length + 0.04) + shift + ends = starts + length + + def vector_field(t, y, args): + y, _ = y + forcing = jnp.sum(jnp.where((t >= starts) & (t < ends), 15.0, 0.0)) + return -20.0 * y + forcing, y**2 + + pid_controller = diffrax.PIDController(rtol=1e-10, atol=1e-10) + stepsize_controller = diffrax.ClipStepSizeController( + pid_controller, + jump_ts=jnp.sort(jnp.concatenate([starts, ends])), + step_ts=jnp.linspace(0.0, 2.0, 101) if with_step_ts else None, + ) + sol = diffrax.diffeqsolve( + diffrax.ODETerm(vector_field), + diffrax.Tsit5(), + 0.0, + 2.0, + None, + (0.0, 0.0), + stepsize_controller=stepsize_controller, + adjoint=adjoint, + ) + _, integral = cast(Array, sol.ys) + (integral,) = integral + return integral + + shift = 0.013 + eps = 1e-6 + finite_diff = (run(shift + eps) - run(shift - eps)) / (2 * eps) + if isinstance(adjoint, diffrax.ForwardMode): + _, autodiff = jax.jit(lambda s: jax.jvp(run, (s,), (1.0,)))(shift) + else: + autodiff = jax.jit(jax.grad(run))(shift) + assert tree_allclose(finite_diff, autodiff, rtol=1e-4) + + def test_pid_meta(): ts = jnp.array([3, 4], dtype=jnp.float64) pid1 = diffrax.PIDController(rtol=1e-4, atol=1e-6)