Pin qwix to 0.1.8 to fix FP8 quantization with current Flax - #499
Merged
copybara-service[bot] merged 1 commit intoOct 2, 2026
Merged
Conversation
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.
There was a problem hiding this comment.
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.
prishajain1
approved these changes
Oct 2, 2026
copybara-service
Bot
merged commit Oct 2, 2026
432301c
into
AI-Hypercomputer:main
27 of 28 checks passed
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
Pins
qwixto the PyPI release0.1.8, replacing the GitHub archive pin408a0f48(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 siblingwan2_1_14b_75600_4x4x4_1passes.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_shardingfromnnx.Conv.__call__, soqwix.quantize_modelfails on the first convolution (Wanpatch_embedding):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 onlyflax>=0.12.0, so no Flax or JAX version change is needed.generated_requirements/requirements.txtandbase_requirements/requirements.txt:qwix==0.1.8extra_deps_from_github.txt: drop qwix, since it now comes from PyPIdependency_versions_table.py: matching entryThe 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 .andsetup.sh(--resolution=lowest) resolve the same version.Companion change (not part of this PR)
With the qwix
TypeErrorfixed, the FP8 nightly exposes a second, independent failure: the uBenchwan_fp8template setsenable_ssim: true(the BF16 template does not). With the recipe'sreplicate_vae=True, vae_spatial=1, the pre-training SSIM sample decodes the full 1280x720x81 video on every chip and OOMs v7x HBM injit_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-sdkapt package). It is not a dependency of this PR, but it is needed to rebuild the runner image frommain.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, runcloud-tpu-gu-ubench-jnuwm3hh:step_time21.21 s. The last 5 nightlies were about 21.98 s, so this is roughly 3.5% faster.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_modelwithget_fp8_config) with the same flags aswan2_1_14b_75600_fp8_4x4x4_1:quantization=fp8_full, the sameqwix_module_pathand calibration, 1280x720x81, flash attention. The only shortcut is zero-filled weights instead of the Hugging Face download; the bug is a trace-timeTypeError, so weight values don't matter.QtProvider.conv_general_dilatedacceptsout_shardingQwix Quantization complete.in 36.8 sTypeError: 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 decodeRun
wan21-fp8-ssim-cached-1001(bodaborg-tpu7x-nap, 16 hosts × 4 chips, same flags as the nightly):Qwix Quantization complete.on all 16 workers, denoising runs — theTypeErroris gone.RESOURCE_EXHAUSTED: Ran out of memory on HBM, HLO temporaries (99.28G) exceeds available HBM (94.74G). HLO module: jit_vae_decode_passgs://ubench-logs/wan21-fp8-ssim-cached-1001/logs/4. FP8 end-to-end on Ironwood 4x4x4 with cl/992225547 applied (
enable_ssim=False) — PASSEDRun
wan21-fp8-nossim-cached-1001(same cluster, image and flags):Completed=True AllJobsCompleted; no errors in any of the 16 worker logs.jit_train_stepHBM 36.2G / 94.7G.gs://ubench-logs/wan21-fp8-nossim-cached-1001/Notes on the harness runs
wan2_1_14b_75600_fp8_4x4x4_1harness 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.toshipahadia_ironwood_runner:latestand runpip install .against the maxdiffusion copy baked into it.:latestwas 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 frommain.