diff --git a/tests/cpp/operator/test_normalization_mxfp8.cu b/tests/cpp/operator/test_normalization_mxfp8.cu index 10b33f8e2a1..84c244ff1ff 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) { 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(use_cudnn); // 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(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,95 @@ 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. +// 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) { + 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 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); + z_swizzled.set_with_gemm_swizzled_scales(true); + fillUniform(&input); + fillUniform(&gamma); + + const auto run = [&](Tensor &out, Tensor &rs, const int sm_count) { + Tensor workspace; + 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(), 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()), + 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 +340,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 +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 = 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); + performTest(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype, + use_cudnn); ); ); } @@ -277,7 +376,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 +386,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..5200ec1074b --- /dev/null +++ b/tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py @@ -0,0 +1,132 @@ +# 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_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) + + +@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/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/normalization/common.h b/transformer_engine/common/normalization/common.h index dd2f3d14591..26fbaf04351 100644 --- a/transformer_engine/common/normalization/common.h +++ b/transformer_engine/common/normalization/common.h @@ -310,6 +310,13 @@ bool use_cudnn_norm_bwd(); bool& use_zero_centered_gamma_in_weight_dtype(); bool use_cudnn_mxfp8_norm_output_in_input_dtype(); +// RMSNorm forward with MXFP8 output by Transformer Engine's fused kernel +// (rmsnorm/rmsnorm_fwd_mxfp8.cu), which also writes GEMM-swizzled scaling factors. +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, 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 f56e2ef1fef..c18e6204f80 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,26 @@ void rmsnorm_fwd(const Tensor &x, const Tensor &gamma, const float epsilon, Tens CheckOutputTensor(*rsigma, "rsigma"); } + // 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; + } + } + 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..1ade354ebc2 --- /dev/null +++ b/transformer_engine/common/normalization/rmsnorm/rmsnorm_fwd_mxfp8.cu @@ -0,0 +1,456 @@ +/************************************************************************* + * 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. 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" + +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; +} + +// 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 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); + 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, { + 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()); + + // 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) { + launch(rmsnorm_mxfp8_quantize_kernel); + } else if (rowwise) { + launch(rmsnorm_mxfp8_quantize_kernel); + } else { + launch(rmsnorm_mxfp8_quantize_kernel); + });); // NOLINT(*) + NVTE_CHECK_CUDA(cudaGetLastError()); +} + +} // namespace + +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(); + 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 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( + 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, multiprocessorCount, 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..e93c01026a9 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -337,6 +337,8 @@ 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 @@ -346,10 +348,17 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w 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 && + 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) { - // cuDNN MXFP8 kernel requires full 128x128 tiles - impl = Impl::FULLY_FUSED; + 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; + } } } else if (detail::IsFloat8CurrentScalingQuantizers(quantizer.ptr()) && !transformer_engine::getenv("NVTE_NORM_FWD_USE_CUDNN")) { @@ -375,9 +384,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) { - // FP8 has no special logic to optimize for GEMM, MXFP8 cuDNN - // kernel does not support GEMM swizzled scales + 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);