Repository navigation
jax 0.11.1: XLA:CPU dynamic-update-slice-in-loop regression (why jax is pinned to 0.11.0) #622
Description
Activity
Filed upstream as jax-ml/jax#40101 — the full report with the minimal reproduction, the variant matrix, the O(buffer x trips) sweeps, the 256-byte cliff, the one-build nightly bisect, and the tentative XLA candidate commit.
That thread is now the one to watch for the unpin: the gate is a jax release whose jitted
lax.fori_loopat n=400,000 on CPU completes in single-digit seconds (0.11.0 measures ~1.6 s including compilation; 0.11.1 measures ~100 s).Status update on the timeline above: #621 merged (
c2589a2, 2026-08-20T04:11Z), so main now carries the pinned install cell. Translation sync fired on the merge and opened the mirrored one-line PRs: lecture-python-programming.fr#39, .fa#161, .zh-cn#96.Merging is not publishing: the live site and the generated notebook mirror still ship the unpinned cell from the 2026-08-18 build, so Colab readers remain exposed until a
publish*tag runs. That publish is the last reader-facing step of this incident.The reader gap is now shut.
publish-2026aug20(run 32331399540, headc2589a2) completed at 2026-08-20T04:2x UTC withSuccessfully installed jax-0.11.0 jaxlib-0.11.0andnumpy_vs_numba_vs_jax.md: Executed notebook in 30.99 seconds— an actual execution, not a cache restore, and incidentally the first time this lecture has executed on the GPU runner at any point in this incident.Verified on the live site after the deploy:
_notebooks/numpy_vs_numba_vs_jax.ipynb(200, 28,752 B) carries!pip install quantecon "jax==0.11.0"and zero occurrences of the old unpinned form; the generated mirror at lecture-python-programming.notebooks is byte-identical, so the Colab launch button now serves the pin. The translation sync PRs (.fr#39, .fa#161, .zh-cn#96) are merged, and the fr/fa cache builds went green on the merge shas under their pinned workflows.Every surface of this incident is now mitigated: 15 pinned workflow files across the family, the pinned lecture cell in all four repos, and a published site + mirror carrying it. What remains is upstream: jax-ml/jax#40101 is the thread to watch, and the unpin gate in this issue's body is the test to run when a candidate release appears.
Correction to two lines above, measured during an independent validation of this incident on hardware the original diagnosis never touched.
Everything in this issue was measured in
python:3.13-slimcontainers on Apple silicon —linux/amd64emulated,linux/aarch64native — plus macOS arm64. Re-running the whole thing on a native x86-64 GitHub-hosted runner (Linux-6.17.0-1022-azure-x86_64-with-glibc2.39, CPython 3.13.15, AVX-512, no container) reproduces the diagnosis but not the 256-byte figure.Sweeping by bytes at B=200,000, m=10,000 under 0.11.1, with 0.11.0 as the control:
bytes float32 W float64 W float16 W 0.11.0 0.11.1 224 56 28 112 ~0.0002 ~0.88 slow 256 64 32 128 ~0.0003 ~0.88 slow 384 96 48 192 ~0.0002 ~0.88 slow 448 112 56 224 ~0.0002 ~0.88 slow 480 120 60 240 ~0.0007 ~0.0007 fast 512 128 64 256 ~0.0006 ~0.0007 fast So the two corrections:
- "a sharp cliff at exactly 256 bytes of update width" — on x86-64 the threshold is between 448 and 480 bytes. The parenthetical that follows it survives and is if anything strengthened: the boundary lands at the same byte count in float16, float32 and float64, so it is a byte threshold and not an element count. Only the value moves.
- "batching writes to >= 256 bytes per iteration" — this mitigation is unsafe as written on x86, where 256 bytes per iteration is still squarely on the slow path (float32 W=64 measures 0.876 s against a fast path of ~0.0007 s). The safe figure there is >= 512 bytes.
Not a harness artefact: three independent DUS constructions were swept — loop-invariant update at unaligned (stride-7) offsets, loop-invariant at W-aligned offsets, and an update computed inside the loop body — and all three put 256 B on the slow side and >= 512 B on the fast side. Under 0.11.0 the same boundary shows as a mild step (~0.0002 s -> ~0.0007 s), i.e. it is where XLA switches DUS emission strategy in both versions, with 0.11.1 regressing only on the below-threshold path. That it tracks a code-path switch rather than a constant makes a vector-width dependence the natural hypothesis — 256 B is 16x128-bit NEON, ~480-512 B is ~8x512-bit AVX-512 — which would explain why the aarch64 and emulated-amd64 measurements agreed with each other and disagree with native x86. That hypothesis is untested.
Everything else transferred, so nothing about the pin decision changes:
- Repro block on 0.11.1:
fori1.9629 / 7.7821 / 31.0894 s at n = 100k/200k/400k (x3.96, x3.99 per doubling),scan_stacked4.0193 / 16.1851 / 64.7680 s (x4.03, x4.00),scan_nostackflat at 0.0005 / 0.0010 / 0.0019 s. 0.11.0 flat at 0.0005-0.0034 s throughout. - O(buffer x trips): fixed m=10,000, B = 10k/50k/200k/1M -> 0.0027 / 0.0143 / 0.0587 / 0.3399 s; fixed B=200,000, m = 5k/10k/20k/40k -> 0.0283 / 0.0567 / 0.1130 / 0.2265 s, exactly x2 per doubling.
- Bisect:
0.11.1.dev202607250.0011 s vs0.11.1.dev202607267.8933 s at n=200,000. - jaxlib and not jax: jax
dev20260717+ jaxlibdev20260816-> 7.1974 s; the reverse pairing -> 0.0011 s. - Magnitudes came within a factor of 1.7 of the emulated numbers (
foriat n=200,000: 7.78 s native vs 13.22 s here), so emulation did not distort the scaling — it cost this report only the threshold value.
The unpin condition is unaffected and was re-measured on the same native runner, since it sits nowhere near the threshold: 0.11.0 = 0.003 s, 0.11.1 = 31.561 s, nightly
0.11.2.dev20260819= 31.628 s — 100.2% of 0.11.1. The regression is still on jax main, so==0.11.0remains the right pin over!=0.11.1, exactly as recorded above.A hardware footnote carrying these numbers has been added to the upstream report at jax-ml/jax#40101, so a maintainer reproducing from its cliff table on x86 is not misled by the "first fast W" rows.
This issue is the public record of why every jax install in this repo is pinned to
0.11.0, what the underlying jax/XLA regression actually is, and what has to be true before the pins can lift. It exists so the upstream report has a public reference for where the bug was found and how it was diagnosed.The regression, in one paragraph
Since jax/jaxlib 0.11.1 (PyPI 2026-08-17), a
dynamic-update-slicethat writes a small slice (< 256 bytes) into a large array inside a loop body costs O(whole buffer) per iteration on the CPU backend instead of O(update). In the standardlax.fori_loop/lax.scanaccumulation idiom the buffer length equals the trip count, so runtime becomes quadratic in the loop length — ~4x per doubling of n. This lecture'snumpy_vs_numba_vs_jax.mdruns exactly that idiom at n = 10,000,000: ~0.06 s under 0.11.0, extrapolated ~18 hours under 0.11.1, so every executing build dies onCellTimeoutErrorat the 600 s myst-nb limit.How it surfaced and what landed
-W, so five timeout builds concluded green before the publish path failed); en's dailyexecution-linux.ymlwent red the same dayjax[cuda13]==0.11.0incache.yml/ci.yml/publish.yml(#617); translations pinnedjax==0.11.0in the same three workflows each (lecture-python-programming.fr#37, .fa#157, .zh-cn#95)execution-{linux,osx,win}.ymlpinned via #620, after which a dispatched execution run installed 0.11.0 and executed the lecture in 5.84 s; #621 opened to pin the lecture's own!pip installcell, which is what Colab readers execute (in CI that cell is a no-op because jax is pre-installed; in Colab it resolves 0.11.1)All six workflows in this repo and all nine translation workflows now pin
0.11.0. Pinning jax alone pins jaxlib too:jax==0.11.0requiresjaxlib<=0.11.0,>=0.11.0on the base requirement and on every accelerator extra, and the pairing is enforced again at import in both directions.Root cause, condensed
The full report with the reproduction script and measurement tables is being filed upstream at jax-ml/jax (link to follow in a comment). The short version, all measured:
0.11.1.dev20260725good (n=200,000 fori 0.0014 s),0.11.1.dev20260726bad (13.24 s) — ~10,000x apart, monotonic on both sides, reproduced on linux/x86_64 and linux/aarch64.6b5d5254...->88e9a7db..., a window of exactly 10 commits, 4 tagged[XLA:CPU].fori_loop,scanandwhile_loopwith scalar bodies are unaffected on 0.11.1; the identicalscanwith its stackedysoutput discarded is unaffected; only bodies writing a small slice into a large buffer regress. Cost is O(buffer x trips) with a sharp cliff at exactly 256 bytes of update width (holds across float16/32/64 — a byte threshold, not an element count).--xla_cpu_use_fusion_emitterswas deprecated inside the same XLA window, and the surviving related flags measurably change nothing. The only mitigations are pinning, discarding stacked outputs, or batching writes to >= 256 bytes per iteration.Unpin condition
A jax release must ship whose jitted
lax.fori_loopat n=400,000 on CPU completes in ~single-digit seconds rather than ~100 s (the two sides differ by ~65x, so the gate is unambiguous; a validated one-command container test exists). As of 2026-08-20 no such release exists — 0.11.1 is latest, and the nightly0.11.2.dev20260819still carries the regression at ~96% of 0.11.1's runtime, so a hypothetical 0.11.2 cut from current main would ship broken. When the upstream fix lands, unpin the six workflows here, the nine translation workflows, and the lecture's install cell together.Diagnosed with assistance from Anthropic's Claude models (Opus 5 and Fable 5).