Skip to content

Pin qwix to 0.1.8 to fix FP8 quantization with current Flax - #499

Merged
copybara-service[bot] merged 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-qwix-0.1.8
Oct 2, 2026
Merged

copybara-service[bot] merged 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-qwix-0.1.8

Conversation

@Toshi-31

@Toshi-31 Toshi-31 commented Sep 30, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Pins qwix to the PyPI release 0.1.8, replacing the GitHub archive pin 408a0f48 (Dec 2025).

Bug: b/537854580 — the Ironwood nightly wan2_1_14b_75600_fp8_4x4x4_1 (Wan 2.1 14B FP8, v7x 4x4x4) has never passed, while the BF16 sibling wan2_1_14b_75600_4x4x4_1 passes.

Issue

The current pin predates google/qwix@98a44ed0 (2026-01-07, "Add out_sharding parameter to Qwix conv_general_dilated functions"). Flax >= 0.12.6 (our minimum; the images ship 0.12.9) always passes out_sharding from nnx.Conv.__call__, so qwix.quantize_model fails on the first convolution (Wan patch_embedding):

wan_pipeline.quantize_transformer -> qwix.quantize_model
  -> transformer_wan.py:789 self.patch_embedding(...)
    -> flax/nnx/nn/linear.py:889 self.conv_general_dilated(..., out_sharding=None)
TypeError: QtProvider.conv_general_dilated() got an unexpected keyword argument 'out_sharding'

This breaks every qwix-quantized run. BF16 runs never call qwix.quantize_model, which is why the BF16 nightly passes and the FP8 nightly does not.

Resolution

qwix==0.1.8 (2026-06-22) includes the fix (QtProvider.conv_general_dilated(..., out_sharding=None)) and requires only flax>=0.12.0, so no Flax or JAX version change is needed.

  • generated_requirements/requirements.txt and base_requirements/requirements.txt: qwix==0.1.8
  • extra_deps_from_github.txt: drop qwix, since it now comes from PyPI
  • dependency_versions_table.py: matching entry

The generated requirements and the deps table are edited by hand. Re-running seed-env would re-resolve every package and add unrelated churn. An exact pin is used so that pip install . and setup.sh (--resolution=lowest) resolve the same version.

Companion change (not part of this PR)

With the qwix TypeError fixed, the FP8 nightly exposes a second, independent failure: the uBench wan_fp8 template sets enable_ssim: true (the BF16 template does not). With the recipe's replicate_vae=True, vae_spatial=1, the pre-training SSIM sample decodes the full 1280x720x81 video on every chip and OOMs v7x HBM in jit_vae_decode_pass (99.28G temporaries > 94.74G). Fixed by cl/992225547, which removes that line. Both changes are needed for the nightly to pass; see Testing §3 and §4.

Related: #490 fixes the Dockerfile build (retired google-cloud-sdk apt package). It is not a dependency of this PR, but it is needed to rebuild the runner image from main.

Testing

Image: gcr.io/cloud-tpu-multipod-dev/toshipahadia_ironwood_runner:qwix018, built from this PR plus #490's Dockerfile fixes.
Versions: qwix 0.1.8, flax 0.12.10, jax/jaxlib 0.11.2, libtpu 0.0.49.

1. BF16 regression check: internal Ironwood nightly benchmark harness (XPK), Ironwood 4x4x4 — PASSED

Test wan2_1_14b_75600_4x4x4_1, run cloud-tpu-gu-ubench-jnuwm3hh:

  • 30 steps, step_time 21.21 s. The last 5 nightlies were about 21.98 s, so this is roughly 3.5% faster.
  • MFU 0.200, 461 TFLOP/s.
  • The new flax, jax and qwix versions don't break the currently passing Wan 2.1 test.

2. FP8 (fp8_full): direct A/B of the failing code path on TPU (tpu7x-2x2x1, same image)

Runs the production code path WanPipeline.load_transformer → WanPipeline.quantize_transformer (qwix.quantize_model with get_fp8_config) with the same flags as wan2_1_14b_75600_fp8_4x4x4_1: quantization=fp8_full, the same qwix_module_path and calibration, 1280x720x81, flash attention. The only shortcut is zero-filled weights instead of the Hugging Face download; the bug is a trace-time TypeError, so weight values don't matter.

