[PyTorch] MXFP4 weight QAT recipes on MXFP8 and FP8 block-scaling hosts#3263
Closed
xiuhu17 wants to merge 9 commits into
Closed
[PyTorch] MXFP4 weight QAT recipes on MXFP8 and FP8 block-scaling hosts#3263xiuhu17 wants to merge 9 commits into
xiuhu17 wants to merge 9 commits into
Conversation
MXFP4QATMXFP8BlockScaling / MXFP4QATFloat8BlockScaling project weights onto the MXFP4 (E2M1, 1x32 power-of-two scale) grid before the host recipe quantizes them; activations and gradients are untouched. The projection is a fused fake-quantization with three bit-identical implementations (CUDA kernel, CuTe DSL, PyTorch reference) dispatched via NVTE_MXFP4_QAT_IMPL, with an identity-STE gradient. - scale contract matches the TileKernels deployment path (floor 2^-126, cap 2^125); integer-bit amax/scale-exponent derivation and non-FTZ PTX keep results bit-exact under --use_fast_math builds - the rowwise MXFP8 and 128x128 blockwise weight encodings of the projected weight are bitwise lossless (raw payload/scale verified); the columnwise 32x1 encoding is bounded - fix ptx.cuh exp2f expansion of UE8M0 code 0 (2^-127) and code 255 (NaN) - reject MXFP4 QAT loudly on surfaces that bypass the weight-quantization hook (te.ops, Userbuffers, fused grouped MLP, quantized primary weights); project the weight in FSDP2/GTP backward rematerialization; invalidate cached weight workspaces on base<->QAT recipe switches - tests: per-step bitwise losslessness, bf16-exhaustive and fp32 bit-fuzz oracles, RTNE midpoint/threshold vectors, fast-math immunity build, STE, misaligned/non-contiguous inputs, and a TileKernels/CuTe-DSL/CUDA/torch four-way bitwise matrix plus a 24-config e2e backward-override matrix
The MXFP4 grid reaches 6*2^125, far beyond fp16 range, so an fp16 pipeline cannot represent the projected weight or its dequantized form. quantize_weight now raises when a QAT recipe is active and the workspace/activation dtype is fp16 (fp16 weights were already rejected by the projection itself); documented in both recipe docstrings.
- exp2f JIT test now checks all 256 codes bitwise including the code-255 NaN payload (0x7FFFFFFF) - add an end-to-end nvte MXFP8 software-dequantize test with planted extreme scale codes through the real kernel: demands code 0 -> 2^-127 and code 255 -> NaN on a fixed build, and pins the pre-fix wheel behavior (flush to zero / Inf) until then - the bf16-dequant losslessness test auto-detects a fixed build instead of hard-asserting the pre-fix flush
xiuhu17
requested review from
Oleg-Goncharov,
ksivaman,
ptrendx and
timmoon10
as code owners
July 25, 2026 18:11
for more information, see https://pre-commit.ci
Contributor
Greptile SummaryThis PR adds MXFP4 weight-only quantization-aware training on MXFP8 and block-scaled FP8 hosts.
Confidence Score: 3/5The PR is not yet safe to merge because Linear can reconstruct its backward weight under a different recipe than the one used for forward. Linear does not retain the forward MXFP4-QAT state and instead queries mutable global recipe state when rebuilding FSDP2 or distributed weight workspaces during backward, so leaving the forward autocast context can cause dgrad to use an incorrectly projected weight. Files Needing Attention: transformer_engine/pytorch/module/linear.py Important Files Changed
Sequence DiagramsequenceDiagram
participant User
participant Module as PyTorch fused module
participant QAT as MXFP4 fake quantizer
participant Host as MXFP8 / FP8 block quantizer
participant GEMM
User->>Module: Forward under MXFP4-QAT recipe
Module->>QAT: Project master weight to MXFP4 grid
QAT-->>Module: BF16/FP32 projected weight
Module->>Host: Encode using host recipe
Host-->>GEMM: Quantized weight workspace
GEMM-->>User: Forward output
User->>Module: Backward
Module->>Host: Reconstruct weight workspace if needed
Host->>GEMM: Dgrad weight operand
Reviews (7): Last reviewed commit: "Drop the TileKernels cross-check tests" | Re-trigger Greptile |
NVFP4-style control: MXFP4 weight QAT is now a field on the MXFP8 and Float8 block-scaling host recipes, defaulting from NVTE_MXFP4_QAT, so a stock --fp8-recipe mxfp8/blockwise launch enables QAT without framework changes. The QAT subclasses remain as the explicit API (field pinned True); blockwise validation moves to the host, guarded on the field.
…t raw-tuple identity The host encoder re-canonicalizes (payload, scale): TE-native normalizes the block amax into (224, 448] while a fixed-shift direct converter keeps the original factorization, so the raw E4M3/UE8M0 bytes generally differ between the two even though every decoded value is identical. 'Losslessly' in the module docstring was readable as raw-bit identity; say exactly what holds.
- Move the composite-torch-ops reference implementation to tests/pytorch/references/mxfp4_qat_reference.py (the numerical oracle, matching the existing blockwise reference layout); the main module no longer carries a runtime fallback path. - transformer_engine/pytorch/mxfp4_qat.py now only validates inputs, applies the straight-through-estimator wrapper, and calls the tex.mxfp4_fake_quantize binding; a missing binding is a hard error and the NVTE_MXFP4_QAT_IMPL/REQUIRE_KERNEL/DISABLE_CUDA_KERNEL knobs are removed. - Drop the CuTe DSL implementation: the cast layer has no DSL precedent and TE carries no cutlass-DSL runtime dependency. - Tests compare the kernel against the reference directly and demand the fixed UE8M0 code-0/255 dequantize semantics unconditionally.
They imported an external TileKernels checkout with a developer-machine default path and can therefore never run in CI; the cross-backend comparison lives on as a standalone harness outside the tree. The in-tree suite keeps the kernel-vs-reference parity, exhaustive bf16 and fp32 fuzz coverage, and the raw-byte encode oracles.
8 tasks
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.
Description
Weight-only MXFP4 quantization-aware training: weights are projected onto the MXFP4 grid (E2M1 payloads, 1x32 power-of-two UE8M0 scales) before the host recipe quantizes them, so training sees exactly the deployment weight while activations and gradients keep the base recipe untouched.
The only new numerics is the fused projection
bf16/fp32 -> MXFP4 grid -> bf16/fp32(tex.mxfp4_fake_quantize: per 1x32 block, scale2^clamp(ceil(log2(amax/6)), -126, 125)derived from the amax bit pattern, RTNE onto E2M1, straight-through-estimator gradient). All MXFP8 and 128x128 blockwise FP8 encoding/decoding of the projected weight runs through the existing quantizer pipeline unchanged. Since every projected value is a grid pointp * 2^e, the existing encodings are decoded-value exact: MXFP8 rowwise unconditionally, 128x128 blockwise for tile scale spreads up to 2^14.Recipes:
MXFP4QATMXFP8BlockScaling(MXFP8BlockScaling)andMXFP4QATFloat8BlockScaling(Float8BlockScaling); QAT can also be enabled on the stock host recipes via themxfp4_qat_weightsfield (NVTE_MXFP4_QAT=1).backward_overridekeeps its base-recipe semantics. Also fixesptx::exp2f(e8m0)to decode UE8M0 code 0 as 2^-127 and code 255 as NaN.Test results (B200,
tests/pytorch/test_mxfp4_qat.py)--use_fast_mathrebuild of the same kernel sourcebackward_override{None, dequantized, high_precision} x fused/unfused wgrad; recipe-switch cache invalidation; loud rejection of unsupported surfaces (ops API, fused Userbuffers/grouped-MLP, fp16, primary FP8 weights)Known limitations
Type of change
Checklist: