Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 22 additions & 1 deletion diffrax/_step_size_controller/clip.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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?
Expand Down
57 changes: 57 additions & 0 deletions test/test_adaptive_stepsize_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down