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..3d00760b52 100644 --- a/transformer_engine/common/normalization/common.cpp +++ b/transformer_engine/common/normalization/common.cpp @@ -304,7 +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 auto ZDtype = _fp8_out ? ctype : otype; + 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; + } + } _z->set_output(!_fp8_out).set_data_type(get_cudnn_fe_dtype(ZDtype)); if (_fp8_out) { @@ -562,6 +572,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