Skip to content

Raise flax floor to 0.12.10 so nightly images import under JAX >= 0.11 - #502

Open
Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-flax-floor-0.12.10
Open

Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-flax-floor-0.12.10

Conversation

@Toshi-31

@Toshi-31 Toshi-31 commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Raises the flax floor from >=0.12.6 to >=0.12.10 so images built from main can import flax.nnx (and therefore maxdiffusion) under JAX >= 0.11.

Bug: b/537854580 (Ironwood Wan 2.1 nightlies). Follow-up to #499 / #490.

Issue

setup.sh installs generated_requirements.txt with uv pip install -U --resolution=lowest, so the floor written in that file is exactly the version that ships in the image. The floor was flax>=0.12.6. That release still subclasses jax.core.Effect (flax/nnx/variablelib.py:289), which JAX removed in 0.11.0. The nightly step then upgrades JAX to 0.11.x, and the result cannot import flax.nnx:

File "/usr/local/lib/python3.12/site-packages/flax/nnx/variablelib.py", line 289, in <module>
    class VariableEffect(jax.core.Effect): ...
AttributeError: jax.core.Effect  was deprecated in JAX v0.10.0 and removed in JAX v0.11.0. Use jax.extend.core.Effect.

Every maxdiffusion entry point imports flax.nnx, so the image fails at startup for every workload (BF16 and FP8 alike).

Why it looked intermittent

