Skip to content

[Common] Fuse RMSNorm forward with MXFP8 quantization - #3657

Merged
phu0ngng merged 9 commits into
NVIDIA:mainfrom
sraman-rgb:sraman/rmsnorm-mxfp8-fused-kernel
Oct 9, 2026
Merged

phu0ngng merged 9 commits into
NVIDIA:mainfrom
sraman-rgb:sraman/rmsnorm-mxfp8-fused-kernel

Conversation

@sraman-rgb

Copy link
Copy Markdown
Collaborator

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

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Change A
  • Change B

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

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>
@greptile-apps

greptile-apps Bot commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 4/5

[High impact] Adds fused RMSNorm kernel with MXFP8 quantization output.

The PR is not ready to merge because column-only MXFP8 forward still throws before launching the kernel.

Findings

  1. P1 Column-only forward throws ▶

Summary

Adds fused RMSNorm forward with MXFP8 output for supported shapes whose dimensions are multiples of 128.

  • Computes rsigma first, then normalizes and quantizes each tile in FP32.
  • Writes compact or GEMM-swizzled scales directly.
  • Keeps requests for 2D quantization on the separate normalization and quantization path.
  • Adds checks for output values, scale layouts, and chunked launches.

Reviews (6) · Last reviewed commit: "Merge branch 'main' into sraman/rmsnorm-..." · Reviewed by Greptile

Comment thread transformer_engine/pytorch/csrc/extensions/normalization.cpp Outdated
Comment thread transformer_engine/pytorch/csrc/extensions/normalization.cpp Outdated
Comment thread transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp Outdated

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Have we profiled this kernel? I estimate its SoL perf to be 0.69x worse than the cuDNN kernel.

Comment thread docs/envvars.rst Outdated
Comment thread transformer_engine/common/include/transformer_engine/normalization.h Outdated
Comment thread transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu Outdated
Comment thread transformer_engine/pytorch/csrc/extensions/normalization.cpp Outdated
sraman-rgb and others added 2 commits October 8, 2026 16:43
…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>
Comment thread tests/cpp/operator/test_normalization_mxfp8.cu Outdated
sraman-rgb and others added 2 commits October 8, 2026 16:18
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>
Comment on lines +354 to +360
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;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Column-only forward throws

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.

@sraman-rgb

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 Pytorch

Comment thread transformer_engine/pytorch/csrc/extensions/normalization.cpp Outdated
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Comment thread transformer_engine/pytorch/csrc/extensions/normalization.cpp Outdated
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
timmoon10
timmoon10 previously approved these changes Oct 9, 2026

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM, pending CI

@vthumbe1503

Copy link
Copy Markdown
Collaborator

/te-ci L1 pytorch

@phu0ngng
phu0ngng merged commit d9beab1 into NVIDIA:main Oct 9, 2026
24 of 29 checks passed
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.

4 participants