From c860887c8cac0000ce2b32a42ad61cbd1dddd555 Mon Sep 17 00:00:00 2001 From: sraman-rgb Date: Thu, 30 Jul 2026 10:24:07 -0700 Subject: [PATCH 1/3] Add opt-in reduced precision output for cuDNN MXFP8 norm Signed-off-by: sraman-rgb --- docs/envvars.rst | 10 ++++++++++ transformer_engine/common/normalization/common.cpp | 10 +++++++++- transformer_engine/common/normalization/common.h | 1 + 3 files changed, 20 insertions(+), 1 deletion(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 466fae9f44..97eaed5ddc 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -379,6 +379,16 @@ Torch Compilation and Fusion LayerNorm/RMSNorm SM Margins ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ +.. envvar:: NVTE_CUDNN_MXFP8_NORM_OUTPUT_IN_INPUT_DTYPE + + :Type: ``int`` (0 or 1) + :Default: ``0`` + :Description: With cuDNN 9.25.0 or later, use the normalization input datatype for the virtual + LayerNorm/RMSNorm output consumed by cuDNN MXFP8 block-scale quantization. This + enables cuDNN's fused MXFP8 normalization engine, which requires matching FP16 or + BF16 input and normalization-output datatypes. When set to ``0``, or with an + earlier cuDNN version, the virtual normalization output uses FP32. + .. envvar:: NVTE_FWD_LAYERNORM_SM_MARGIN :Type: ``int`` diff --git a/transformer_engine/common/normalization/common.cpp b/transformer_engine/common/normalization/common.cpp index 375b109c23..a664dc5a8d 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -304,7 +304,9 @@ CudnnNormalizationPlan::CudnnNormalizationPlan(NVTE_Norm_Type NormType, NVTE_Nor if (_training) _rsigma->set_output(true).set_data_type(get_cudnn_fe_dtype(ctype)); - const auto ZDtype = _fp8_out ? ctype : otype; + const bool use_input_dtype = cudnnGetVersion() >= 92500 && _fp8_out && _ndim_scale_block == 1 && + use_cudnn_mxfp8_norm_output_in_input_dtype(); + const auto ZDtype = use_input_dtype ? itype : (_fp8_out ? ctype : otype); _z->set_output(!_fp8_out).set_data_type(get_cudnn_fe_dtype(ZDtype)); if (_fp8_out) { @@ -562,6 +564,12 @@ bool& _zero_centered_gamma_in_weight_dtype() { bool& use_zero_centered_gamma_in_weight_dtype() { return _zero_centered_gamma_in_weight_dtype(); } +bool use_cudnn_mxfp8_norm_output_in_input_dtype() { + static bool flag = + transformer_engine::getenv("NVTE_CUDNN_MXFP8_NORM_OUTPUT_IN_INPUT_DTYPE"); + return flag; +} + } // namespace normalization } // namespace transformer_engine diff --git a/transformer_engine/common/normalization/common.h b/transformer_engine/common/normalization/common.h index 31a547f62c..dd2f3d1459 100644 --- a/transformer_engine/common/normalization/common.h +++ b/transformer_engine/common/normalization/common.h @@ -308,6 +308,7 @@ bool use_cudnn_norm_fwd(); bool use_cudnn_norm_bwd(); bool& use_zero_centered_gamma_in_weight_dtype(); +bool use_cudnn_mxfp8_norm_output_in_input_dtype(); } // namespace normalization } // namespace transformer_engine From 2f85aff33b9646fa9f838487c41614e266a4d774 Mon Sep 17 00:00:00 2001 From: sraman-rgb Date: Fri, 31 Jul 2026 11:23:31 -0700 Subject: [PATCH 2/3] Warn when cuDNN MXFP8 norm changes intermediate dtype Signed-off-by: sraman-rgb --- transformer_engine/common/normalization/common.cpp | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/transformer_engine/common/normalization/common.cpp b/transformer_engine/common/normalization/common.cpp index a664dc5a8d..745d55ed74 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -306,6 +306,11 @@ CudnnNormalizationPlan::CudnnNormalizationPlan(NVTE_Norm_Type NormType, NVTE_Nor const bool use_input_dtype = cudnnGetVersion() >= 92500 && _fp8_out && _ndim_scale_block == 1 && use_cudnn_mxfp8_norm_output_in_input_dtype(); + if (use_input_dtype) { + NVTE_WARN( + "The cuDNN MXFP8 normalization intermediate output uses the input dtype (itype) " + "instead of the compute dtype; otype still applies to the final quantized output."); + } const auto ZDtype = use_input_dtype ? itype : (_fp8_out ? ctype : otype); _z->set_output(!_fp8_out).set_data_type(get_cudnn_fe_dtype(ZDtype)); From 63829a2340084fbb770fb21044e210f6de046365 Mon Sep 17 00:00:00 2001 From: sraman-rgb Date: Fri, 31 Jul 2026 12:01:23 -0700 Subject: [PATCH 3/3] Scope cuDNN MXFP8 dtype override to FP8 output Signed-off-by: sraman-rgb --- .../common/normalization/common.cpp | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/transformer_engine/common/normalization/common.cpp b/transformer_engine/common/normalization/common.cpp index 745d55ed74..3d00760b52 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -304,14 +304,17 @@ CudnnNormalizationPlan::CudnnNormalizationPlan(NVTE_Norm_Type NormType, NVTE_Nor if (_training) _rsigma->set_output(true).set_data_type(get_cudnn_fe_dtype(ctype)); - const bool use_input_dtype = cudnnGetVersion() >= 92500 && _fp8_out && _ndim_scale_block == 1 && - use_cudnn_mxfp8_norm_output_in_input_dtype(); - if (use_input_dtype) { - NVTE_WARN( - "The cuDNN MXFP8 normalization intermediate output uses the input dtype (itype) " - "instead of the compute dtype; otype still applies to the final quantized output."); + auto ZDtype = _fp8_out ? ctype : otype; + if (_fp8_out) { + const bool use_input_dtype = cudnnGetVersion() >= 92500 && _ndim_scale_block == 1 && + use_cudnn_mxfp8_norm_output_in_input_dtype(); + if (use_input_dtype) { + NVTE_WARN( + "The cuDNN MXFP8 normalization intermediate output uses the input dtype (itype) " + "instead of the compute dtype; otype still applies to the final quantized output."); + ZDtype = itype; + } } - const auto ZDtype = use_input_dtype ? itype : (_fp8_out ? ctype : otype); _z->set_output(!_fp8_out).set_data_type(get_cudnn_fe_dtype(ZDtype)); if (_fp8_out) {