Repository navigation
[Common] Fuse RMSNorm forward with MXFP8 quantization - #3657
Conversation
RMSNorm forward with MXFP8 output used cuDNN's fused kernel or, in PyTorch without NVTE_NORM_FWD_USE_CUDNN, a BF16 normalization followed by a separate quantization of the rounded values. cuDNN's kernel runs 128 CTAs for 4096 rows and cannot write GEMM-swizzled scaling factors, so a swizzle kernel follows before the GEMM. Add a fused kernel. A pre-pass computes rsigma per row, then one CTA per 32x128 tile normalizes in FP32, quantizes row- and/or column-wise, and writes the scaling factors in compact or GEMM-swizzled layout. The output is the MXFP8 quantizer applied to the FP32 normalized values, as cuDNN computes it; the two can differ only through the rounding of rsigma. Use the kernel for RMSNorm with MXFP8 output whenever both dimensions are multiples of 128, and let the PyTorch binding request swizzled scaling factors in that case. NVTE_NORM_FWD_MXFP8_USE_CUDNN=1, or nvte_enable_cudnn_norm_fwd_mxfp8(true), restores the previous paths. Signed-off-by: Siddhartha Raman S <sraman@nvidia.com> Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
timmoon10
left a comment
There was a problem hiding this comment.
Have we profiled this kernel? I estimate its SoL perf to be 0.69x worse than the cuDNN kernel.
…tion.h Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: Siddhartha Raman Sundara Raman <sraman@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: Siddhartha Raman Sundara Raman <sraman@nvidia.com>
Address review of the fused RMSNorm + MXFP8 forward: - Select the kernel in the suggested order: GEMM-swizzled scales always use it, use_cudnn_norm_fwd() selects cuDNN, tensors the kernel does not support fall back to cuDNN, and otherwise the kernel runs. Drop the NVTE_NORM_FWD_MXFP8_USE_CUDNN flag and nvte_enable_cudnn_norm_fwd_mxfp8. - PyTorch: add Impl::FUSED_NORM_QUANT_UNSWIZZLED for cuDNN, and keep MXFP8 quantizers with 2D quantization on the unfused path, since both fused kernels quantize 1D blocks. - Honor the SM margin: with SMs reserved, launch the kernels over chunks of 128-row groups that fit on the remaining SMs. Without a margin, a single launch covers the tensor, as before. - Tests: switch backends with nvte_enable_cudnn_norm_fwd, compute the swizzled output on a single SM, and cover an SM margin and 2D quantization in PyTorch. Signed-off-by: Siddhartha Raman S <sraman@nvidia.com> Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
| if (!mxfp8_quantizer_cpp->with_2d_quantization && outer_size % 128 == 0 && | ||
| inner_size % 128 == 0) { | ||
| // cuDNN MXFP8 kernel requires full 128x128 tiles | ||
| impl = Impl::FULLY_FUSED; | ||
| if (transformer_engine::getenv<bool>("NVTE_NORM_FWD_USE_CUDNN")) { | ||
| impl = Impl::FUSED_NORM_QUANT_UNSWIZZLED; | ||
| } else { | ||
| // Transformer Engine's fused RMSNorm + MXFP8 kernel | ||
| impl = Impl::FULLY_FUSED; |
There was a problem hiding this comment.
With cuDNN disabled, shapes whose dimensions are multiples of 128 now select FULLY_FUSED even for MXFP8Quantizer(rowwise=False, columnwise=True). That quantizer leaves z->data.shape at {0}, but nvte_rmsnorm_fwd requires it to match the input shape before reaching the new kernel. Forward therefore throws during the workspace query instead of returning column-scaled output.
Keep column-only quantizers on UNFUSED, or update the common API to check the logical output shape.
|
/te-ci L1 Pytorch |
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
for more information, see https://pre-commit.ci
|
/te-ci L1 pytorch |
RMSNorm forward with MXFP8 output used cuDNN's fused kernel or, in PyTorch without NVTE_NORM_FWD_USE_CUDNN, a BF16 normalization followed by a separate quantization of the rounded values. cuDNN's kernel runs 128 CTAs for 4096 rows and cannot write GEMM-swizzled scaling factors, so a swizzle kernel follows before the GEMM.
Add a fused kernel. A pre-pass computes rsigma per row, then one CTA per 32x128 tile normalizes in FP32, quantizes row- and/or column-wise, and writes the scaling factors in compact or GEMM-swizzled layout. The output is the MXFP8 quantizer applied to the FP32 normalized values, as cuDNN computes it; the two can differ only through the rounding of rsigma.
Use the kernel for RMSNorm with MXFP8 output whenever both dimensions are multiples of 128, and let the PyTorch binding request swizzled scaling factors in that case. NVTE_NORM_FWD_MXFP8_USE_CUDNN=1, or nvte_enable_cudnn_norm_fwd_mxfp8(true), restores the previous paths.
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: