From 5b3653af3ab042d8086ede4d958ce2c2c5799f96 Mon Sep 17 00:00:00 2001 From: Siddhartha Raman S Date: Thu, 8 Oct 2026 11:28:48 -0700 Subject: [PATCH 1/7] [Common] Fuse RMSNorm forward with MXFP8 quantization 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 Co-Authored-By: Claude Opus 5.5 --- docs/envvars.rst | 12 + .../cpp/operator/test_normalization_mxfp8.cu | 143 +++++- tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py | 70 +++ transformer_engine/common/CMakeLists.txt | 1 + .../transformer_engine/normalization.h | 10 + .../common/normalization/common.cpp | 12 + .../common/normalization/common.h | 7 + .../normalization/rmsnorm/rmsnorm_api.cpp | 21 +- .../rmsnorm/rmsnorm_fwd_mxfp8.cu | 408 ++++++++++++++++++ .../pytorch/csrc/extensions/normalization.cpp | 17 +- 10 files changed, 685 insertions(+), 16 deletions(-) create mode 100644 tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py create mode 100644 transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu diff --git a/docs/envvars.rst b/docs/envvars.rst index 46b70bbe46a..64d1b92cc73 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -418,6 +418,18 @@ LayerNorm/RMSNorm SM Margins BF16 input and normalization-output datatypes. When set to ``0``, or with an earlier cuDNN version, the virtual normalization output uses FP32. +.. envvar:: NVTE_NORM_FWD_MXFP8_USE_CUDNN + + :Type: ``int`` (0 or 1) + :Default: ``0`` + :Description: Use cuDNN for RMSNorm forward with MXFP8 output. By default, when both dimensions + of the input are multiples of 128, Transformer Engine's fused kernel normalizes, + quantizes the FP32 normalized values, and writes the scaling factors in the layout + of the output tensor, including the GEMM-swizzled layout. When set to ``1``, and + for other shapes, the forward uses cuDNN's fused kernel when + ``NVTE_NORM_FWD_USE_CUDNN=1`` and otherwise normalizes and quantizes in separate + kernels. + .. envvar:: NVTE_FWD_LAYERNORM_SM_MARGIN :Type: ``int`` diff --git a/tests/cpp/operator/test_normalization_mxfp8.cu b/tests/cpp/operator/test_normalization_mxfp8.cu index 10b33f8e2a1..212318fc614 100644 --- a/tests/cpp/operator/test_normalization_mxfp8.cu +++ b/tests/cpp/operator/test_normalization_mxfp8.cu @@ -102,7 +102,8 @@ void dequantize_2x(Tensor& input, Tensor& output, bool is_training) } template -void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, NormType norm_type, bool is_training, const bool zero_centered_gamma_in_weight_dtype) { +void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, NormType norm_type, bool is_training, const bool zero_centered_gamma_in_weight_dtype, + const bool use_cudnn_mxfp8) { cudaDeviceProp prop; cudaGetDeviceProperties(&prop, 0); @@ -134,6 +135,9 @@ void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, } else { nvte_enable_zero_centered_gamma_in_weight_dtype(false); } + // RMSNorm shapes that are multiples of 128 use Transformer Engine's fused MXFP8 kernel + // unless cuDNN is requested. + nvte_enable_cudnn_norm_fwd_mxfp8(use_cudnn_mxfp8); // Forward kernel float epsilon = 1e-5; @@ -163,6 +167,7 @@ void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, if (zero_centered_gamma_in_weight_dtype) { nvte_enable_zero_centered_gamma_in_weight_dtype(false); } + nvte_enable_cudnn_norm_fwd_mxfp8(false); Tensor dequantized_output("dequantized_output", std::vector{ N, H }, DType::kFloat32, true, true); @@ -197,7 +202,7 @@ void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, nullptr, // amax 1.f, // scale zero_centered_gamma, - true, // CuDNN is the only MXFP8 backend currently + true, // The MXFP8 backends add 1 to a zero-centered gamma as cuDNN does zero_centered_gamma_in_weight_dtype); cudaDeviceSynchronize(); @@ -235,6 +240,88 @@ std::vector> test_cases = { {2048, 12288}, }; +// Index of compact scaling factor (i, j) in the GEMM-swizzled layout: 128x4 tiles of i x j, +// num_tiles_j tiles along j. +size_t swizzled_scale_idx(size_t i, size_t j, size_t num_tiles_j) { + return ((i / 128) * num_tiles_j + j / 4) * 512 + (i % 32) * 16 + ((i % 128) / 32) * 4 + j % 4; +} + +// RMSNorm with MXFP8 output written with GEMM-swizzled scaling factors must give the same +// data as with compact scaling factors, and the compact scaling factors in swizzled order. +template +void performSwizzledScalesTest(const size_t N, const size_t H, const bool is_training) { + if (getDeviceComputeCapability() < blackwellComputeCapability) { + GTEST_SKIP(); + } + cudaDeviceProp prop; + cudaGetDeviceProperties(&prop, 0); + + const DType itype = TypeInfo::dtype; + const DType otype = TypeInfo::dtype; + Tensor input("input", std::vector{ N, H }, itype); + Tensor gamma("gamma", std::vector{ H }, itype); + Tensor rsigma("rsigma", std::vector{ N }, DType::kFloat32); + Tensor z("z", std::vector{ N, H }, otype, true, is_training, NVTE_MXFP8_1D_SCALING); + Tensor z_swizzled("z_swizzled", std::vector{ N, H }, otype, true, is_training, + NVTE_MXFP8_1D_SCALING); + z_swizzled.set_with_gemm_swizzled_scales(true); + fillUniform(&input); + fillUniform(&gamma); + + for (Tensor *out : {&z, &z_swizzled}) { + Tensor workspace; + nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out->data(), rsigma.data(), + workspace.data(), prop.multiProcessorCount, false, 0); + workspace = Tensor("workspace", workspace.rowwise_shape(), workspace.dtype()); + nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out->data(), rsigma.data(), + workspace.data(), prop.multiProcessorCount, false, 0); + } + cudaDeviceSynchronize(); + auto err = cudaGetLastError(); + ASSERT_EQ(err, cudaSuccess) << cudaGetErrorString(err); + z.to_cpu(); + z_swizzled.to_cpu(); + + const auto bytes = [](const OutputType *p) { return reinterpret_cast(p); }; + compareResults("rowwise_data", bytes(z_swizzled.rowwise_cpu_dptr()), + bytes(z.rowwise_cpu_dptr()), N * H); + const NVTEShape row_shape = z.rowwise_scale_inv_shape(); + const size_t row_dim0 = row_shape.data[0], row_dim1 = row_shape.data[1]; + std::vector ref(row_dim0 * row_dim1); + const uint8_t *compact = z.rowwise_cpu_scale_inv_ptr(); + for (size_t i = 0; i < row_dim0; ++i) { + for (size_t j = 0; j < row_dim1; ++j) { + ref[swizzled_scale_idx(i, j, row_dim1 / 4)] = compact[i * row_dim1 + j]; + } + } + compareResults("rowwise_scales", z_swizzled.rowwise_cpu_scale_inv_ptr(), ref.data(), + ref.size()); + + if (is_training) { + compareResults("colwise_data", bytes(z_swizzled.columnwise_cpu_dptr()), + bytes(z.columnwise_cpu_dptr()), N * H); + // Compact column-wise scaling factors are [rows / 32, cols]; the swizzled layout tiles + // columns by 128 and scale rows by 4. + const NVTEShape col_shape = z.columnwise_scale_inv_shape(); + const size_t col_dim0 = col_shape.data[0], col_dim1 = col_shape.data[1]; + ref.assign(col_dim0 * col_dim1, 0); + compact = z.columnwise_cpu_scale_inv_ptr(); + for (size_t i = 0; i < col_dim0; ++i) { + for (size_t j = 0; j < col_dim1; ++j) { + ref[swizzled_scale_idx(j, i, col_dim0 / 4)] = compact[i * col_dim1 + j]; + } + } + compareResults("colwise_scales", z_swizzled.columnwise_cpu_scale_inv_ptr(), + ref.data(), ref.size()); + } +} + +std::vector> swizzled_scales_test_cases = { + {128, 128}, + {768, 2304}, + {256, 7168}, +}; + std::vector norms = { NormType::LayerNorm, NormType::RMSNorm @@ -246,7 +333,7 @@ class MxNormTestSuite : public ::testing::TestWithParam< std::tuple, - bool, bool, bool>> {}; + bool, bool, bool, bool>> {}; TEST_P(MxNormTestSuite, TestMxNorm) { using namespace transformer_engine; @@ -259,10 +346,15 @@ TEST_P(MxNormTestSuite, TestMxNorm) { const bool zero_centered_gamma = std::get<4>(GetParam()); const bool is_training = std::get<5>(GetParam()); const bool zero_centered_gamma_in_weight_dtype = std::get<6>(GetParam()); + const bool use_cudnn_mxfp8 = std::get<7>(GetParam()); + if (norm_type == NormType::LayerNorm && use_cudnn_mxfp8) { + GTEST_SKIP() << "LayerNorm with MXFP8 output always uses cuDNN"; + } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, TRANSFORMER_ENGINE_TYPE_SWITCH_FP8_ONLY(output_type, OutputType, - performTest(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype); + performTest(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype, + use_cudnn_mxfp8); ); ); } @@ -277,7 +369,8 @@ INSTANTIATE_TEST_SUITE_P( ::testing::ValuesIn(test_cases), ::testing::Values(true, false), ::testing::Values(true, false), - ::testing::Values(true, false)), + ::testing::Values(true, false), + ::testing::Values(false, true)), [](const testing::TestParamInfo& info) { std::string name = normToString.at(std::get<0>(info.param)) + "_" + test::typeName(std::get<1>(info.param)) + "X" + @@ -286,6 +379,44 @@ INSTANTIATE_TEST_SUITE_P( std::to_string(std::get<3>(info.param).second) + "X" + std::to_string(std::get<4>(info.param)) + "out" + std::to_string(int(std::get<5>(info.param)) + 1) + "x" + - std::to_string(std::get<6>(info.param)); + std::to_string(std::get<6>(info.param)) + + (std::get<7>(info.param) ? "Xcudnn" : ""); return name; }); + +class MxNormSwizzledScalesTestSuite + : public ::testing::TestWithParam, bool>> {}; + +TEST_P(MxNormSwizzledScalesTestSuite, TestMxRMSNormSwizzledScales) { + using namespace transformer_engine; + using namespace test; + + const DType input_type = std::get<0>(GetParam()); + const DType output_type = std::get<1>(GetParam()); + const auto size = std::get<2>(GetParam()); + const bool is_training = std::get<3>(GetParam()); + + TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, + TRANSFORMER_ENGINE_TYPE_SWITCH_FP8_ONLY(output_type, OutputType, + performSwizzledScalesTest(size.first, size.second, is_training); + ); + ); +} + +INSTANTIATE_TEST_SUITE_P( + OperatorTest, + MxNormSwizzledScalesTestSuite, + ::testing::Combine( + ::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16), + ::testing::Values(DType::kFloat8E5M2, DType::kFloat8E4M3), + ::testing::ValuesIn(swizzled_scales_test_cases), + ::testing::Values(true, false)), + [](const testing::TestParamInfo& info) { + return test::typeName(std::get<0>(info.param)) + "X" + + test::typeName(std::get<1>(info.param)) + "X" + + std::to_string(std::get<2>(info.param).first) + "X" + + std::to_string(std::get<2>(info.param).second) + "X" + + std::to_string(int(std::get<3>(info.param)) + 1) + "x"; + }); diff --git a/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py b/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py new file mode 100644 index 00000000000..ed16fa0ac6e --- /dev/null +++ b/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py @@ -0,0 +1,70 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""RMSNorm forward with MXFP8 output from Transformer Engine's fused kernel.""" + +import os + +import pytest +import torch + +import transformer_engine.pytorch as te +import transformer_engine_torch as tex +from transformer_engine.pytorch import MXFP8Quantizer +from transformer_engine.pytorch.constants import TE_DType + +from mxfp8_utils import swizzle_mxfp8_scale + +recipe_available, reason_for_no_recipe = te.is_mxfp8_available(return_reason=True) + + +@pytest.mark.skipif(not recipe_available, reason=reason_for_no_recipe) +@pytest.mark.skipif( + os.getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN", "0") == "1", + reason="Tests Transformer Engine's fused kernel, not cuDNN", +) +@pytest.mark.parametrize("shape", [(128, 128), (256, 1024), (384, 7168)]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("fp8_dtype", [tex.DType.kFloat8E4M3, tex.DType.kFloat8E5M2]) +@pytest.mark.parametrize("zero_centered_gamma", [False, True]) +@pytest.mark.parametrize("optimize_for_gemm", [False, True]) +def test_rmsnorm_fwd_mxfp8(shape, dtype, fp8_dtype, zero_centered_gamma, optimize_for_gemm): + """The fused output is the MXFP8 quantization of the FP32 normalized values.""" + torch.manual_seed(1234) + rows, cols = shape + eps = 1e-5 + x = torch.randn(rows, cols, dtype=dtype, device="cuda") + weight = torch.randn(cols, dtype=dtype, device="cuda") + + quantizer = MXFP8Quantizer(fp8_dtype) + quantizer.optimize_for_gemm = optimize_for_gemm + out, _, rsigma = tex.rmsnorm_fwd( + x, weight, eps, None, quantizer, TE_DType[dtype], 0, zero_centered_gamma + ) + assert out._with_gemm_swizzled_scales == optimize_for_gemm + + torch.testing.assert_close( + rsigma, torch.rsqrt(x.float().square().mean(dim=1) + eps), atol=0, rtol=1e-5 + ) + + # Reference: the MXFP8 quantizer on the FP32 normalized values, with the kernel's rsigma. + gamma = weight.float() + 1 if zero_centered_gamma else weight.float() + y = (x.float() * rsigma.unsqueeze(1)) * gamma + ref_quantizer = MXFP8Quantizer(fp8_dtype) + ref_quantizer.optimize_for_gemm = False + ref = ref_quantizer(y) + + for attr in ("_rowwise_data", "_columnwise_data"): + torch.testing.assert_close( + getattr(out, attr).view(torch.uint8), + getattr(ref, attr).view(torch.uint8), + atol=0, + rtol=0, + msg=attr, + ) + for attr, columnwise in (("_rowwise_scale_inv", False), ("_columnwise_scale_inv", True)): + expected = getattr(ref, attr) + if optimize_for_gemm: + expected = swizzle_mxfp8_scale(rows, cols, expected, columnwise=columnwise) + torch.testing.assert_close(getattr(out, attr), expected, atol=0, rtol=0, msg=attr) diff --git a/transformer_engine/common/CMakeLists.txt b/transformer_engine/common/CMakeLists.txt index 843ec412c3d..ded6b6cb801 100644 --- a/transformer_engine/common/CMakeLists.txt +++ b/transformer_engine/common/CMakeLists.txt @@ -373,6 +373,7 @@ list(APPEND transformer_engine_cuda_arch_specific_sources hadamard_transform/group_row_cast_col_hadamard_transform_cast_fusion.cu hadamard_transform/graph_safe_group_row_cast_col_hadamard_transform_cast_fusion.cu multi_tensor/compute_scale.cu + normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu recipe/mxfp8_scaling.cu recipe/nvfp4.cu transpose/quantize_transpose_square_blockwise.cu diff --git a/transformer_engine/common/include/transformer_engine/normalization.h b/transformer_engine/common/include/transformer_engine/normalization.h index b72b8abe32a..5c963c1c6c2 100644 --- a/transformer_engine/common/include/transformer_engine/normalization.h +++ b/transformer_engine/common/include/transformer_engine/normalization.h @@ -181,6 +181,16 @@ void nvte_enable_cudnn_norm_fwd(bool enable); */ void nvte_enable_cudnn_norm_bwd(bool enable); +/*! \brief Set whether RMSNorm forward with MXFP8 output uses cuDNN. + * + * By default, RMSNorm forward with MXFP8 output uses Transformer Engine's fused kernel + * whenever both dimensions are multiples of 128; it can also write GEMM-swizzled scaling + * factors. Otherwise, and when enabled here, the cuDNN backend is used. + * + * \param[in] enable Whether to use cuDNN for RMSNorm forward with MXFP8 output. + */ +void nvte_enable_cudnn_norm_fwd_mxfp8(bool enable); + /*! \brief Control whether norm computes `gamma += 1.0` for zero-centered gamma * in weight dtype. If set to false, it will compute in compute dtype. * diff --git a/transformer_engine/common/normalization/common.cpp b/transformer_engine/common/normalization/common.cpp index a7095aa0570..84c1fe93482 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -578,6 +578,13 @@ bool use_cudnn_mxfp8_norm_output_in_input_dtype() { return flag; } +bool& _cudnn_norm_fwd_mxfp8_flag() { + static bool flag = transformer_engine::getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN"); + return flag; +} + +bool use_cudnn_norm_fwd_mxfp8() { return _cudnn_norm_fwd_mxfp8_flag(); } + } // namespace normalization } // namespace transformer_engine @@ -586,6 +593,11 @@ void nvte_enable_cudnn_norm_fwd(bool enable) { transformer_engine::normalization::_cudnn_norm_fwd_flag() = enable; } +void nvte_enable_cudnn_norm_fwd_mxfp8(bool enable) { + NVTE_API_CALL(nvte_enable_cudnn_norm_fwd_mxfp8); + transformer_engine::normalization::_cudnn_norm_fwd_mxfp8_flag() = enable; +} + void nvte_enable_cudnn_norm_bwd(bool enable) { NVTE_API_CALL(nvte_enable_cudnn_norm_bwd); transformer_engine::normalization::_cudnn_norm_bwd_flag() = enable; diff --git a/transformer_engine/common/normalization/common.h b/transformer_engine/common/normalization/common.h index dd2f3d14591..19eed8cfea1 100644 --- a/transformer_engine/common/normalization/common.h +++ b/transformer_engine/common/normalization/common.h @@ -309,6 +309,13 @@ bool use_cudnn_norm_bwd(); bool& use_zero_centered_gamma_in_weight_dtype(); bool use_cudnn_mxfp8_norm_output_in_input_dtype(); +bool use_cudnn_norm_fwd_mxfp8(); + +// RMSNorm forward with MXFP8 output by Transformer Engine's fused kernel +// (rmsnorm/rmsnorm_fwd_mxfp8.cu), which also writes GEMM-swizzled scaling factors. +bool use_te_rmsnorm_fwd_mxfp8(const Tensor& x, const Tensor& gamma, const Tensor& z); +void rmsnorm_fwd_mxfp8(const Tensor& x, const Tensor& gamma, float epsilon, Tensor* z, + Tensor* rsigma, bool zero_centered_gamma, cudaStream_t stream); } // namespace normalization } // namespace transformer_engine diff --git a/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp b/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp index f56e2ef1fef..e074f7e72c4 100644 --- a/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp +++ b/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp @@ -27,11 +27,6 @@ void rmsnorm_fwd(const Tensor &x, const Tensor &gamma, const float epsilon, Tens !is_mxfp8_scaling(z->scaling_mode)) { NVTE_ERROR("Not implemented scaling mode: " + to_string(z->scaling_mode) + "."); } - if (is_mxfp8_scaling(z->scaling_mode)) { - NVTE_CHECK(!z->with_gemm_swizzled_scales, - "MXFP8 output must have scales in compact format, not swizzled for GEMM."); - } - NVTE_CHECK(!x.data.shape.empty(), "x must have at least one dimension."); const auto [rows, cols] = x.flat_2d_dims(); @@ -53,6 +48,22 @@ void rmsnorm_fwd(const Tensor &x, const Tensor &gamma, const float epsilon, Tens CheckOutputTensor(*rsigma, "rsigma"); } + if (use_te_rmsnorm_fwd_mxfp8(x, gamma, *z)) { + if (workspace->data.numel() == 0) { + // The fused kernel needs no workspace, but a zero size denotes a workspace query. + workspace->data.shape = {1}; + workspace->data.dtype = DType::kByte; + return; + } + rmsnorm_fwd_mxfp8(x, gamma, epsilon, z, rsigma, zero_centered_gamma, stream); + return; + } + if (is_mxfp8_scaling(z->scaling_mode)) { + NVTE_CHECK(!z->with_gemm_swizzled_scales, + "MXFP8 output must have scales in compact format, not swizzled for GEMM, " + "unless Transformer Engine's fused RMSNorm + MXFP8 kernel is used."); + } + NVTE_Norm_Backend norm_backend; bool is_aligned = true; bool cudnn_backend = use_cudnn_norm_fwd() || is_mxfp8_scaling(z->scaling_mode); diff --git a/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu b/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu new file mode 100644 index 00000000000..103bb3c043a --- /dev/null +++ b/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu @@ -0,0 +1,408 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +/*! \file rmsnorm_fwd_mxfp8.cu + * \brief RMSNorm forward fused with MXFP8 quantization. + * + * Normalizes in FP32 and quantizes the FP32 result, row-wise and/or column-wise, with + * the scaling factors written in compact or GEMM-swizzled layout. The output is that of + * the MXFP8 quantizer applied to the FP32 normalized values, which is also what cuDNN's + * fused kernel produces; the two can differ only through the rounding of rsigma. + * + * Two kernels: a pre-pass computes rsigma per row (one warp per row), then one CTA per + * 32x128 tile normalizes the tile, quantizes it and writes its scaling factors. A tile + * only sees 128 columns of a row, so it cannot reduce the row itself; doing so with whole + * rows per CTA would limit the grid to rows / 32 CTAs. + */ + +#include + +#include +#include + +#include "../../cast/mxfp8/swizzle.cuh" +#include "../../common.h" +#include "../../util/ptx_arch_spec.cuh" +#include "../../utils.cuh" +#include "../common.h" + +namespace transformer_engine { +namespace normalization { +namespace { + +constexpr int kScaleBlock = 32; // MXFP8 scaling block size. +constexpr int kTileRows = 32; // One column-wise scaling block. +constexpr int kTileCols = 128; // Four row-wise scaling blocks. +constexpr int kThreads = 128; +constexpr int kWarps = kThreads / THREADS_PER_WARP; +constexpr int kRowsPerWarp = kTileRows / kWarps; +constexpr int kColsPerLane = kTileCols / THREADS_PER_WARP; +constexpr int kLanesPerBlock = kScaleBlock / kColsPerLane; +constexpr int kRsigmaRowsPerCTA = 4; + +static_assert(kRowsPerWarp == 8); +static_assert(kColsPerLane == 4); +static_assert(kLanesPerBlock == 8); + +template +__device__ __forceinline__ void load_elements(T (&dst)[N], const T *src) { + if constexpr (sizeof(T) * N == 16) { + *reinterpret_cast(dst) = *reinterpret_cast(src); + } else if constexpr (sizeof(T) * N == 8) { + *reinterpret_cast(dst) = *reinterpret_cast(src); + } else { +#pragma unroll + for (int i = 0; i < N; ++i) { + dst[i] = src[i]; + } + } +} + +// rsigma = 1 / sqrt(mean(x^2) + epsilon), one warp per row. +template +__global__ void __launch_bounds__(kRsigmaRowsPerCTA *THREADS_PER_WARP) + rmsnorm_mxfp8_rsigma_kernel(const IType *const x, float *const rsigma, const int rows, + const int cols, const float epsilon) { + const int row = blockIdx.x * kRsigmaRowsPerCTA + threadIdx.x / THREADS_PER_WARP; + const int lane = threadIdx.x % THREADS_PER_WARP; + if (row >= rows) { + return; + } + const IType *const x_row = x + static_cast(row) * cols; + float sum = 0.0f; + if constexpr (kAligned) { + constexpr int kVec = 16 / sizeof(IType); + for (int c = lane * kVec; c < cols; c += THREADS_PER_WARP * kVec) { + IType values[kVec]; + load_elements(values, x_row + c); +#pragma unroll + for (int i = 0; i < kVec; ++i) { + const float value = static_cast(values[i]); + sum += value * value; + } + } + } else { + for (int c = lane; c < cols; c += THREADS_PER_WARP) { + const float value = static_cast(x_row[c]); + sum += value * value; + } + } +#pragma unroll + for (int offset = THREADS_PER_WARP / 2; offset > 0; offset /= 2) { + sum += __shfl_xor_sync(0xffffffff, sum, offset); + } + if (lane == 0) { + rsigma[row] = 1.0f / sqrtf(sum / static_cast(cols) + epsilon); + } +} + +// Four FP32 values times per-value multipliers, converted to four FP8 values. +template +__device__ __forceinline__ uint32_t quantize_4x(const float (&values)[kColsPerLane], + const float (&multipliers)[kColsPerLane]) { + using OTypex2 = + std::conditional_t, ptx::fp8e4m3x2, ptx::fp8e5m2x2>; + OTypex2 out[2]; + ptx::mul_cvt_2x(out[0], ptx::floatx2{values[0], values[1]}, + ptx::floatx2{multipliers[0], multipliers[1]}); + ptx::mul_cvt_2x(out[1], ptx::floatx2{values[2], values[3]}, + ptx::floatx2{multipliers[2], multipliers[3]}); + return static_cast(reinterpret_cast(out[0])) | + (static_cast(reinterpret_cast(out[1])) << 16); +} + +struct QuantizeArgs { + void *rowwise_data; + e8m0_t *rowwise_scale_inv; + void *colwise_data; + e8m0_t *colwise_scale_inv; + // Compact layouts: row strides of the scaling-factor tensors. Swizzled layouts: number of + // 4-wide scaling-factor tiles along the blocked dimension. + size_t rowwise_scale_stride; + size_t colwise_scale_stride; +}; + +// One CTA per 32x128 tile: warp w normalizes rows 8w..8w+7, lane l columns 4l..4l+3. +template +__global__ void __launch_bounds__(kThreads) + rmsnorm_mxfp8_quantize_kernel(const IType *const x, const WType *const gamma, + const float *const rsigma, const QuantizeArgs args, + const int cols, const bool zero_centered_gamma, + const bool gamma_in_weight_dtype) { +#if (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + using dispatch::mxfp8::swizzle::gemm_swizzled_scale_idx; + constexpr float kMaxNormRcp = Quantized_Limits::max_norm_rcp; + __shared__ float colwise_amax[kWarps][kTileCols]; + + const int warp = threadIdx.x / THREADS_PER_WARP; + const int lane = threadIdx.x % THREADS_PER_WARP; + const int row_base = blockIdx.x * kTileRows; + const int row0 = row_base + warp * kRowsPerWarp; + const int col0 = blockIdx.y * kTileCols + lane * kColsPerLane; + + // Gamma as applied by the other normalization backends: with a zero-centered gamma, 1 is + // added in FP32, or in the weight type if requested. + float g[kColsPerLane]; + { + WType w[kColsPerLane]; + if constexpr (kAligned) { + load_elements(w, gamma + col0); + } else { +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + w[i] = gamma[col0 + i]; + } + } +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + float value = static_cast(w[i]); + if (zero_centered_gamma) { + value = gamma_in_weight_dtype ? static_cast(static_cast(value + 1.0f)) + : value + 1.0f; + } + g[i] = value; + } + } + + float y[kRowsPerWarp][kColsPerLane]; +#pragma unroll + for (int r = 0; r < kRowsPerWarp; ++r) { + const int row = row0 + r; + const float rs = rsigma[row]; + IType values[kColsPerLane]; + const IType *const src = x + static_cast(row) * cols + col0; + if constexpr (kAligned) { + load_elements(values, src); + } else { +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + values[i] = src[i]; + } + } +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + y[r][i] = (static_cast(values[i]) * rs) * g[i]; + } + } + + if constexpr (kRowwise) { + OType *const data = reinterpret_cast(args.rowwise_data); +#pragma unroll + for (int r = 0; r < kRowsPerWarp; ++r) { + const int row = row0 + r; + // Each 32-column scaling block spans 8 consecutive lanes. + float amax = 0.0f; +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + amax = fmaxf(amax, fabsf(y[r][i])); + } +#pragma unroll + for (int offset = 1; offset < kLanesPerBlock; offset *= 2) { + amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, offset)); + } + const e8m0_t biased_exponent = ptx::float_to_e8m0(amax * kMaxNormRcp); + const float multiplier = ptx::exp2f_rcp(biased_exponent); + const float multipliers[kColsPerLane] = {multiplier, multiplier, multiplier, multiplier}; + const uint32_t quantized = quantize_4x(y[r], multipliers); + OType *const dst = data + static_cast(row) * cols + col0; + if constexpr (kAligned) { + *reinterpret_cast(dst) = quantized; + } else { +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + reinterpret_cast(dst)[i] = static_cast(quantized >> (8 * i)); + } + } + if (lane % kLanesPerBlock == 0) { + const size_t scale_col = col0 / kScaleBlock; + const size_t idx = kSwizzled + ? gemm_swizzled_scale_idx(row, scale_col, args.rowwise_scale_stride) + : static_cast(row) * args.rowwise_scale_stride + scale_col; + args.rowwise_scale_inv[idx] = biased_exponent; + } + } + } + + if constexpr (kColwise) { + // Each column's scaling block is the tile's 32 rows: reduce over this warp's 8 rows, then + // over the 4 warps through shared memory. +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + float amax = 0.0f; +#pragma unroll + for (int r = 0; r < kRowsPerWarp; ++r) { + amax = fmaxf(amax, fabsf(y[r][i])); + } + colwise_amax[warp][lane * kColsPerLane + i] = amax; + } + __syncthreads(); + e8m0_t biased_exponents[kColsPerLane]; + float multipliers[kColsPerLane]; +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + float amax = colwise_amax[0][lane * kColsPerLane + i]; +#pragma unroll + for (int w = 1; w < kWarps; ++w) { + amax = fmaxf(amax, colwise_amax[w][lane * kColsPerLane + i]); + } + biased_exponents[i] = ptx::float_to_e8m0(amax * kMaxNormRcp); + multipliers[i] = ptx::exp2f_rcp(biased_exponents[i]); + } + OType *const data = reinterpret_cast(args.colwise_data); +#pragma unroll + for (int r = 0; r < kRowsPerWarp; ++r) { + const uint32_t quantized = quantize_4x(y[r], multipliers); + OType *const dst = data + static_cast(row0 + r) * cols + col0; + if constexpr (kAligned) { + *reinterpret_cast(dst) = quantized; + } else { +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + reinterpret_cast(dst)[i] = static_cast(quantized >> (8 * i)); + } + } + } + if (warp == 0) { + const size_t scale_row = row_base / kScaleBlock; +#pragma unroll + for (int i = 0; i < kColsPerLane; ++i) { + const size_t col = col0 + i; + const size_t idx = kSwizzled + ? gemm_swizzled_scale_idx(col, scale_row, args.colwise_scale_stride) + : scale_row * args.colwise_scale_stride + col; + args.colwise_scale_inv[idx] = biased_exponents[i]; + } + } + } +#else + NVTE_DEVICE_THREAD0_ERROR("RMSNorm forward with MXFP8 output requires SM 10.0+."); +#endif // (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) +} + +bool is_aligned_to(const void *ptr, size_t alignment) { + return reinterpret_cast(ptr) % alignment == 0; +} + +template +void launch_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float epsilon, Tensor *z, + Tensor *rsigma, const bool zero_centered_gamma, + const bool gamma_in_weight_dtype, cudaStream_t stream) { + const auto [rows_size_t, cols_size_t] = x.flat_2d_dims(); + const int rows = static_cast(rows_size_t); + const int cols = static_cast(cols_size_t); + const bool rowwise = z->has_data(); + const bool colwise = z->has_columnwise_data(); + const bool swizzled = z->with_gemm_swizzled_scales; + + QuantizeArgs args{}; + if (rowwise) { + args.rowwise_data = z->data.dptr; + args.rowwise_scale_inv = reinterpret_cast(z->scale_inv.dptr); + const size_t scale_cols = z->scale_inv.shape.back(); + args.rowwise_scale_stride = swizzled ? scale_cols / 4 : scale_cols; + } + if (colwise) { + args.colwise_data = z->columnwise_data.dptr; + args.colwise_scale_inv = reinterpret_cast(z->columnwise_scale_inv.dptr); + args.colwise_scale_stride = + swizzled ? z->columnwise_scale_inv.shape.front() / 4 : z->columnwise_scale_inv.shape.back(); + } + + const IType *const x_ptr = reinterpret_cast(x.data.dptr); + const WType *const gamma_ptr = reinterpret_cast(gamma.data.dptr); + float *const rsigma_ptr = reinterpret_cast(rsigma->data.dptr); + + // Vector accesses: 16 bytes of a row of x in the pre-pass, 4 elements of x and gamma and 4 + // bytes of each output in the main kernel. Every row starts aligned since cols % 128 == 0. + const bool aligned = is_aligned_to(x_ptr, 16) && is_aligned_to(gamma_ptr, 4 * sizeof(WType)) && + (!rowwise || is_aligned_to(args.rowwise_data, 4)) && + (!colwise || is_aligned_to(args.colwise_data, 4)); + + TRANSFORMER_ENGINE_SWITCH_CONDITION( + aligned, kAligned, + rmsnorm_mxfp8_rsigma_kernel + <<>>( + x_ptr, rsigma_ptr, rows, cols, epsilon);); // NOLINT(*) + NVTE_CHECK_CUDA(cudaGetLastError()); + + const dim3 grid(rows / kTileRows, cols / kTileCols); + TRANSFORMER_ENGINE_SWITCH_CONDITION( + aligned, kAligned, + TRANSFORMER_ENGINE_SWITCH_CONDITION( + swizzled, kSwizzled, + if (rowwise && colwise) { + rmsnorm_mxfp8_quantize_kernel + <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, + zero_centered_gamma, gamma_in_weight_dtype); + } else if (rowwise) { + rmsnorm_mxfp8_quantize_kernel + <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, + zero_centered_gamma, gamma_in_weight_dtype); + } else { + rmsnorm_mxfp8_quantize_kernel + <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, + zero_centered_gamma, gamma_in_weight_dtype); + });); // NOLINT(*) + NVTE_CHECK_CUDA(cudaGetLastError()); +} + +} // namespace + +bool use_te_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const Tensor &z) { + if (!is_mxfp8_scaling(z.scaling_mode) || use_cudnn_norm_fwd_mxfp8() || + !is_supported_by_CC_100()) { + return false; + } + const bool rowwise = z.has_data(); + const bool colwise = z.has_columnwise_data(); + if (!rowwise && !colwise) { + return false; + } + if ((rowwise && z.scale_inv.shape.size() != 2) || + (colwise && z.columnwise_scale_inv.shape.size() != 2)) { + return false; + } + const auto [rows, cols] = x.flat_2d_dims(); + // Full 128x128 tiles, so the scaling factors carry no padding in either layout. + if (rows == 0 || cols == 0 || rows % 128 != 0 || cols % 128 != 0 || + rows > static_cast(std::numeric_limits::max()) || cols / kTileCols > 65535) { + return false; + } + const auto is_supported_input = [](DType t) { + return t == DType::kFloat32 || t == DType::kBFloat16 || t == DType::kFloat16; + }; + if (!is_supported_input(x.data.dtype) || !is_supported_input(gamma.data.dtype)) { + return false; + } + const auto is_supported_output = [](DType t) { + return t == DType::kFloat8E4M3 || t == DType::kFloat8E5M2; + }; + if ((rowwise && !is_supported_output(z.data.dtype)) || + (colwise && !is_supported_output(z.columnwise_data.dtype)) || + (rowwise && colwise && z.data.dtype != z.columnwise_data.dtype)) { + return false; + } + return true; +} + +void rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float epsilon, Tensor *z, + Tensor *rsigma, const bool zero_centered_gamma, cudaStream_t stream) { + const DType otype = z->has_data() ? z->data.dtype : z->columnwise_data.dtype; + const bool gamma_in_weight_dtype = use_zero_centered_gamma_in_weight_dtype(); + TRANSFORMER_ENGINE_TYPE_SWITCH_FLOAT( + x.data.dtype, IType, + TRANSFORMER_ENGINE_TYPE_SWITCH_FLOAT( + gamma.data.dtype, WType, + TRANSFORMER_ENGINE_TYPE_SWITCH_FP8ONLY( + otype, OType, + launch_rmsnorm_fwd_mxfp8( + x, gamma, epsilon, z, rsigma, zero_centered_gamma, gamma_in_weight_dtype, + stream);););); // NOLINT(*) +} + +} // namespace normalization +} // namespace transformer_engine diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index 43f1d32b8a3..ebef28d6341 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -343,13 +343,20 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w FUSED_NORM_AMAX_NVFP4 }; Impl impl = Impl::UNFUSED; + // Whether the fused kernel can write GEMM-swizzled scaling factors + bool fused_kernel_supports_swizzled_scales = false; if (quantizer.is_none() || IsFloat8Quantizers(quantizer.ptr())) { impl = Impl::FULLY_FUSED; } else if (IsMXFP8Quantizers(quantizer.ptr())) { - if (transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN") && outer_size % 128 == 0 && - inner_size % 128 == 0) { - // cuDNN MXFP8 kernel requires full 128x128 tiles - impl = Impl::FULLY_FUSED; + if (outer_size % 128 == 0 && inner_size % 128 == 0) { + if (!transformer_engine::getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN")) { + // Transformer Engine's fused RMSNorm + MXFP8 kernel + impl = Impl::FULLY_FUSED; + fused_kernel_supports_swizzled_scales = true; + } else if (transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN")) { + // cuDNN MXFP8 kernel requires full 128x128 tiles + impl = Impl::FULLY_FUSED; + } } } else if (detail::IsFloat8CurrentScalingQuantizers(quantizer.ptr()) && !transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN")) { @@ -375,7 +382,7 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // Output tensor TensorWrapper out_nvte; if (out.is_none()) { - if (impl == Impl::FULLY_FUSED) { + if (impl == Impl::FULLY_FUSED && !fused_kernel_supports_swizzled_scales) { // FP8 has no special logic to optimize for GEMM, MXFP8 cuDNN // kernel does not support GEMM swizzled scales quantizer_cpp->optimize_for_gemm = false; From 1d07cc3427337c473054de29579b812e8dd24c18 Mon Sep 17 00:00:00 2001 From: Siddhartha Raman Sundara Raman Date: Thu, 8 Oct 2026 16:43:37 -0500 Subject: [PATCH 2/7] Update transformer_engine/common/include/transformer_engine/normalization.h Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: Siddhartha Raman Sundara Raman --- .../common/include/transformer_engine/normalization.h | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/transformer_engine/common/include/transformer_engine/normalization.h b/transformer_engine/common/include/transformer_engine/normalization.h index 5c963c1c6c2..b72b8abe32a 100644 --- a/transformer_engine/common/include/transformer_engine/normalization.h +++ b/transformer_engine/common/include/transformer_engine/normalization.h @@ -181,16 +181,6 @@ void nvte_enable_cudnn_norm_fwd(bool enable); */ void nvte_enable_cudnn_norm_bwd(bool enable); -/*! \brief Set whether RMSNorm forward with MXFP8 output uses cuDNN. - * - * By default, RMSNorm forward with MXFP8 output uses Transformer Engine's fused kernel - * whenever both dimensions are multiples of 128; it can also write GEMM-swizzled scaling - * factors. Otherwise, and when enabled here, the cuDNN backend is used. - * - * \param[in] enable Whether to use cuDNN for RMSNorm forward with MXFP8 output. - */ -void nvte_enable_cudnn_norm_fwd_mxfp8(bool enable); - /*! \brief Control whether norm computes `gamma += 1.0` for zero-centered gamma * in weight dtype. If set to false, it will compute in compute dtype. * From e1991d0695a9ab16061fcdbe8a992f77e58de8a1 Mon Sep 17 00:00:00 2001 From: Siddhartha Raman Sundara Raman Date: Thu, 8 Oct 2026 16:44:29 -0500 Subject: [PATCH 3/7] Update docs/envvars.rst Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: Siddhartha Raman Sundara Raman --- docs/envvars.rst | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 64d1b92cc73..46b70bbe46a 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -418,18 +418,6 @@ LayerNorm/RMSNorm SM Margins BF16 input and normalization-output datatypes. When set to ``0``, or with an earlier cuDNN version, the virtual normalization output uses FP32. -.. envvar:: NVTE_NORM_FWD_MXFP8_USE_CUDNN - - :Type: ``int`` (0 or 1) - :Default: ``0`` - :Description: Use cuDNN for RMSNorm forward with MXFP8 output. By default, when both dimensions - of the input are multiples of 128, Transformer Engine's fused kernel normalizes, - quantizes the FP32 normalized values, and writes the scaling factors in the layout - of the output tensor, including the GEMM-swizzled layout. When set to ``1``, and - for other shapes, the forward uses cuDNN's fused kernel when - ``NVTE_NORM_FWD_USE_CUDNN=1`` and otherwise normalizes and quantizes in separate - kernels. - .. envvar:: NVTE_FWD_LAYERNORM_SM_MARGIN :Type: ``int`` From 59b913ba7bf4fc26ab2d7eb8c3660f18096f6aa4 Mon Sep 17 00:00:00 2001 From: Siddhartha Raman S Date: Thu, 8 Oct 2026 16:18:29 -0700 Subject: [PATCH 4/7] [Common] Select the RMSNorm + MXFP8 kernel with NVTE_NORM_FWD_USE_CUDNN 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 Co-Authored-By: Claude Opus 5.5 --- .../cpp/operator/test_normalization_mxfp8.cu | 31 +++--- tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py | 64 ++++++++++++- .../common/normalization/common.cpp | 12 --- .../common/normalization/common.h | 6 +- .../normalization/rmsnorm/rmsnorm_api.cpp | 28 +++--- .../rmsnorm/rmsnorm_fwd_mxfp8.cu | 94 ++++++++++++++----- .../pytorch/csrc/extensions/normalization.cpp | 21 +++-- 7 files changed, 184 insertions(+), 72 deletions(-) diff --git a/tests/cpp/operator/test_normalization_mxfp8.cu b/tests/cpp/operator/test_normalization_mxfp8.cu index 212318fc614..84c244ff1ff 100644 --- a/tests/cpp/operator/test_normalization_mxfp8.cu +++ b/tests/cpp/operator/test_normalization_mxfp8.cu @@ -103,7 +103,7 @@ void dequantize_2x(Tensor& input, Tensor& output, bool is_training) template void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, NormType norm_type, bool is_training, const bool zero_centered_gamma_in_weight_dtype, - const bool use_cudnn_mxfp8) { + const bool use_cudnn) { cudaDeviceProp prop; cudaGetDeviceProperties(&prop, 0); @@ -137,7 +137,7 @@ void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, } // RMSNorm shapes that are multiples of 128 use Transformer Engine's fused MXFP8 kernel // unless cuDNN is requested. - nvte_enable_cudnn_norm_fwd_mxfp8(use_cudnn_mxfp8); + nvte_enable_cudnn_norm_fwd(use_cudnn); // Forward kernel float epsilon = 1e-5; @@ -167,7 +167,7 @@ void performTest(const size_t N, const size_t H, const bool zero_centered_gamma, if (zero_centered_gamma_in_weight_dtype) { nvte_enable_zero_centered_gamma_in_weight_dtype(false); } - nvte_enable_cudnn_norm_fwd_mxfp8(false); + nvte_enable_cudnn_norm_fwd(false); Tensor dequantized_output("dequantized_output", std::vector{ N, H }, DType::kFloat32, true, true); @@ -248,6 +248,8 @@ size_t swizzled_scale_idx(size_t i, size_t j, size_t num_tiles_j) { // RMSNorm with MXFP8 output written with GEMM-swizzled scaling factors must give the same // data as with compact scaling factors, and the compact scaling factors in swizzled order. +// The swizzled output is computed on a single SM, as with an SM margin, which launches the +// kernels over chunks of rows. template void performSwizzledScalesTest(const size_t N, const size_t H, const bool is_training) { if (getDeviceComputeCapability() < blackwellComputeCapability) { @@ -261,6 +263,7 @@ void performSwizzledScalesTest(const size_t N, const size_t H, const bool is_tra Tensor input("input", std::vector{ N, H }, itype); Tensor gamma("gamma", std::vector{ H }, itype); Tensor rsigma("rsigma", std::vector{ N }, DType::kFloat32); + Tensor rsigma_swizzled("rsigma_swizzled", std::vector{ N }, DType::kFloat32); Tensor z("z", std::vector{ N, H }, otype, true, is_training, NVTE_MXFP8_1D_SCALING); Tensor z_swizzled("z_swizzled", std::vector{ N, H }, otype, true, is_training, NVTE_MXFP8_1D_SCALING); @@ -268,19 +271,23 @@ void performSwizzledScalesTest(const size_t N, const size_t H, const bool is_tra fillUniform(&input); fillUniform(&gamma); - for (Tensor *out : {&z, &z_swizzled}) { + const auto run = [&](Tensor &out, Tensor &rs, const int sm_count) { Tensor workspace; - nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out->data(), rsigma.data(), - workspace.data(), prop.multiProcessorCount, false, 0); + nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out.data(), rs.data(), + workspace.data(), sm_count, false, 0); workspace = Tensor("workspace", workspace.rowwise_shape(), workspace.dtype()); - nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out->data(), rsigma.data(), - workspace.data(), prop.multiProcessorCount, false, 0); - } + nvte_rmsnorm_fwd(input.data(), gamma.data(), 1e-5f, out.data(), rs.data(), + workspace.data(), sm_count, false, 0); + }; + run(z, rsigma, prop.multiProcessorCount); + run(z_swizzled, rsigma_swizzled, 1); cudaDeviceSynchronize(); auto err = cudaGetLastError(); ASSERT_EQ(err, cudaSuccess) << cudaGetErrorString(err); z.to_cpu(); z_swizzled.to_cpu(); + rsigma.to_cpu(); + compareResults("rsigma", rsigma_swizzled, rsigma.rowwise_cpu_dptr(), true, 0, 0); const auto bytes = [](const OutputType *p) { return reinterpret_cast(p); }; compareResults("rowwise_data", bytes(z_swizzled.rowwise_cpu_dptr()), @@ -346,15 +353,15 @@ TEST_P(MxNormTestSuite, TestMxNorm) { const bool zero_centered_gamma = std::get<4>(GetParam()); const bool is_training = std::get<5>(GetParam()); const bool zero_centered_gamma_in_weight_dtype = std::get<6>(GetParam()); - const bool use_cudnn_mxfp8 = std::get<7>(GetParam()); - if (norm_type == NormType::LayerNorm && use_cudnn_mxfp8) { + const bool use_cudnn = std::get<7>(GetParam()); + if (norm_type == NormType::LayerNorm && use_cudnn) { GTEST_SKIP() << "LayerNorm with MXFP8 output always uses cuDNN"; } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, TRANSFORMER_ENGINE_TYPE_SWITCH_FP8_ONLY(output_type, OutputType, performTest(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype, - use_cudnn_mxfp8); + use_cudnn); ); ); } diff --git a/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py b/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py index ed16fa0ac6e..5200ec1074b 100644 --- a/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py +++ b/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py @@ -21,7 +21,7 @@ @pytest.mark.skipif(not recipe_available, reason=reason_for_no_recipe) @pytest.mark.skipif( - os.getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN", "0") == "1", + os.getenv("NVTE_NORM_FWD_USE_CUDNN", "0") == "1", reason="Tests Transformer Engine's fused kernel, not cuDNN", ) @pytest.mark.parametrize("shape", [(128, 128), (256, 1024), (384, 7168)]) @@ -68,3 +68,65 @@ def test_rmsnorm_fwd_mxfp8(shape, dtype, fp8_dtype, zero_centered_gamma, optimiz if optimize_for_gemm: expected = swizzle_mxfp8_scale(rows, cols, expected, columnwise=columnwise) torch.testing.assert_close(getattr(out, attr), expected, atol=0, rtol=0, msg=attr) + + +@pytest.mark.skipif(not recipe_available, reason=reason_for_no_recipe) +@pytest.mark.skipif( + os.getenv("NVTE_NORM_FWD_USE_CUDNN", "0") == "1", + reason="Tests Transformer Engine's fused kernel, not cuDNN", +) +@pytest.mark.parametrize("optimize_for_gemm", [False, True]) +def test_rmsnorm_fwd_mxfp8_sm_margin(optimize_for_gemm): + """With all SMs but one reserved, the kernels run over chunks of rows with the same output.""" + torch.manual_seed(1234) + x = torch.randn(1024, 7168, dtype=torch.bfloat16, device="cuda") + weight = torch.randn(7168, dtype=torch.bfloat16, device="cuda") + sm_count = torch.cuda.get_device_properties(x.device).multi_processor_count + + outputs = [] + for sm_margin in (0, sm_count - 1): + quantizer = MXFP8Quantizer(tex.DType.kFloat8E4M3) + quantizer.optimize_for_gemm = optimize_for_gemm + outputs.append( + tex.rmsnorm_fwd(x, weight, 1e-5, None, quantizer, TE_DType[x.dtype], sm_margin, False) + ) + (out, _, rsigma), (out_margin, _, rsigma_margin) = outputs + torch.testing.assert_close(rsigma_margin, rsigma, atol=0, rtol=0) + for attr in ( + "_rowwise_data", + "_columnwise_data", + "_rowwise_scale_inv", + "_columnwise_scale_inv", + ): + torch.testing.assert_close( + getattr(out_margin, attr), getattr(out, attr), atol=0, rtol=0, msg=attr + ) + + +@pytest.mark.skipif(not recipe_available, reason=reason_for_no_recipe) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_rmsnorm_fwd_mxfp8_2d_quantization(dtype): + """The fused kernels quantize 1D blocks, so with 2D quantization the forward normalizes and + then quantizes 32x32 blocks.""" + torch.manual_seed(1234) + x = torch.randn(256, 1024, dtype=dtype, device="cuda") + weight = torch.randn(1024, dtype=dtype, device="cuda") + + quantizer = MXFP8Quantizer(tex.DType.kFloat8E4M3, with_2d_quantization=True) + out, _, _ = tex.rmsnorm_fwd(x, weight, 1e-5, None, quantizer, TE_DType[dtype], 0, False) + + y, _, _ = tex.rmsnorm_fwd(x, weight, 1e-5, None, None, TE_DType[dtype], 0, False) + ref = MXFP8Quantizer(tex.DType.kFloat8E4M3, with_2d_quantization=True)(y) + for attr in ( + "_rowwise_data", + "_columnwise_data", + "_rowwise_scale_inv", + "_columnwise_scale_inv", + ): + torch.testing.assert_close( + getattr(out, attr).view(torch.uint8), + getattr(ref, attr).view(torch.uint8), + atol=0, + rtol=0, + msg=attr, + ) diff --git a/transformer_engine/common/normalization/common.cpp b/transformer_engine/common/normalization/common.cpp index 84c1fe93482..a7095aa0570 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -578,13 +578,6 @@ bool use_cudnn_mxfp8_norm_output_in_input_dtype() { return flag; } -bool& _cudnn_norm_fwd_mxfp8_flag() { - static bool flag = transformer_engine::getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN"); - return flag; -} - -bool use_cudnn_norm_fwd_mxfp8() { return _cudnn_norm_fwd_mxfp8_flag(); } - } // namespace normalization } // namespace transformer_engine @@ -593,11 +586,6 @@ void nvte_enable_cudnn_norm_fwd(bool enable) { transformer_engine::normalization::_cudnn_norm_fwd_flag() = enable; } -void nvte_enable_cudnn_norm_fwd_mxfp8(bool enable) { - NVTE_API_CALL(nvte_enable_cudnn_norm_fwd_mxfp8); - transformer_engine::normalization::_cudnn_norm_fwd_mxfp8_flag() = enable; -} - void nvte_enable_cudnn_norm_bwd(bool enable) { NVTE_API_CALL(nvte_enable_cudnn_norm_bwd); transformer_engine::normalization::_cudnn_norm_bwd_flag() = enable; diff --git a/transformer_engine/common/normalization/common.h b/transformer_engine/common/normalization/common.h index 19eed8cfea1..26fbaf04351 100644 --- a/transformer_engine/common/normalization/common.h +++ b/transformer_engine/common/normalization/common.h @@ -309,13 +309,13 @@ bool use_cudnn_norm_bwd(); bool& use_zero_centered_gamma_in_weight_dtype(); bool use_cudnn_mxfp8_norm_output_in_input_dtype(); -bool use_cudnn_norm_fwd_mxfp8(); // RMSNorm forward with MXFP8 output by Transformer Engine's fused kernel // (rmsnorm/rmsnorm_fwd_mxfp8.cu), which also writes GEMM-swizzled scaling factors. -bool use_te_rmsnorm_fwd_mxfp8(const Tensor& x, const Tensor& gamma, const Tensor& z); +bool is_supported_by_te_rmsnorm_fwd_mxfp8(const Tensor& x, const Tensor& gamma, const Tensor& z); void rmsnorm_fwd_mxfp8(const Tensor& x, const Tensor& gamma, float epsilon, Tensor* z, - Tensor* rsigma, bool zero_centered_gamma, cudaStream_t stream); + Tensor* rsigma, int multiprocessorCount, bool zero_centered_gamma, + cudaStream_t stream); } // namespace normalization } // namespace transformer_engine diff --git a/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp b/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp index e074f7e72c4..c18e6204f80 100644 --- a/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp +++ b/transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp @@ -48,20 +48,24 @@ void rmsnorm_fwd(const Tensor &x, const Tensor &gamma, const float epsilon, Tens CheckOutputTensor(*rsigma, "rsigma"); } - if (use_te_rmsnorm_fwd_mxfp8(x, gamma, *z)) { - if (workspace->data.numel() == 0) { - // The fused kernel needs no workspace, but a zero size denotes a workspace query. - workspace->data.shape = {1}; - workspace->data.dtype = DType::kByte; + // MXFP8 output: Transformer Engine's fused kernel, which alone writes GEMM-swizzled scales, + // unless cuDNN is requested or the kernel does not support the tensors. + if (is_mxfp8_scaling(z->scaling_mode)) { + const bool te_supported = is_supported_by_te_rmsnorm_fwd_mxfp8(x, gamma, *z); + if (z->with_gemm_swizzled_scales || (!use_cudnn_norm_fwd() && te_supported)) { + NVTE_CHECK(te_supported, + "RMSNorm forward with GEMM-swizzled MXFP8 scales requires FP32, BF16 or FP16 " + "input and weight, and rows and columns that are multiples of 128."); + if (workspace->data.numel() == 0) { + // The fused kernel needs no workspace, but a zero size denotes a workspace query. + workspace->data.shape = {1}; + workspace->data.dtype = DType::kByte; + return; + } + rmsnorm_fwd_mxfp8(x, gamma, epsilon, z, rsigma, multiprocessorCount, zero_centered_gamma, + stream); return; } - rmsnorm_fwd_mxfp8(x, gamma, epsilon, z, rsigma, zero_centered_gamma, stream); - return; - } - if (is_mxfp8_scaling(z->scaling_mode)) { - NVTE_CHECK(!z->with_gemm_swizzled_scales, - "MXFP8 output must have scales in compact format, not swizzled for GEMM, " - "unless Transformer Engine's fused RMSNorm + MXFP8 kernel is used."); } NVTE_Norm_Backend norm_backend; diff --git a/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu b/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu index 103bb3c043a..1ade354ebc2 100644 --- a/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu +++ b/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu @@ -15,16 +15,19 @@ * Two kernels: a pre-pass computes rsigma per row (one warp per row), then one CTA per * 32x128 tile normalizes the tile, quantizes it and writes its scaling factors. A tile * only sees 128 columns of a row, so it cannot reduce the row itself; doing so with whole - * rows per CTA would limit the grid to rows / 32 CTAs. + * rows per CTA would limit the grid to rows / 32 CTAs. When the caller reserves SMs for other + * work, both kernels are launched over chunks of rows that fit on the remaining SMs. */ #include +#include #include #include #include "../../cast/mxfp8/swizzle.cuh" #include "../../common.h" +#include "../../util/cuda_runtime.h" #include "../../util/ptx_arch_spec.cuh" #include "../../utils.cuh" #include "../common.h" @@ -287,9 +290,24 @@ bool is_aligned_to(const void *ptr, size_t alignment) { return reinterpret_cast(ptr) % alignment == 0; } +// Rows per launch of a kernel with grid rows of grid_cols CTAs, each grid row covering cta_rows +// rows: all rows, unless the caller reserves SMs for other work. Then a launch runs at most as +// many CTAs as fit on the remaining num_sms SMs, in multiples of 128 rows (at least 128). +template +int rows_per_launch(Kernel kernel, const int threads, const int num_sms, const int rows, + const int cta_rows, const int grid_cols) { + if (num_sms <= 0 || num_sms >= cuda::sm_count()) { + return rows; + } + int ctas_per_sm = 0; + NVTE_CHECK_CUDA(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&ctas_per_sm, kernel, threads, 0)); + const int grid_rows = std::max(ctas_per_sm, 1) * num_sms / grid_cols; + return std::min(rows, std::max(grid_rows * cta_rows / 128 * 128, 128)); +} + template void launch_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float epsilon, Tensor *z, - Tensor *rsigma, const bool zero_centered_gamma, + Tensor *rsigma, const int num_sms, const bool zero_centered_gamma, const bool gamma_in_weight_dtype, cudaStream_t stream) { const auto [rows_size_t, cols_size_t] = x.flat_2d_dims(); const int rows = static_cast(rows_size_t); @@ -322,39 +340,68 @@ void launch_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float (!rowwise || is_aligned_to(args.rowwise_data, 4)) && (!colwise || is_aligned_to(args.colwise_data, 4)); - TRANSFORMER_ENGINE_SWITCH_CONDITION( - aligned, kAligned, - rmsnorm_mxfp8_rsigma_kernel - <<>>( - x_ptr, rsigma_ptr, rows, cols, epsilon);); // NOLINT(*) + TRANSFORMER_ENGINE_SWITCH_CONDITION(aligned, kAligned, { + const auto kernel = rmsnorm_mxfp8_rsigma_kernel; + constexpr int threads = kRsigmaRowsPerCTA * THREADS_PER_WARP; + const int launch_rows = rows_per_launch(kernel, threads, num_sms, rows, kRsigmaRowsPerCTA, 1); + for (int row = 0; row < rows; row += launch_rows) { + const int n = std::min(launch_rows, rows - row); + kernel<<>>( + x_ptr + static_cast(row) * cols, rsigma_ptr + row, n, cols, epsilon); + } + }); // NOLINT(*) NVTE_CHECK_CUDA(cudaGetLastError()); - const dim3 grid(rows / kTileRows, cols / kTileCols); + // Arguments for the rows from row on, a multiple of 128. The scaling factors of these rows + // start at a fixed offset in each layout: row-wise, after row rows of scales (swizzled: row / + // 128 rows of 128x4 tiles, the same bytes); column-wise, after row / 32 rows of scales + // (swizzled: row / 128 tiles into each column of tiles). + const auto args_from = [&](const int row) { + QuantizeArgs a = args; + const size_t data_offset = static_cast(row) * cols; + if (rowwise) { + a.rowwise_data = static_cast(args.rowwise_data) + data_offset; + a.rowwise_scale_inv += static_cast(row) * z->scale_inv.shape.back(); + } + if (colwise) { + a.colwise_data = static_cast(args.colwise_data) + data_offset; + a.colwise_scale_inv += + swizzled ? static_cast(row / 128) * 512 + : static_cast(row / kScaleBlock) * z->columnwise_scale_inv.shape.back(); + } + return a; + }; + const auto launch = [&](auto kernel) { + const int launch_rows = + rows_per_launch(kernel, kThreads, num_sms, rows, kTileRows, cols / kTileCols); + for (int row = 0; row < rows; row += launch_rows) { + const int n = std::min(launch_rows, rows - row); + kernel<<>>( + x_ptr + static_cast(row) * cols, gamma_ptr, rsigma_ptr + row, args_from(row), + cols, zero_centered_gamma, gamma_in_weight_dtype); + } + }; TRANSFORMER_ENGINE_SWITCH_CONDITION( aligned, kAligned, TRANSFORMER_ENGINE_SWITCH_CONDITION( swizzled, kSwizzled, if (rowwise && colwise) { - rmsnorm_mxfp8_quantize_kernel - <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, - zero_centered_gamma, gamma_in_weight_dtype); + launch(rmsnorm_mxfp8_quantize_kernel); } else if (rowwise) { - rmsnorm_mxfp8_quantize_kernel - <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, - zero_centered_gamma, gamma_in_weight_dtype); + launch(rmsnorm_mxfp8_quantize_kernel); } else { - rmsnorm_mxfp8_quantize_kernel - <<>>(x_ptr, gamma_ptr, rsigma_ptr, args, cols, - zero_centered_gamma, gamma_in_weight_dtype); + launch(rmsnorm_mxfp8_quantize_kernel); });); // NOLINT(*) NVTE_CHECK_CUDA(cudaGetLastError()); } } // namespace -bool use_te_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const Tensor &z) { - if (!is_mxfp8_scaling(z.scaling_mode) || use_cudnn_norm_fwd_mxfp8() || - !is_supported_by_CC_100()) { +bool is_supported_by_te_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const Tensor &z) { + if (!is_mxfp8_scaling(z.scaling_mode) || !is_supported_by_CC_100()) { return false; } const bool rowwise = z.has_data(); @@ -390,7 +437,8 @@ bool use_te_rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const Tensor } void rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float epsilon, Tensor *z, - Tensor *rsigma, const bool zero_centered_gamma, cudaStream_t stream) { + Tensor *rsigma, const int multiprocessorCount, + const bool zero_centered_gamma, cudaStream_t stream) { const DType otype = z->has_data() ? z->data.dtype : z->columnwise_data.dtype; const bool gamma_in_weight_dtype = use_zero_centered_gamma_in_weight_dtype(); TRANSFORMER_ENGINE_TYPE_SWITCH_FLOAT( @@ -400,8 +448,8 @@ void rmsnorm_fwd_mxfp8(const Tensor &x, const Tensor &gamma, const float epsilon TRANSFORMER_ENGINE_TYPE_SWITCH_FP8ONLY( otype, OType, launch_rmsnorm_fwd_mxfp8( - x, gamma, epsilon, z, rsigma, zero_centered_gamma, gamma_in_weight_dtype, - stream);););); // NOLINT(*) + x, gamma, epsilon, z, rsigma, multiprocessorCount, zero_centered_gamma, + gamma_in_weight_dtype, stream);););); // NOLINT(*) } } // namespace normalization diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index ebef28d6341..a50b588e6a7 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -337,25 +337,27 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w UNFUSED, // Compute norm directly FULLY_FUSED, + // Compute norm directly, with quantization scales in compact format + FUSED_NORM_QUANT_UNSWIZZLED, // Compute norm and amax in high precision, then quantize to FP8 FUSED_NORM_AMAX_FP8, // Compute norm and amax in high precision, then quantize to NVFP4 FUSED_NORM_AMAX_NVFP4 }; Impl impl = Impl::UNFUSED; - // Whether the fused kernel can write GEMM-swizzled scaling factors - bool fused_kernel_supports_swizzled_scales = false; if (quantizer.is_none() || IsFloat8Quantizers(quantizer.ptr())) { impl = Impl::FULLY_FUSED; } else if (IsMXFP8Quantizers(quantizer.ptr())) { - if (outer_size % 128 == 0 && inner_size % 128 == 0) { - if (!transformer_engine::getenv("NVTE_NORM_FWD_MXFP8_USE_CUDNN")) { + auto mxfp8_quantizer_cpp = dynamic_cast(quantizer_cpp.get()); + NVTE_CHECK(mxfp8_quantizer_cpp != nullptr, "Could not cast to MXFP8 quantizer"); + // The fused MXFP8 kernels quantize 1D blocks and require full 128x128 tiles + if (!mxfp8_quantizer_cpp->with_2d_quantization && outer_size % 128 == 0 && + inner_size % 128 == 0) { + if (transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN")) { + impl = Impl::FUSED_NORM_QUANT_UNSWIZZLED; + } else { // Transformer Engine's fused RMSNorm + MXFP8 kernel impl = Impl::FULLY_FUSED; - fused_kernel_supports_swizzled_scales = true; - } else if (transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN")) { - // cuDNN MXFP8 kernel requires full 128x128 tiles - impl = Impl::FULLY_FUSED; } } } else if (detail::IsFloat8CurrentScalingQuantizers(quantizer.ptr()) && @@ -382,7 +384,8 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // Output tensor TensorWrapper out_nvte; if (out.is_none()) { - if (impl == Impl::FULLY_FUSED && !fused_kernel_supports_swizzled_scales) { + if (impl == Impl::FUSED_NORM_QUANT_UNSWIZZLED || + (impl == Impl::FULLY_FUSED && !IsMXFP8Quantizers(quantizer.ptr()))) { // FP8 has no special logic to optimize for GEMM, MXFP8 cuDNN // kernel does not support GEMM swizzled scales quantizer_cpp->optimize_for_gemm = false; From 0aef2410b72acb3fcfef44990f37ac35522c4c91 Mon Sep 17 00:00:00 2001 From: Tim Moon <4406448+timmoon10@users.noreply.github.com> Date: Thu, 8 Oct 2026 19:09:31 -0700 Subject: [PATCH 5/7] Apply suggestion from @timmoon10 Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> --- transformer_engine/pytorch/csrc/extensions/normalization.cpp | 2 ++ 1 file changed, 2 insertions(+) diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index a50b588e6a7..6256cd07ad7 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -389,6 +389,8 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // FP8 has no special logic to optimize for GEMM, MXFP8 cuDNN // kernel does not support GEMM swizzled scales quantizer_cpp->optimize_for_gemm = false; + if (impl == Impl::FUSED_NORM_QUANT_UNSWIZZLED) { + quantizer_cpp->optimize_for_gemm = false; } std::tie(out_nvte, out) = quantizer_cpp->create_tensor(shape, out_dtype); } else { From e27db7fbb5c173c1e3f9379cd660f57474c591fd Mon Sep 17 00:00:00 2001 From: Tim Moon <4406448+timmoon10@users.noreply.github.com> Date: Thu, 8 Oct 2026 19:10:45 -0700 Subject: [PATCH 6/7] Apply suggestion from @timmoon10 Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> --- transformer_engine/pytorch/csrc/extensions/normalization.cpp | 5 ----- 1 file changed, 5 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index 6256cd07ad7..e62743fe7f4 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -384,11 +384,6 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // Output tensor TensorWrapper out_nvte; if (out.is_none()) { - if (impl == Impl::FUSED_NORM_QUANT_UNSWIZZLED || - (impl == Impl::FULLY_FUSED && !IsMXFP8Quantizers(quantizer.ptr()))) { - // FP8 has no special logic to optimize for GEMM, MXFP8 cuDNN - // kernel does not support GEMM swizzled scales - quantizer_cpp->optimize_for_gemm = false; if (impl == Impl::FUSED_NORM_QUANT_UNSWIZZLED) { quantizer_cpp->optimize_for_gemm = false; } From 06d34c64549eed732a87b6fc1a09b0957d9115e1 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 9 Oct 2026 02:12:36 +0000 Subject: [PATCH 7/7] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/pytorch/csrc/extensions/normalization.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index e62743fe7f4..e93c01026a9 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -385,7 +385,7 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w TensorWrapper out_nvte; if (out.is_none()) { if (impl == Impl::FUSED_NORM_QUANT_UNSWIZZLED) { - quantizer_cpp->optimize_for_gemm = false; + quantizer_cpp->optimize_for_gemm = false; } std::tie(out_nvte, out) = quantizer_cpp->create_tensor(shape, out_dtype); } else {