qwix QtProvider.conv_general_dilated accepts out_sharding Result
0.1.8 (this PR) yes Qwix Quantization complete. in 36.8 s
408a0f48 (current pin) no TypeError: QtProvider.conv_general_dilated() got an unexpected keyword argument 'out_sharding'. This is the exact nightly failure.

3. FP8 end-to-end on Ironwood 4x4x4, recipe as-is (enable_ssim=True) — gets past qwix, then OOMs in VAE decode

Run wan21-fp8-ssim-cached-1001 (bodaborg-tpu7x-nap, 16 hosts × 4 chips, same flags as the nightly):

  • Weights load, Qwix Quantization complete. on all 16 workers, denoising runs — the TypeError is gone.
  • All 16 workers then fail in the pre-training SSIM sample:
    RESOURCE_EXHAUSTED: Ran out of memory on HBM, HLO temporaries (99.28G) exceeds available HBM (94.74G). HLO module: jit_vae_decode_pass
  • This is the second failure described above, fixed by cl/992225547. Logs: gs://ubench-logs/wan21-fp8-ssim-cached-1001/logs/

4. FP8 end-to-end on Ironwood 4x4x4 with cl/992225547 applied (enable_ssim=False) — PASSED

Run wan21-fp8-nossim-cached-1001 (same cluster, image and flags):

  • JobSet Completed=True AllJobsCompleted; no errors in any of the 16 worker logs.
  • 29/29 steps. Steady-state (steps 2–28, n=27): 20.77 s/step (σ ≈ 0.02 s), 235.6 TFLOP/s/device.
  • jit_train_step HBM 36.2G / 94.7G.
  • vs the BF16 reference in §1: 20.77 s vs 21.21 s (−2.1%).
  • Logs + XProf trace: gs://ubench-logs/wan21-fp8-nossim-cached-1001/
completed step: 0,  seconds: 82.076, TFLOP/s/device:  59.619, loss: 2.949   # compile
completed step: 1,  seconds: 72.495, TFLOP/s/device:  67.498, loss: 3.103   # compile
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

Notes on the harness runs

  • Earlier attempts at the full wan2_1_14b_75600_fp8_4x4x4_1 harness run stalled on Hugging Face downloads (text encoder and transformer) on all 16 hosts during daytime; that hang happens before any qwix code runs and is unrelated to this change. For §3 and §4 the weights were pre-staged to local disk to avoid it; this is test setup only.
  • The Ironwood Guitar recipes pin the prebuilt image toshipahadia_ironwood_runner:latest and run pip install . against the maxdiffusion copy baked into it. :latest was rebuilt from this branch on 2026-09-30 (same digest as :qwix018), so the nightly already carries this pin; after merge the image should be rebuilt from main.

The qwix pin (commit 408a0f48, Dec 2025) predates google/qwix@98a44ed0
(2026-01-07), which added `out_sharding` to QtProvider.conv_general_dilated.
Flax >= 0.12.6 (our minimum) always passes `out_sharding` from nnx.Conv, so
qwix.quantize_model fails on the first conv (Wan patch_embedding) with:

  TypeError: QtProvider.conv_general_dilated() got an unexpected keyword
  argument 'out_sharding'

This breaks every qwix-quantized run (e.g. Wan 2.1 FP8 training,
b/537854580). qwix 0.1.8 (PyPI, 2026-06-22) includes the fix and requires
only flax>=0.12.0.

Switching from a GitHub archive URL to a PyPI version also keeps direct URL
references out of the package metadata. The generated requirements and deps
table are edited by hand to avoid unrelated churn from re-running seed-env.
@Toshi-31
Toshi-31 requested a review from entrpn as a code owner September 30, 2026 06:51

@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 qwix dependency from a GitHub archive URL to a pinned PyPI version (qwix==0.1.8) across multiple requirements files and the dependency versions table, while also removing it from the extra GitHub dependencies list. There are no review comments, and I have no feedback to provide.

@Toshi-31
Toshi-31 requested a review from prishajain1 October 2, 2026 11:06
@copybara-service
copybara-service Bot merged commit 432301c into AI-Hypercomputer:main Oct 2, 2026
27 of 28 checks passed
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