Image Built from flax resolved jax Result
toshipahadia_ironwood_runner:qwix018 (2026-09-30) #499 branch with the floor raised locally 0.12.10 0.11.2 Wan 2.1 BF16 + FP8 pass on 4x4x4 (cloud-tpu-gu-ubench-jnuwm3hh, wan21-fp8-nossim-cached-1001)
maxdiffusion_jax_nightly:2026-10-03/04/05 (official, post #490+#499) main (flax>=0.12.6) 0.12.6 0.11.2 import flax.nnx → AttributeError

Same Dockerfile, same JAX, same qwix 0.1.8. The only difference is which floor the resolver was given. So the nightly "sometimes passes, sometimes fails" depending on which tree the image was built from, not on the day. This PR removes that dependency: with the floor at 0.12.10, --resolution=lowest always yields a JAX-0.11-compatible flax.

Why 0.12.10

flax uses jax.core.Effect declares
0.12.6 yes jax>=0.8.1
0.12.7 / 0.12.8 no jax>=0.10.0
0.12.9 / 0.12.10 no jax>=0.11.1

0.12.7 is the first release that imports, but 0.12.10 is the version that has actually been verified end-to-end on Ironwood hardware (see Testing), and under --resolution=lowest the floor is the shipped version — so the floor should be the tested one.

Changes

  • dependencies/requirements/generated_requirements/requirements.txt: flax>=0.12.6 → flax>=0.12.10
  • dependencies/requirements/base_requirements/requirements.txt: flax → flax>=0.12.10
  • src/maxdiffusion/dependency_versions_table.py: matching entry

Hand-edited to avoid seed-env churn, same as #499. No code changes; qwix stays at 0.1.8 (it requires only flax>=0.12.0).

Testing

All runs on the uBench cluster bodaborg-tpu7x-nap (namespace ubench-regression-tests), TPU v7x, FP8 Wan 2.1 14B recipe flags (wan2_1_14b_75600_fp8_4x4x4_1, enable_ssim=False per cl/992225547, max_train_steps=30).

1. Official image from main (flax 0.12.6 + jax 0.11.2) — FAILS at import

Image: gcr.io/tpu-prod-env-multipod/maxdiffusion_jax_nightly@sha256:3f7baa3c05e301a7bd8953676282a8a194a71c659a9ce64ba3a5ac5ce27d9aac (= :latest / :2026-10-05)
Run: wan21-fp8-off-0126-1host (2026-10-05 09:42:39–09:43:29 UTC; single tpu7x host — the failure is at import, before any device/topology code runs, so 1 host is sufficient and avoided tying up a 16-host slice while the cluster was capacity-constrained)

  • Result: FAILED, EXIT_CODE=1, JobSet Failed=True at 09:43:32 UTC. Weights staged fine ([stage] done in 14s; 75G staged), then:
    File "/deps/src/maxdiffusion/train_wan.py", line 37, in train
        from maxdiffusion.trainers.wan_trainer import WanTrainer
    File ".../maxdiffusion/trainers/wan_trainer.py", line 19, in <module>
        from flax import nnx
    ...
    File "/usr/local/lib/python3.12/site-packages/flax/nnx/variablelib.py", line 289, in <module>
        class VariableEffect(jax.core.Effect): ...
    AttributeError: jax.core.Effect  was deprecated in JAX v0.10.0 and removed in JAX v0.11.0. Use jax.extend.core.Effect.
    
  • Logs: gs://ubench-logs/wan21-fp8-off-0126-1host/logs/train-0.log · Cloud Logging

2. Image built from this PR (flax 0.12.10 + jax 0.11.2) — PASSES

Image: gcr.io/cloud-tpu-multipod-dev/toshipahadia_ironwood_runner:flaxfloor (sha256:b0b84925afe2c9f7fd59e31f78c1c1ead9c5b48d286fd5f9cafb7f0afd6bb4e5; Cloud Build 0b9c91cd from this branch, MODE=nightly, same Dockerfiles as the official image; build includes a sanity step that printed SANITY OK: flax 0.12.10, jax 0.11.2; flax.nnx + qwix + wan pipeline/trainer import)
Run: wan21-fp8-ff-01210-v3 — tpu7x 4x4x4 (16 hosts × 4 chips), weights pre-staged from GCS, 2026-10-05 09:35:42–09:51:58 UTC

  • Result: PASSED — JobSet Completed=True AllJobsCompleted, 16/16 pods Completed, 0 tracebacks in any of the 16 logs. Steps 0–28 ran; steady state (steps 2–28, n=27): 20.77 s/step, 235.6 TFLOP/s/device — identical to the earlier :qwix018 verification run (wan21-fp8-nossim-cached-1001), as expected since both images carry flax 0.12.10.
    completed step: 0, seconds: 26.170, TFLOP/s/device: 186.979, loss: 2.831
    completed step: 1, seconds: 22.080, TFLOP/s/device: 221.613, loss: 2.961
    completed step: 2, seconds: 20.763, TFLOP/s/device: 235.672, loss: 2.780
    ...
    completed step: 28, seconds: 20.738, TFLOP/s/device: 235.958, loss: 2.819
    
  • Logs: gs://ubench-logs/wan21-fp8-ff-01210-v3/logs/train-{0..15}.log (JAX process 0 = train-4.log) · Cloud Logging

Deployment note: the Ironwood Guitar recipes currently pin toshipahadia_ironwood_runner:latest (cl/988719265, a temporary measure while the official image build was broken). Once this PR merges and maxdiffusion_jax_nightly:latest is rebuilt from main, that CL can be rolled back so the nightlies run on the official image.

setup.sh installs generated_requirements.txt with `uv pip install
--resolution=lowest`, so whatever floor is written there is exactly what
ships in the image. The floor was `flax>=0.12.6`; that release still
subclasses `jax.core.Effect` (flax/nnx/variablelib.py:289), which JAX
removed in 0.11.0. The nightly step then upgrades JAX to 0.11.x, and the
result cannot import flax.nnx at all:

  AttributeError: jax.core.Effect was deprecated in JAX v0.10.0 and
  removed in JAX v0.11.0. Use jax.extend.core.Effect.

Every maxdiffusion entry point imports flax.nnx, so the official nightly
image built from main (maxdiffusion_jax_nightly:2026-10-03 .. :2026-10-05,
flax 0.12.6 + jax 0.11.2) fails at startup for BF16 and FP8 alike. Images
built from a tree where the floor had been raised resolved flax 0.12.10
and passed the Ironwood 4x4x4 Wan 2.1 runs, which is why the breakage
looked intermittent: it depends on which floor the image was built from,
not on the day.

flax 0.12.7 is the first release without jax.core.Effect; 0.12.9+ declare
jax>=0.11.1. 0.12.10 is the version verified end-to-end on hardware, and
with --resolution=lowest the floor is the shipped version, so pin the
floor to the tested one.

Hand-edited (generated + base requirements, deps table) to avoid seed-env
churn, same as AI-Hypercomputer#499.

b/537854580
@Toshi-31
Toshi-31 requested a review from entrpn as a code owner October 5, 2026 09:58

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request updates the dependency requirement for flax to version 0.12.10 or higher across the base requirements, generated requirements, and the dependency versions table. There are no review comments, and I have no feedback to provide.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants