Repository navigation
Conversation
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
prishajain1
approved these changes
Oct 6, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Raises the
flaxfloor from>=0.12.6to>=0.12.10so images built frommaincan importflax.nnx(and therefore maxdiffusion) under JAX >= 0.11.Bug: b/537854580 (Ironwood Wan 2.1 nightlies). Follow-up to #499 / #490.
Issue
setup.shinstallsgenerated_requirements.txtwithuv pip install -U --resolution=lowest, so the floor written in that file is exactly the version that ships in the image. The floor wasflax>=0.12.6. That release still subclassesjax.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 importflax.nnx:Every maxdiffusion entry point imports
flax.nnx, so the image fails at startup for every workload (BF16 and FP8 alike).Why it looked intermittent
toshipahadia_ironwood_runner:qwix018(2026-09-30)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)import flax.nnx→AttributeErrorSame 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=lowestalways yields a JAX-0.11-compatible flax.Why 0.12.10
jax.core.Effectjax>=0.8.1jax>=0.10.0jax>=0.11.10.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=lowestthe 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.10dependencies/requirements/base_requirements/requirements.txt:flax→flax>=0.12.10src/maxdiffusion/dependency_versions_table.py: matching entryHand-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(namespaceubench-regression-tests), TPU v7x, FP8 Wan 2.1 14B recipe flags (wan2_1_14b_75600_fp8_4x4x4_1,enable_ssim=Falseper cl/992225547,max_train_steps=30).1. Official image from
main(flax 0.12.6 + jax 0.11.2) — FAILS at importImage:
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)EXIT_CODE=1, JobSetFailed=Trueat 09:43:32 UTC. Weights staged fine ([stage] done in 14s; 75G staged), then:gs://ubench-logs/wan21-fp8-off-0126-1host/logs/train-0.log· Cloud Logging2. 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 Build0b9c91cdfrom this branch, MODE=nightly, same Dockerfiles as the official image; build includes a sanity step that printedSANITY 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 UTCCompleted=True AllJobsCompleted, 16/16 podsCompleted, 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:qwix018verification run (wan21-fp8-nossim-cached-1001), as expected since both images carry flax 0.12.10.gs://ubench-logs/wan21-fp8-ff-01210-v3/logs/train-{0..15}.log(JAX process 0 =train-4.log) · Cloud Logging