Skip to content

Loop-free search of the next step/jump time in ClipStepSizeController - #777

Open
etiferrier wants to merge 1 commit into
patrick-kidger:mainfrom
etiferrier:loop-free-jump-search
Open

etiferrier wants to merge 1 commit into
patrick-kidger:mainfrom
etiferrier:loop-free-jump-search

Conversation

@etiferrier

Copy link
Copy Markdown

Drafted with AI assistance (Claude). Reproduced and checked by me.

Fixes #776.

_find_idx_with_hint runs on every step of a solve with ClipStepSizeController(step_ts=..., jump_ts=...). It walked from the hint with two lax.while_loops. On GPU each iteration of a while_loop costs a device-to-host sync. This PR replaces the walk with an unrolled binary search:

i = jnp.searchsorted(ts, t, side="right", method="scan_unrolled")
return i.astype(jnp.result_type(hint))

ts is already sorted (_none_or_sorted_array, and wrap), and init already uses searchsorted for the first index. So the result is the same, and only the cost changes: O(log len(ts)) fused comparisons and no loop. The hint now only sets the dtype; I kept the signature so the call sites stay the same.

Results (details and the benchmark script in #776): on an A100, a Tsit5 solve with 10–100 jump_ts is 17–30% faster. The steps and ys are bit-identical. On CPU it is neutral (×0.95–×1.01).

Independent of #773: it edits the neighbouring lines of clip.py, and the two apply together cleanly.

Tests (test/test_adaptive_stepsize_controller.py):

  • test_find_idx_with_hint also covers t below, at and above the ends of ts;
  • test_find_idx_with_hint_is_loop_free (new) checks that the jaxpr has no while, and that the result matches and keeps the hint's dtype.

Checks:

  • CI in my fork: Python 3.11 all green (pre-commit, tests, docs build). Python 3.13: tests and docs build pass; pre-commit fails only on pyright's diffrax/_misc.py:175 (unnecessary # type: ignore), which unmodified main also reports with jax 0.11.2, unrelated to this change.
  • Locally (Python 3.13, jax 0.11.2): pytest test/test_adaptive_stepsize_controller.py, 20 passed; ruff format --check and ruff check pass, and pyright reports no new errors on the changed files.

🤖 Generated with Claude Code

`_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 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

ClipStepSizeController's jump search costs two while_loops per step (slow on GPU)

1 participant