Skip to content
150 changes: 144 additions & 6 deletions tests/cpp/operator/test_normalization_mxfp8.cu
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,8 @@ void dequantize_2x(Tensor& input, Tensor& output, bool is_training)
}

template <typename InputType, typename OutputType>
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);
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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<size_t>{ N, H }, DType::kFloat32, true, true);

Expand Down Expand Up @@ -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();
Expand Down Expand Up @@ -235,6 +240,95 @@ std::vector<std::pair<size_t, size_t>> 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 <typename InputType, typename OutputType>
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<InputType>::dtype;
const DType otype = TypeInfo<OutputType>::dtype;
Tensor input("input", std::vector<size_t>{ N, H }, itype);
Tensor gamma("gamma", std::vector<size_t>{ H }, itype);
Tensor rsigma("rsigma", std::vector<size_t>{ N }, DType::kFloat32);
Tensor rsigma_swizzled("rsigma_swizzled", std::vector<size_t>{ N }, DType::kFloat32);
Tensor z("z", std::vector<size_t>{ N, H }, otype, true, is_training, NVTE_MXFP8_1D_SCALING);
Tensor z_swizzled("z_swizzled", std::vector<size_t>{ 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<float>(), true, 0, 0);

const auto bytes = [](const OutputType *p) { return reinterpret_cast<const uint8_t *>(p); };
compareResults("rowwise_data", bytes(z_swizzled.rowwise_cpu_dptr<OutputType>()),
bytes(z.rowwise_cpu_dptr<OutputType>()), 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<uint8_t> ref(row_dim0 * row_dim1);
const uint8_t *compact = z.rowwise_cpu_scale_inv_ptr<uint8_t>();
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<uint8_t>(), ref.data(),
ref.size());

if (is_training) {
compareResults("colwise_data", bytes(z_swizzled.columnwise_cpu_dptr<OutputType>()),
bytes(z.columnwise_cpu_dptr<OutputType>()), 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<uint8_t>();
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<uint8_t>(),
ref.data(), ref.size());
}
}

std::vector<std::pair<size_t, size_t>> swizzled_scales_test_cases = {
{128, 128},
{768, 2304},
{256, 7168},
};

std::vector<NormType> norms = {
NormType::LayerNorm,
NormType::RMSNorm
Expand All @@ -246,7 +340,7 @@ class MxNormTestSuite : public ::testing::TestWithParam< std::tuple<NormType,
transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>,
bool, bool, bool>> {};
bool, bool, bool, bool>> {};

TEST_P(MxNormTestSuite, TestMxNorm) {
using namespace transformer_engine;
Expand All @@ -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<InputType, OutputType>(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma, norm_type, is_training, zero_centered_gamma_in_weight_dtype,
use_cudnn);
);
);
}
Expand All @@ -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<MxNormTestSuite::ParamType>& info) {
std::string name = normToString.at(std::get<0>(info.param)) + "_" +
test::typeName(std::get<1>(info.param)) + "X" +
Expand All @@ -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<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>, 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<InputType, OutputType>(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<MxNormSwizzledScalesTestSuite::ParamType>& 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";
});
132 changes: 132 additions & 0 deletions tests/pytorch/mxfp8/test_mxfp8_rmsnorm_fwd.py
Original file line number Diff line number Diff line change
@@ -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,
)
1 change: 1 addition & 0 deletions transformer_engine/common/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 7 additions & 0 deletions transformer_engine/common/normalization/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading