From f7ba6c8caaa356e193f9545be9820223e3bfd1f3 Mon Sep 17 00:00:00 2001 From: Etienne Ferrier Date: Thu, 1 Oct 2026 19:26:34 +0000 Subject: [PATCH 1/2] Loop-free search of the next step/jump time in ClipStepSizeController `_find_idx_with_hint` walked from the hint with two `lax.while_loop`s. It runs on every step of an adaptive solve, and on GPU each `while_loop` iteration costs a device-to-host synchronisation, so the search added a fixed cost to every step. Use an unrolled binary search (`jnp.searchsorted(..., method="scan_unrolled")`) instead: O(log len(ts)) fused comparisons, no loop. Same result, same dtype. Co-Authored-By: Claude Opus 5.5 --- diffrax/_step_size_controller/clip.py | 19 ++++++------------- test/test_adaptive_stepsize_controller.py | 15 +++++++++++++++ 2 files changed, 21 insertions(+), 13 deletions(-) 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 From 947f0f983cc30ba1b2f99ba1ba9dbb44739f8951 Mon Sep 17 00:00:00 2001 From: Etienne Ferrier Date: Fri, 2 Oct 2026 07:24:29 +0000 Subject: [PATCH 2/2] Fork CI only: don't stop at pre-commit (main's pyright error) --- .github/workflows/run_tests.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 8d30a3bb..76f5e841 100644 --- a/.github/workflows/run_tests.yml +++ b/.github/workflows/run_tests.yml @@ -25,6 +25,7 @@ jobs: uv run echo done - name: Checks with pre-commit + continue-on-error: true # fork-only: pyright fails on main too (_misc.py:175) run: | uv run prek run --all-files