diff --git a/diffrax/_step_size_controller/clip.py b/diffrax/_step_size_controller/clip.py index 4899da4b..7cc30312 100644 --- a/diffrax/_step_size_controller/clip.py +++ b/diffrax/_step_size_controller/clip.py @@ -100,21 +100,14 @@ def _bump_next_t0(next_t0, ts): 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 - # `{0, 1, ..., len(ts)}` + # The value `hint` is assumed to be in `{0, 1, ..., len(ts)}`; it is only used for + # its dtype. We use an unrolled binary search rather than a `while_loop` walking + # from `hint`: each iteration of a `while_loop` costs a device-to-host + # synchronisation on GPU, whilst this is a handful of fused comparisons. if ts is None: return 0 - - def cond_up(_i): - return (_i < len(ts)) & (ts[_i] <= t) - - def cond_down(_i): - return (_i > 0) & (ts[_i - 1] > t) - - i = hint - i = jax.lax.while_loop(cond_up, lambda _i: _i + 1, i) - i = jax.lax.while_loop(cond_down, lambda _i: _i - 1, i) - return i + i = jnp.searchsorted(ts, t, side="right", method="scan_unrolled") + return i.astype(jnp.result_type(hint)) class ClipStepSizeController( diff --git a/test/test_adaptive_stepsize_controller.py b/test/test_adaptive_stepsize_controller.py index 4a51336f..8e643ff9 100644 --- a/test/test_adaptive_stepsize_controller.py +++ b/test/test_adaptive_stepsize_controller.py @@ -313,6 +313,21 @@ def test_find_idx_with_hint(): assert idx == 3 # not 2; we want the first value *strictly* greater. idx = _find_idx_with_hint(1.9, ts, hint) assert idx == 2 + assert _find_idx_with_hint(-1.0, ts, hint) == 0 + assert _find_idx_with_hint(4.0, ts, hint) == 5 + assert _find_idx_with_hint(7.0, ts, hint) == 5 + + +def test_find_idx_with_hint_is_loop_free(): + # A `while_loop` costs a device-to-host synchronisation per iteration on GPU, and + # this search runs on every step of an adaptive solve. + ts = jnp.linspace(0.0, 1.0, 100) + hint = jnp.array(3) + jaxpr = jax.make_jaxpr(lambda t: _find_idx_with_hint(t, ts, hint))(0.5) + assert "while" not in str(jaxpr) + idx = jax.jit(lambda t: _find_idx_with_hint(t, ts, hint))(0.5) + assert idx == 50 + assert jnp.result_type(idx) == hint.dtype # https://github.com/patrick-kidger/diffrax/issues/607