Skip to content

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

Open
Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-qwix-0.1.8
Open

Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:bump-qwix-0.1.8

Conversation

@Toshi-31

Copy link
Copy Markdown
Collaborator

Summary

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

Why

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. For example, the Ironwood nightly wan2_1_14b_75600_fp8_4x4x4_1 (b/537854580) has never passed.

qwix==0.1.8 (2026-06-22) includes the fix (QtProvider.conv_general_dilated(..., out_sharding=None)) and requires only flax>=0.12.0.

Changes

  • 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.

Testing

Validation is in progress: a runner image built from this branch, run through the Ironwood ubench FP8 workload. Results will be posted here.

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

Copy link
Copy Markdown
Collaborator Author

Test results

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

Test wan2_1_14b_75600_4x4x4_1: PASSED

  • 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)

This runs the production code path WanPipeline.load_transformer → WanPipeline.quantize_transformer (qwix.quantize_model with get_fp8_config). It uses 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.

Note on the full FP8 nightly-harness run

I tried the full wan2_1_14b_75600_fp8_4x4x4_1 harness run several times today. It never reached training: the 16 hosts stalled on anonymous Hugging Face downloads (text encoder and transformer) during daytime runs. That hang happens before any qwix code runs and is unrelated to this change. The real nightly does get past the downloads (that's how it hits the TypeError), so the nightly after merge will be the end-to-end confirmation.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant