PyTorch has added an Inductor optimization that folds eligible NVFP4 output scaling into scaled matrix-multiplication implementations, enabling native NVGEMM output-scale candidates to participate in autotuning. For ModelOpt NVFP4 linear layers that scale results before QKV fan-out, the change can remove one standalone scaling launch per layer.
The rewrite recognizes an aten._scaled_mm(...) result multiplied by an output-scale tensor in either operand order. It runs after gradients but before lowering, candidate construction and scheduling, which PyTorch says preserves the graph structure needed to evaluate GEMM implementations that apply the scale directly.
The target case is a matmul result scaled before it is chunked, split or viewed into QKV branches. Folding the factor into the GEMM produces an already-scaled buffer for every consumer; leaving the operation until scheduling would require a transformation across multiple consumers. A complete pointwise suffix without that fan-out boundary is left for the existing scheduler fusion path.
Eligibility is limited to packed FP4 inputs using E4M3 1×16 block scales and a zero-dimensional CUDA FP32 multiplier. Graphs with existing bias or public scale_result are excluded, as are fast accumulation, shared outputs, incompatible scale recipes, pipelined autotuning and tensor-parallel structures that the micro-pipeline pass can use. Direct all-gather forms eligible for that micro-pipeline path remain unchanged.
For mixed ATEN and NVGEMM tuning, PyTorch retains an ATen candidate that performs the scaled matmul and multiplication separately. If native folded choices fail, lowering restores the ordinary sequence, and ineligible patterns continue through ordinary lowering.
On GB200 with Llama-3.1-8B NVFP4 at batch 64, PyTorch reports 32 fewer launches per decode step and a captured graph reduced from 481 to 449 nodes. Focused GB200 validation passed 18 output-scale cases and three tensor-parallel preservation cases. The project also cites later complete-stack spans of 2.175–2.180 ms versus 2.232–2.237 ms before the fold, but says those figures include later PDL overlap and do not isolate this change.
Definition. NVFP4 output-scale folding applies an eligible output-scale factor directly in a scaled GEMM implementation rather than launching a separate multiplication.
| Metric | Reported result |
|---|---|
| Decode-step launches | 32 fewer launches |
| Captured graph nodes | 481 to 449 |
| Complete-stack span | 2.175–2.180 ms versus 2.232–2.237 ms before the fold; PyTorch says later PDL overlap is included |
| Focused validation | 18 output-scale cases and three tensor-parallel preservation cases passed |
Key takeaways
- The rewrite recognizes an aten._scaled_mm result multiplied by an output-scale tensor in either operand order.
- It targets scaling before chunk, split, or view operations that fan out into QKV branches.
- Eligibility requires packed FP4 inputs with E4M3 1x16 block scales and a zero-dimensional CUDA FP32 multiplier.
- ATen fallback candidates retain the separate scaled matmul and multiplication sequence.
- PyTorch excludes several cases, including bias, public scale_result, fast accumulation, shared outputs, incompatible scale recipes, pipelined autotuning, and protected tensor-parallel structures.
FAQ
What does the PyTorch optimization fold?
It folds an eligible NVFP4 output-scale multiplication into a scaled matrix-multiplication implementation.
Why is QKV fan-out important?
Applying the scale in the GEMM creates an already-scaled buffer for every QKV consumer, avoiding a transformation across multiple consumers later in scheduling.
What happens when native folded candidates fail?
Lowering restores the ordinary sequence, while ineligible patterns continue through ordinary lowering.
What performance-related change did PyTorch report?
For GB200 Llama-3.1-8B NVFP4 at batch 64, PyTorch reported 32 fewer launches per decode step and a graph reduction from 481 to 449 nodes.