Skip to content

jax 0.11.1: XLA:CPU dynamic-update-slice-in-loop regression (why jax is pinned to 0.11.0) #622

Description

@mmcky

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-slice that 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 standard lax.fori_loop / lax.scan accumulation idiom the buffer length equals the trip count, so runtime becomes quadratic in the loop length — ~4x per doubling of n. This lecture's numpy_vs_numba_vs_jax.md runs 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 on CellTimeoutError at the 600 s myst-nb limit.

How it surfaced and what landed

date (UTC) event
2026-08-17 20:29 jax/jaxlib 0.11.1 hit PyPI
2026-08-18 fr/fa translation cache builds began timing out (their builds omit -W, so five timeout builds concluded green before the publish path failed); en's daily execution-linux.yml went red the same day
2026-08-19 bisected to the jax version and pinned jax[cuda13]==0.11.0 in cache.yml / ci.yml / publish.yml (#617); translations pinned jax==0.11.0 in the same three workflows each (lecture-python-programming.fr#37, .fa#157, .zh-cn#95)
2026-08-20 the three remaining unpinned workflows execution-{linux,osx,win}.yml pinned 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 install cell, 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.0 requires jaxlib<=0.11.0,>=0.11.0 on 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:

  • Bisected to a single nightly build: 0.11.1.dev20260725 good (n=200,000 fori 0.0014 s), 0.11.1.dev20260726 bad (13.24 s) — ~10,000x apart, monotonic on both sides, reproduced on linux/x86_64 and linux/aarch64.
  • It is jaxlib, not jax: mixing nightly packages, runtime tracks the jaxlib version only, and the optimized HLO is byte-identical between fast and slow pairings. The boundary corresponds to the XLA roll 6b5d5254... -> 88e9a7db..., a window of exactly 10 commits, 4 tagged [XLA:CPU].
  • The trigger is the DUS, not the loop: fori_loop, scan and while_loop with scalar bodies are unaffected on 0.11.1; the identical scan with its stacked ys output 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).
  • No flag-level workaround exists: --xla_cpu_use_fusion_emitters was 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_loop at 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 nightly 0.11.2.dev20260819 still 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).

Activity

  1. mmcky commented on Aug 20, 2026

    @mmcky
    ContributorAuthor

    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_loop at 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).

  2. mmcky commented on Aug 20, 2026

    @mmcky
    ContributorAuthor

    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.

  3. mmcky commented on Aug 20, 2026

    @mmcky
    ContributorAuthor

    The reader gap is now shut. publish-2026aug20 (run 32331399540, head c2589a2) completed at 2026-08-20T04:2x UTC with Successfully installed jax-0.11.0 jaxlib-0.11.0 and numpy_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.

  4. mmcky commented on Aug 20, 2026

    @mmcky
    ContributorAuthor

    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-slim containers on Apple silicon — linux/amd64 emulated, linux/aarch64 native — 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: fori 1.9629 / 7.7821 / 31.0894 s at n = 100k/200k/400k (x3.96, x3.99 per doubling), scan_stacked 4.0193 / 16.1851 / 64.7680 s (x4.03, x4.00), scan_nostack flat 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.dev20260725 0.0011 s vs 0.11.1.dev20260726 7.8933 s at n=200,000.
    • jaxlib and not jax: jax dev20260717 + jaxlib dev20260816 -> 7.1974 s; the reverse pairing -> 0.0011 s.
    • Magnitudes came within a factor of 1.7 of the emulated numbers (fori at 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.0 remains 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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions