From aeff92bda7d60ace0ea19ed4fb939f4a0f7887d0 Mon Sep 17 00:00:00 2001 From: zhongboz Date: Tue, 29 Sep 2026 03:45:21 -0700 Subject: [PATCH 1/3] [PyTorch] Deprecate grouped linear and MLP environment gates Signed-off-by: zhongboz --- .../benchmark_graph_safe_grouped_mlp.py | 2 - docs/envvars.rst | 29 +++ docs/examples/te_mixtral/run_finetune_ep.py | 12 -- docs/examples/te_mixtral/te_mixtral_mxfp8.py | 6 +- ...torial_accelerate_hf_mixtral_with_te.ipynb | 2 +- docs/examples/te_mixtral/utils.py | 15 +- qa/L0_pytorch_debug_unittest/test.sh | 2 +- qa/L0_pytorch_unittest/test.sh | 6 +- tests/pytorch/test_grouped_linear.py | 66 +++++-- tests/pytorch/test_grouped_mlp.py | 173 ++++++++++++++---- tests/pytorch/test_grouped_tensor.py | 5 - ...odule_grouped_linear_distributed_weight.py | 3 +- tests/pytorch/test_sanity.py | 8 +- .../pytorch/module/grouped_linear.py | 19 +- .../pytorch/ops/basic/grouped_linear.py | 16 +- .../pytorch/ops/fused/grouped_mlp.py | 67 ++++--- transformer_engine/pytorch/utils.py | 21 +-- 17 files changed, 293 insertions(+), 159 deletions(-) diff --git a/benchmarks/linear/benchmark_graph_safe_grouped_mlp.py b/benchmarks/linear/benchmark_graph_safe_grouped_mlp.py index 00f7f516c88..973329a19c3 100644 --- a/benchmarks/linear/benchmark_graph_safe_grouped_mlp.py +++ b/benchmarks/linear/benchmark_graph_safe_grouped_mlp.py @@ -35,7 +35,6 @@ os.environ.setdefault("CUDA_DEVICE_MAX_CONNECTIONS", "1") os.environ.setdefault("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1") -os.environ.setdefault("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "1") os.environ.setdefault("CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL", "1") import argparse @@ -318,7 +317,6 @@ def main() -> None: for name in ( "CUDA_DEVICE_MAX_CONNECTIONS", "NVTE_ALLOW_NONDETERMINISTIC_ALGO", - "NVTE_CUTEDSL_FUSED_GROUPED_MLP", "CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL", ): print(f" {name}={os.environ.get(name)}") diff --git a/docs/envvars.rst b/docs/envvars.rst index 46b70bbe46a..61d80dd0e3e 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -134,6 +134,35 @@ Runtime Environment Variables These environment variables control the behavior of Transformer Engine during execution. +Deprecated Grouped Linear and MLP Controls +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. envvar:: NVTE_GROUPED_LINEAR_SINGLE_PARAM + + :Status: **DEPRECATED**. No longer read. + :Historical default: ``0`` (before deprecation). + :Description: Previously permitted grouped-linear weights and biases to use + a single grouped parameter. This variable is now ignored, + including when set to ``0`` or ``1``. + :Migration: Remove this variable from launch scripts. Select the parameter + layout with the ``single_grouped_weight`` and + ``single_grouped_bias`` constructor arguments; both still + default to ``False``. + +.. envvar:: NVTE_CUTEDSL_FUSED_GROUPED_MLP + + :Status: **DEPRECATED**. No longer read. + :Historical default: ``0`` (before deprecation). + :Description: Previously enabled CuTeDSL grouped MLP fusion in the PyTorch + operation fuser. This variable is now ignored, including when + set to ``0`` or ``1``; setting it to ``0`` no longer disables + fusion. + :Migration: Remove this variable from launch scripts. The operation fuser + automatically selects grouped MLP fusion for compatible + operation sequences when the GPU, recipe, and installed + dependencies support it. No environment variable needs to be + set before importing TE to enable this fusion. + General ^^^^^^^ diff --git a/docs/examples/te_mixtral/run_finetune_ep.py b/docs/examples/te_mixtral/run_finetune_ep.py index e389b7986b5..fdf4392ed8e 100644 --- a/docs/examples/te_mixtral/run_finetune_ep.py +++ b/docs/examples/te_mixtral/run_finetune_ep.py @@ -19,18 +19,6 @@ import argparse import os -import sys - -# Improvement 3 needs ``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1`` set before TE is imported, -# because the fused-grouped-MLP fusion is registered at module-import time -# inside ``if ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8.is_supported(): ...``. -for _i, _arg in enumerate(sys.argv[1:]): - if _arg == "--improvement" and _i + 2 < len(sys.argv) and sys.argv[_i + 2] == "3": - os.environ["NVTE_CUTEDSL_FUSED_GROUPED_MLP"] = "1" - break - if _arg == "--improvement=3": - os.environ["NVTE_CUTEDSL_FUSED_GROUPED_MLP"] = "1" - break from utils import HyperParameters, run_hf_baseline_finetune, run_te_mixtral_finetune diff --git a/docs/examples/te_mixtral/te_mixtral_mxfp8.py b/docs/examples/te_mixtral/te_mixtral_mxfp8.py index cca7d23085c..c35effeeac9 100644 --- a/docs/examples/te_mixtral/te_mixtral_mxfp8.py +++ b/docs/examples/te_mixtral/te_mixtral_mxfp8.py @@ -11,9 +11,9 @@ MXFP8. HF gate (``w1``) and up (``w3``) weights are row-interleaved in blocks of 32 to match the GLU interleaved layout that fused kernel reads. -The fused kernel is enabled by ``utils._enable_fused_mxfp8_grouped_mlp()`` -(sets ``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1`` and patches the SM-version / -cudnn-frontend signature checks). Requires +Supported grouped-MLP sequences are fused automatically. +``utils._enable_fused_mxfp8_grouped_mlp()`` provides legacy SM-version / +cudnn-frontend signature compatibility patches. Requires ``nvidia-cudnn-frontend >= 1.23.0`` and SM>=10 (Blackwell B100/B200/B300+). """ diff --git a/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb b/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb index 8463c387b4c..e757085862b 100644 --- a/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb +++ b/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb @@ -352,7 +352,7 @@ "\n", "Note\n", "\n", - "`NVTE_CUTEDSL_FUSED_GROUPED_MLP=1` must be set before TE imports the fused op registration. In this tutorial, `run_finetune_ep.py` already does that automatically for improvement 3.\n", + "TE automatically selects grouped-MLP fusion when the operation sequence, GPU, recipe, and installed dependencies support it. No environment variable needs to be set before importing TE to enable this fusion.\n", "\n", "" ] diff --git a/docs/examples/te_mixtral/utils.py b/docs/examples/te_mixtral/utils.py index 559bab885ee..8701afcb6b5 100644 --- a/docs/examples/te_mixtral/utils.py +++ b/docs/examples/te_mixtral/utils.py @@ -149,21 +149,16 @@ def init_baseline_model(hyperparams: HyperParameters): def _enable_fused_mxfp8_grouped_mlp() -> None: - """Improvement 3: enable the fused ``ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8`` and - backward kernel in the installed TE without recompiling. + """Improvement 3: adapt legacy fused grouped-MLP kernels without recompiling. - ``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1`` must be set *before* - ``transformer_engine.pytorch.ops`` is imported — the fusion is registered - at TE module-import-time. ``run_finetune_ep.py`` sniffs ``--improvement 3`` - and sets the env var before importing ``utils``. + Current TE selects supported grouped-MLP fusions automatically. The + compatibility patches below target older TE implementations. - We also (a) relax the SM-version check from ``!= 10`` to ``>= 10`` so + We (a) relax the SM-version check from ``!= 10`` to ``>= 10`` so SM>=11 successors of B300 fire the kernel, and (b) wrap the cudnn-frontend grouped-GEMM wrappers so the installed TE's ``c_dtype`` kwarg (dropped by cudnn-frontend 1.23.0) is silently filtered out. """ - os.environ["NVTE_CUTEDSL_FUSED_GROUPED_MLP"] = "1" - import inspect import cudnn # type: ignore from transformer_engine.pytorch.ops.fused import forward_grouped_mlp as _fwd_mod @@ -172,8 +167,6 @@ def _enable_fused_mxfp8_grouped_mlp() -> None: def _make_is_supported(kernel_method_names): def _is_supported(cls) -> bool: - if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - return False if get_device_compute_capability()[0] < 10: return False try: diff --git a/qa/L0_pytorch_debug_unittest/test.sh b/qa/L0_pytorch_debug_unittest/test.sh index 36efe485f5c..9e024e22e32 100644 --- a/qa/L0_pytorch_debug_unittest/test.sh +++ b/qa/L0_pytorch_debug_unittest/test.sh @@ -41,7 +41,7 @@ NVTE_TORCH_COMPILE=0 pytest -v -s --junitxml=$XML_LOG_DIR/test_api_features.xml pytest -v -s --junitxml=$XML_LOG_DIR/test_perf.xml $TE_PATH/tests/pytorch/debug/test_perf.py --feature_dirs=$NVTE_TEST_NVINSPECT_FEATURE_DIRS --configs_dir=$NVTE_TEST_NVINSPECT_CONFIGS_DIR || test_fail "test_perf.py" # standard sanity and numerics tests with initialized debug -NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_TEST_NVINSPECT_ENABLED=1 NVTE_TEST_NVINSPECT_CONFIG_FILE=$NVTE_TEST_NVINSPECT_DUMMY_CONFIG_FILE NVTE_TEST_NVINSPECT_FEATURE_DIRS=$NVTE_TEST_NVINSPECT_FEATURE_DIRS PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 pytest -v -s --junitxml=$XML_LOG_DIR/test_sanity_2.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "debug test_sanity.py" +NVTE_TEST_NVINSPECT_ENABLED=1 NVTE_TEST_NVINSPECT_CONFIG_FILE=$NVTE_TEST_NVINSPECT_DUMMY_CONFIG_FILE NVTE_TEST_NVINSPECT_FEATURE_DIRS=$NVTE_TEST_NVINSPECT_FEATURE_DIRS PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 pytest -v -s --junitxml=$XML_LOG_DIR/test_sanity_2.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "debug test_sanity.py" NVTE_TEST_NVINSPECT_ENABLED=1 NVTE_TEST_NVINSPECT_CONFIG_FILE=$NVTE_TEST_NVINSPECT_DUMMY_CONFIG_FILE NVTE_TEST_NVINSPECT_FEATURE_DIRS=$NVTE_TEST_NVINSPECT_FEATURE_DIRS PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED_ATTN=0 pytest -v -s --junitxml=$XML_LOG_DIR/test_numerics_2.xml $TE_PATH/tests/pytorch/test_numerics.py || test_fail "debug test_numerics.py" if [ "$RET" -ne 0 ]; then diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index b3b6ccacac7..d0a2855a427 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -30,7 +30,7 @@ export NVTE_FLASH_ATTN_V4=0 pip3 install pytest==8.2.1 || error_exit "Failed to install pytest" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_optional_flash_attn_import.xml $TE_PATH/tests/pytorch/test_optional_flash_attn_import.py || test_fail "test_optional_flash_attn_import.py" -NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_sanity.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "test_sanity.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_sanity.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "test_sanity.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_recipe.xml $TE_PATH/tests/pytorch/test_recipe.py || test_fail "test_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_custom_recipe.xml $TE_PATH/tests/pytorch/test_custom_recipe.py || test_fail "test_custom_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_deferred_init.xml $TE_PATH/tests/pytorch/test_deferred_init.py || test_fail "test_deferred_init.py" @@ -47,7 +47,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_float8blockwiset python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_float8_blockwise_scaling_exact.xml $TE_PATH/tests/pytorch/test_float8_blockwise_scaling_exact.py || test_fail "test_float8_blockwise_scaling_exact.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_float8_blockwise_gemm_exact.xml $TE_PATH/tests/pytorch/test_float8_blockwise_gemm_exact.py || test_fail "test_float8_blockwise_gemm_exact.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_float8_current_scaling_exact.xml $TE_PATH/tests/pytorch/test_float8_current_scaling_exact.py || test_fail "test_float8_current_scaling_exact.py" -NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/test_grouped_tensor.xml $TE_PATH/tests/pytorch/test_grouped_tensor.py || test_fail "test_grouped_tensor.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/test_grouped_tensor.xml $TE_PATH/tests/pytorch/test_grouped_tensor.py || test_fail "test_grouped_tensor.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gqa.xml $TE_PATH/tests/pytorch/test_gqa.py || test_fail "test_gqa.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_qk_norm.xml $TE_PATH/tests/pytorch/test_qk_norm.py || test_fail "test_qk_norm.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_optimizer.xml $TE_PATH/tests/pytorch/test_fused_optimizer.py || test_fail "test_fused_optimizer.py" @@ -88,7 +88,7 @@ NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_DISABLE_TRITON_AUTOTUNING=1 NVIDIA_TF32_ PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_linear.xml $TE_PATH/tests/pytorch/test_grouped_linear.py || test_fail "test_grouped_linear.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_ops_grouped_linear_distributed_weight.xml $TE_PATH/tests/pytorch/test_ops_grouped_linear_distributed_weight.py || test_fail "test_ops_grouped_linear_distributed_weight.py" PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_module_grouped_linear_distributed_weight.xml $TE_PATH/tests/pytorch/test_module_grouped_linear_distributed_weight.py || test_fail "test_module_grouped_linear_distributed_weight.py" -NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_grouped_mlp.xml $TE_PATH/tests/pytorch/test_grouped_mlp.py || test_fail "test_grouped_mlp.py" if [ "$RET" -ne 0 ]; then echo "Error in the following test cases:$FAILED_CASES" diff --git a/tests/pytorch/test_grouped_linear.py b/tests/pytorch/test_grouped_linear.py index e294da0d718..774e758a637 100644 --- a/tests/pytorch/test_grouped_linear.py +++ b/tests/pytorch/test_grouped_linear.py @@ -1698,12 +1698,53 @@ def _reset_fp8_state(monkeypatch): monkeypatch.delenv(_FUSED_GROUPED_GEMM_ENV, raising=False) +@pytest.mark.parametrize("use_op_fuser", (False, True), ids=("module", "op")) +@pytest.mark.parametrize( + "single_grouped_weight,single_grouped_bias", + [(None, None), (False, False), (True, False), (False, True), (True, True)], + ids=("default", "discrete", "single-weight", "single-bias", "single-weight-and-bias"), +) +def test_grouped_parameter_layout_selected_by_constructor( + use_op_fuser, single_grouped_weight, single_grouped_bias +): + """Constructor options alone select registered parameters; defaults stay per-expert.""" + kwargs = {} + if single_grouped_weight is not None: + kwargs.update( + single_grouped_weight=single_grouped_weight, + single_grouped_bias=single_grouped_bias, + ) + if use_op_fuser: + linear = te.ops.GroupedLinear(2, 64, 64, device="cuda", dtype=torch.bfloat16, **kwargs) + else: + linear = GroupedLinear( + 2, + 64, + 64, + device="cuda", + params_dtype=torch.bfloat16, + use_grouped_tensor=True, + **kwargs, + ) + + assert linear.single_grouped_weight is bool(single_grouped_weight) + assert linear.single_grouped_bias is bool(single_grouped_bias) + parameters = dict(linear.named_parameters()) + expected_weights = {"weight"} if single_grouped_weight else {"weight0", "weight1"} + expected_biases = {"bias"} if single_grouped_bias else {"bias0", "bias1"} + assert parameters.keys() == expected_weights | expected_biases + if single_grouped_weight: + assert isinstance(parameters["weight"], GroupedTensor) + if single_grouped_bias: + assert isinstance(parameters["bias"], GroupedTensor) + + @pytest.mark.parametrize( "m_splits,exception", [([256, 256], ValueError), (torch.tensor([256, 256]), ValueError)], ids=["python-list", "cpu-tensor"], ) -def test_single_grouped_weight_rejects_host_m_splits(monkeypatch, m_splits, exception): +def test_single_grouped_weight_rejects_host_m_splits(m_splits, exception): """A single parent parameter must never fall back to host-split per-expert GEMMs.""" if not is_module_grouped_tensor_path_supported( None, @@ -1711,7 +1752,6 @@ def test_single_grouped_weight_rejects_host_m_splits(monkeypatch, m_splits, exce ): pytest.skip("Native GroupedTensor GEMM is unavailable on this system.") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") grouped_linear = GroupedLinear( 2, 64, @@ -1751,7 +1791,7 @@ def test_single_grouped_weight_rejects_host_m_splits(monkeypatch, m_splits, exce ), ], ) -def test_single_grouped_weight_matches_discrete_grouped_tensor_path(monkeypatch, fp8_recipe): +def test_single_grouped_weight_matches_discrete_grouped_tensor_path(fp8_recipe): """Match single and discrete weights while both use CUDA m_splits and grouped GEMM.""" if not is_module_grouped_tensor_path_supported( fp8_recipe, @@ -1759,7 +1799,6 @@ def test_single_grouped_weight_matches_discrete_grouped_tensor_path(monkeypatch, ): pytest.skip("Recipe is not supported with a single grouped weight on this system.") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") FP8GlobalStateManager.reset() num_gemms = 3 @@ -1820,7 +1859,7 @@ def test_single_grouped_weight_matches_discrete_grouped_tensor_path(monkeypatch, @pytest.mark.skipif(not _mxfp8_available, reason=_reason_for_no_mxfp8) -def test_single_grouped_weight_mxfp8_workspace_cache(monkeypatch): +def test_single_grouped_weight_mxfp8_workspace_cache(): """BF16 primary weights update one persistent MXFP8 grouped workspace per iteration.""" mxfp8_recipe = recipe.MXFP8BlockScaling() if not is_module_grouped_tensor_path_supported( @@ -1828,7 +1867,6 @@ def test_single_grouped_weight_mxfp8_workspace_cache(monkeypatch): torch.bfloat16, ): pytest.skip("MXFP8 single-weight GroupedTensor path is unavailable on this system.") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") FP8GlobalStateManager.reset() grouped_linear = GroupedLinear( 2, @@ -1873,9 +1911,8 @@ def test_single_grouped_weight_mxfp8_workspace_cache(monkeypatch): @pytest.mark.skipif(not _mxfp8_available, reason=_reason_for_no_mxfp8) @pytest.mark.parametrize("fp8_recipe", [recipe.MXFP8BlockScaling()], ids=recipe_id) -def test_single_grouped_weight_with_disabled_weight_preswizzle(monkeypatch, fp8_recipe): +def test_single_grouped_weight_with_disabled_weight_preswizzle(fp8_recipe): """Grouped weight preparation preserves a disabled preswizzle decision.""" - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") FP8GlobalStateManager.reset() with quantized_model_init(enabled=True, recipe=fp8_recipe): grouped_linear = GroupedLinear( @@ -1920,7 +1957,7 @@ def test_single_grouped_weight_with_disabled_weight_preswizzle(monkeypatch, fp8_ @pytest.mark.skipif(not _mxfp8_available, reason=_reason_for_no_mxfp8) -def test_single_grouped_primary_mxfp8_bypasses_weight_workspace(monkeypatch): +def test_single_grouped_primary_mxfp8_bypasses_weight_workspace(): """An MXFP8 primary grouped parameter is already GEMM-ready and is not requantized.""" mxfp8_recipe = recipe.MXFP8BlockScaling() if not is_module_grouped_tensor_path_supported( @@ -1928,7 +1965,6 @@ def test_single_grouped_primary_mxfp8_bypasses_weight_workspace(monkeypatch): torch.bfloat16, ): pytest.skip("MXFP8 single-weight GroupedTensor path is unavailable on this system.") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") FP8GlobalStateManager.reset() with quantized_model_init(enabled=True, recipe=mxfp8_recipe): grouped_linear = GroupedLinear( @@ -2136,7 +2172,6 @@ def _run_grouped_parameter_layout( @pytest.mark.parametrize("delay_wgrad_compute", _ALL_BOOLEAN) @pytest.mark.parametrize("fuse_wgrad_accumulation", _ALL_BOOLEAN) def test_grouped_parameter_layout_matches_cpu_m_splits( - monkeypatch, use_bias, single_grouped_weight, single_grouped_bias, @@ -2170,7 +2205,6 @@ def test_grouped_parameter_layout_matches_cpu_m_splits( biases = (0.1 * torch.randn(num_gemms, out_features, device="cuda")).to(torch.bfloat16) # The CPU m_splits baseline is explicitly the legacy, discrete-parameter contract. - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") reference = _run_grouped_parameter_layout( use_grouped_tensor=False, fp8_recipe=fp8_recipe, @@ -2186,9 +2220,7 @@ def test_grouped_parameter_layout_matches_cpu_m_splits( m_splits=m_splits, ) - # Enable single parameters only for the CUDA m_splits target. The explicit layout flags - # below still decide whether this particular case uses discrete or grouped parameters. - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") + # The explicit layout flags select discrete or grouped parameters for the CUDA target. result = _run_grouped_parameter_layout( use_grouped_tensor=True, fp8_recipe=fp8_recipe, @@ -2269,7 +2301,6 @@ def reject_split_fallback(*_args, **_kwargs): "transformer_engine.pytorch.module._split_quantization._split_quantize", reject_split_fallback, ) - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") torch.manual_seed(1234) num_gemms = 2 @@ -2527,7 +2558,7 @@ def test_grouped_linear_grouped_tensor_path_single_grouped_bias_delay_wgrad(monk grouped_linear.backward_dw() -def test_grouped_linear_returns_single_grouped_bias_parameter(monkeypatch): +def test_grouped_linear_returns_single_grouped_bias_parameter(): """return_bias preserves the grouped parent and accumulates dbias into it. This mirrors how MCore applies a returned MoE bias:: @@ -2552,7 +2583,6 @@ def test_grouped_linear_returns_single_grouped_bias_parameter(monkeypatch): torch.bfloat16, ): pytest.skip("BF16 GroupedTensor path is unavailable on this system.") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") dtype = torch.bfloat16 num_gemms = 2 diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 0dd63bf2e81..aabe2aed34a 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -7,6 +7,7 @@ from collections.abc import Iterable import contextlib import functools +import importlib.util import os import math import random @@ -138,6 +139,142 @@ def _reset_rng_states_per_test(): yield +@pytest.fixture +def isolated_grouped_mlp_module(monkeypatch): + """Import the module with a private fusion registry and fresh capability caches.""" + from transformer_engine.pytorch import utils as te_utils + from transformer_engine.pytorch.ops.fuser import OperationFuser + + monkeypatch.delenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP", raising=False) + + def reject_cuda_query(*args, **kwargs): + raise AssertionError("Importing grouped MLP must not query or initialize CUDA") + + spec = importlib.util.spec_from_file_location( + f"{grouped_mlp_module.__package__}._test_grouped_mlp", grouped_mlp_module.__file__ + ) + module = importlib.util.module_from_spec(spec) + with monkeypatch.context() as import_patch: + import_patch.setattr(te_utils, "get_device_compute_capability", reject_cuda_query) + import_patch.setattr(torch.cuda, "is_available", reject_cuda_query) + import_patch.setattr(torch.cuda, "_lazy_init", reject_cuda_query) + import_patch.setattr(OperationFuser, "forward_backward_fusion_functions", []) + spec.loader.exec_module(module) + registered_fusions = tuple(OperationFuser.forward_backward_fusion_functions) + return module, registered_fusions + + +def test_grouped_mlp_registers_without_cuda_query(isolated_grouped_mlp_module) -> None: + """Both default fusion callbacks must be registered before capability is known.""" + module, registered_fusions = isolated_grouped_mlp_module + assert set(registered_fusions) == {module.fuse_glu_ops, module.fuse_unary_activation_ops} + + +@pytest.mark.parametrize( + "pipeline,quantization", + ( + ("unrelated", None), + ("unrelated", "mxfp8"), + ("unrelated", "nvfp4"), + ("glu", "mxfp8_hybrid"), + ("unary", "mxfp8_hybrid"), + ("glu", "nvfp4_no_rht"), + ("unary", "nvfp4_no_rht"), + ), +) +def test_grouped_mlp_ineligible_pipeline_skips_capability_queries( + monkeypatch, isolated_grouped_mlp_module, pipeline: str, quantization: Optional[str] +) -> None: + """Unrelated ops and unsupported recipes must not trigger device or cuDNN probes.""" + from transformer_engine.common.recipe import Format, MXFP8BlockScaling, NVFP4BlockScaling + + module, registered_fusions = isolated_grouped_mlp_module + + def reject_capability_query(*args, **kwargs): + raise AssertionError("Ineligible pipelines must bypass grouped MLP capability checks") + + monkeypatch.setattr(torch.cuda, "is_available", reject_capability_query) + monkeypatch.setattr(module, "get_pkg_version", reject_capability_query) + recipe = None + if quantization in ("mxfp8", "mxfp8_hybrid"): + fp8_format = Format.HYBRID if quantization == "mxfp8_hybrid" else Format.E4M3 + recipe = MXFP8BlockScaling(fp8_format=fp8_format) + elif quantization in ("nvfp4", "nvfp4_no_rht"): + recipe = NVFP4BlockScaling(disable_rht=quantization == "nvfp4_no_rht") + + if pipeline == "unrelated": + ops = [te.ops.ReLU() for _ in range(3)] + else: + is_glu = pipeline == "glu" + kwargs = {"bias": False, "device": "meta", "dtype": torch.bfloat16} + ops = [ + te.ops.GroupedLinear(2, 64, 128 if is_glu else 64, **kwargs), + te.ops.ScaledSwiGLU(glu_interleave_size=32) if is_glu else te.ops.ScaledSReLU(), + te.ops.GroupedLinear(2, 64, 64, **kwargs), + ] + for fuse in registered_fusions: + assert fuse(ops, recipe=recipe) is ops + + +@pytest.mark.parametrize("activation", ("glu", "unary")) +@pytest.mark.parametrize( + "capability", ("supported", "cuda", "architecture", "frontend", "wrapper", "recipe", "shape") +) +def test_grouped_mlp_default_fusion_capabilities( + monkeypatch, isolated_grouped_mlp_module, activation: str, capability: str +) -> None: + """Default fusion still respects hardware, dependency, recipe, and shape support.""" + from transformer_engine.common.recipe import Format, MXFP8BlockScaling + + module, registered_fusions = isolated_grouped_mlp_module + monkeypatch.setattr(torch.cuda, "is_available", lambda: capability != "cuda") + monkeypatch.setattr( + module, + "get_device_compute_capability", + lambda: (9, 0) if capability == "architecture" else (10, 0), + ) + monkeypatch.setattr( + module, "get_pkg_version", lambda _: "1.22.0" if capability == "frontend" else "1.24.0" + ) + + def unused_kernel(): + raise AssertionError("Fusion planning must not execute a kernel") + + fake_cudnn = types.ModuleType("cudnn") + for name in ("glu", "dglu", "srelu", "dsrelu", "quant", "wgrad"): + setattr(fake_cudnn, f"grouped_gemm_{name}_wrapper_sm100", unused_kernel) + if capability == "wrapper": + del fake_cudnn.grouped_gemm_quant_wrapper_sm100 + monkeypatch.setitem(sys.modules, "cudnn", fake_cudnn) + + hidden_size = 96 if capability == "shape" else 64 + if activation == "glu": + activation_op = te.ops.ScaledSwiGLU(glu_interleave_size=32) + fc1_out_features = 128 + fuse = module.fuse_glu_ops + fused_cls = module.GroupedMLP_CuTeGEMMGLU + else: + activation_op = te.ops.ScaledSReLU() + fc1_out_features = 64 + fuse = module.fuse_unary_activation_ops + fused_cls = module.GroupedMLP_CuTeGEMMUnary + assert fuse in registered_fusions + kwargs = {"bias": False, "device": "meta", "dtype": torch.bfloat16} + ops = [ + te.ops.GroupedLinear(2, hidden_size, fc1_out_features, **kwargs), + activation_op, + te.ops.GroupedLinear(2, 64, hidden_size, **kwargs), + ] + recipe = None if capability == "recipe" else MXFP8BlockScaling(fp8_format=Format.E4M3) + fused_ops = fuse(ops, recipe=recipe) + if capability == "supported": + assert len(fused_ops) == 1 + assert isinstance(fused_ops[0], fused_cls) + assert list(fused_ops[0].basic_ops) == ops + else: + assert fused_ops == ops + + @pytest.mark.parametrize( "unsupported_wrapper", ( @@ -445,10 +582,8 @@ def backward(ctx, grad_output): # pylint: disable=arguments-differ class TestGroupedLinearOp: """Tests for advanced features with grouped linear basic op""" - def test_meta_single_grouped_weight_with_delayed_wgrad(self, monkeypatch) -> None: + def test_meta_single_grouped_weight_with_delayed_wgrad(self) -> None: """A deferred op shell must not access its grouped parent before it is attached.""" - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") - op = te.ops.GroupedLinear( 2, 16, @@ -472,9 +607,8 @@ def test_meta_single_grouped_weight_with_delayed_wgrad(self, monkeypatch) -> Non assert grouped_weight.skip_backward_post_hook - def test_single_grouped_bias_uses_registered_packed_storage(self, monkeypatch) -> None: + def test_single_grouped_bias_uses_registered_packed_storage(self) -> None: """The grouped bias compute view must alias the registered trainable parent.""" - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") op = te.ops.GroupedLinear( 2, 128, @@ -521,13 +655,6 @@ def test_grouped_linear( single_grouped_bias: bool, ) -> None: """Grouped GEMM""" - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and ( - single_grouped_weight or single_grouped_bias - ): - pytest.skip( - "single_grouped_weight/single_grouped_bias requires" - " NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) # Split sizes split_sizes = [split_alignment * i for i in range(group_size)] random.shuffle(split_sizes) @@ -923,13 +1050,6 @@ def test_grouped_linear_cuda_graph_safe( """ # Skip invalid configurations - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and ( - single_grouped_weight - ): - pytest.skip( - "single_grouped_weight/single_grouped_bias requires" - " NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) if quantization is None and quantized_weight: pytest.skip("quantized_weight requires a quantization recipe") if ( @@ -1213,13 +1333,6 @@ def test_grouped_mlp( maybe_skip_quantization(quantization, dims=in_shape, device=device, dtype=dtype) if dtype == torch.bfloat16 and not is_bf16_available(): pytest.skip("BF16 requires SM 8.0+") - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and ( - single_grouped_weight or single_grouped_bias - ): - pytest.skip( - "single_grouped_weight/single_grouped_bias requires" - " NVTE_GROUPED_LINEAR_SINGLE_PARAM=1" - ) if single_grouped_weight and quantization != "mxfp8": pytest.skip("single_grouped_weight is only supported for MXFP8 quantization") if single_grouped_bias and not bias: @@ -1707,8 +1820,6 @@ def test_grouped_mlp_glu_mxfp8_real_cudnn_fusion( assert fused_cls.is_supported() # FC2 bias-gradient accumulation uses an atomic Triton reduction. monkeypatch.setenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1") - if single_grouped_weight: - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") self.test_grouped_mlp( group_size=4, @@ -2182,8 +2293,6 @@ def test_grouped_mlp_single_weight_numerics( ) -> None: """single_grouped_weight=True/False should match exactly for fused MXFP8 grouped MLP.""" - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0": - pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") if activation == "scaled_clamped_qgeglu": @@ -2688,8 +2797,6 @@ def test_grouped_mlp_overwrite_main_grad( that read ``.grad`` don't see stale bytes from the cached dummy). """ - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: - pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") @@ -2821,8 +2928,6 @@ def test_grouped_mlp_cuda_graph_safe_mxfp8( ) -> None: """Grouped MLP forward+backward should be CUDA graph capturable (MXFP8).""" - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_grouped_weight: - pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported(): pytest.skip("MXFP8 fused grouped MLP is not supported on this system") if dtype not in (torch.bfloat16, torch.float16): diff --git a/tests/pytorch/test_grouped_tensor.py b/tests/pytorch/test_grouped_tensor.py index 388b3bb7041..d4eba350a66 100644 --- a/tests/pytorch/test_grouped_tensor.py +++ b/tests/pytorch/test_grouped_tensor.py @@ -4,7 +4,6 @@ """Tests for GroupedTensor class""" -import os from types import SimpleNamespace from typing import List, Optional, Tuple @@ -1944,8 +1943,6 @@ def test_grouped_linear_load_state_dict_multi_to_single_param(self, tmp_path) -> in_features = 64 out_features = 32 dtype = torch.float32 - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0": - pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") src = te.GroupedLinear( num_gemms=num_gemms, in_features=in_features, @@ -1996,8 +1993,6 @@ def test_grouped_linear_load_state_dict_multi_to_single_param(self, tmp_path) -> def test_grouped_linear_load_state_dict_single_to_multi_param(self, tmp_path) -> None: """Load grouped-parameter checkpoint from disk into per-GEMM parameter format.""" - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0": - pytest.skip("single_grouped_weight requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") num_gemms = 3 in_features = 64 out_features = 32 diff --git a/tests/pytorch/test_module_grouped_linear_distributed_weight.py b/tests/pytorch/test_module_grouped_linear_distributed_weight.py index 38c600b3e9a..05fba219a99 100644 --- a/tests/pytorch/test_module_grouped_linear_distributed_weight.py +++ b/tests/pytorch/test_module_grouped_linear_distributed_weight.py @@ -236,7 +236,7 @@ def test_unfused_wgrad_returns_grads_through_finalize(use_grouped_tensor): assert torch.count_nonzero(w.main_grad) == 0, "main_grad was written without fusion" -def test_single_grouped_weight_dispatches(monkeypatch): +def test_single_grouped_weight_dispatches(): """A distributed weight must also work when the group is one packed GroupedTensor. single_grouped_weight requires use_grouped_tensor, and makes the group a single ``weight`` @@ -245,7 +245,6 @@ def test_single_grouped_weight_dispatches(monkeypatch): if not torch.cuda.is_available(): pytest.skip("requires CUDA") - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") torch.manual_seed(0) num_gemms = 2 module = te.GroupedLinear( diff --git a/tests/pytorch/test_sanity.py b/tests/pytorch/test_sanity.py index 7db3002427c..b5fc7df2744 100644 --- a/tests/pytorch/test_sanity.py +++ b/tests/pytorch/test_sanity.py @@ -605,8 +605,6 @@ def test_sanity_grouped_linear( # Small batch size used to catch bug from https://github.com/NVIDIA/TransformerEngine/pull/1527. bs = bs * 16 num_tokens = bs * config.max_seqlen_q * (num_gemms - 1) - if os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0") == "0" and single_param: - pytest.skip("single parameter grouped linear requires NVTE_GROUPED_LINEAR_SINGLE_PARAM=1") skip_unsupported_backward_override("grouped_linear", fp8_recipe, backward_override) if fp8_recipe is not None: fp8_recipe = copy.deepcopy(fp8_recipe) @@ -1358,9 +1356,8 @@ def test_quantized_param_attrs_survive_apply(move): @pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) -def test_grouped_linear_single_param_preserves_high_precision_init(monkeypatch): +def test_grouped_linear_single_param_preserves_high_precision_init(): """Grouped MXFP8 and discrete weights produce identical FP32 master initialization.""" - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") num_gemms = 3 def make_module(single_grouped_weight): @@ -1411,9 +1408,8 @@ def make_module(single_grouped_weight): @pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) -def test_grouped_linear_rejects_partial_high_precision_init(monkeypatch): +def test_grouped_linear_rejects_partial_high_precision_init(): """Packing fails rather than mixing preserved and dequantized initialization.""" - monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") with quantized_model_init( enabled=True, recipe=recipe.MXFP8BlockScaling(), diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index e403630fd0a..1b82b82b8e7 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -42,7 +42,7 @@ init_method_constant, mark_grouped_tensor, requires_grad, - resolve_grouped_linear_single_param_flags, + warn_if_single_grouped_parameters, get_nvtx_range_context, ) from ..distributed import ( @@ -1680,15 +1680,13 @@ class GroupedLinear(TransformerEngineBaseModule): single_grouped_weight : bool, default = False If set to ``True``, grouped weights are stored as a single grouped parameter instead of one parameter per GEMM. - EXPERIMENTAL and subject to change. Gated by the - ``NVTE_GROUPED_LINEAR_SINGLE_PARAM`` environment variable: if the env var - is not set this argument is forced to ``False`` with a warning. + EXPERIMENTAL and subject to change. Requires ``use_grouped_tensor=True`` + and a supported device, dtype, and quantization recipe. single_grouped_bias : bool, default = False If set to ``True``, grouped biases are stored as a single grouped bias instead of one bias per GEMM. - EXPERIMENTAL and subject to change. Gated by the - ``NVTE_GROUPED_LINEAR_SINGLE_PARAM`` environment variable: if the env var - is not set this argument is forced to ``False`` with a warning. + EXPERIMENTAL and subject to change. Requires ``use_grouped_tensor=True`` + and a supported device, dtype, and quantization recipe. use_grouped_tensor : bool or None, default = None Prefer the native GroupedTensor grouped GEMM path. Discrete parameters fall back to split-quantize when the path is unsupported. Single grouped @@ -1763,9 +1761,7 @@ def __init__( f"use_grouped_tensor must be a bool or None, got {type(use_grouped_tensor)}." ) self.use_grouped_tensor = use_grouped_tensor - single_grouped_weight, single_grouped_bias = resolve_grouped_linear_single_param_flags( - single_grouped_weight, single_grouped_bias - ) + warn_if_single_grouped_parameters(single_grouped_weight, single_grouped_bias) self.single_grouped_weight = single_grouped_weight self.single_grouped_bias = single_grouped_bias if self.use_bias and self.single_grouped_weight and not self.single_grouped_bias: @@ -1979,8 +1975,7 @@ def make_grouped_weights(self, defer_init=False) -> None: raise NotImplementedError( "GroupedLinear(single_grouped_weight=True) does not support " f"{quantizer_names} weight quantizers yet. Set " - "single_grouped_weight=False or unset " - "NVTE_GROUPED_LINEAR_SINGLE_PARAM. See #3158." + "single_grouped_weight=False. See #3158." ) recipe = ( diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 9a31fab4e82..b92b136fe8e 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -39,8 +39,8 @@ clear_tensor_data, devices_match, get_device_compute_capability, - resolve_grouped_linear_single_param_flags, round_up_to_nearest_multiple, + warn_if_single_grouped_parameters, ) from .._common import ( get_accumulate_flag_in_param, @@ -169,17 +169,15 @@ class GroupedLinear(BasicOperation): ``main_grad`` instead of accumulating. single_grouped_weight : bool, default = ``False`` Store all expert weights as one ``GroupedTensor`` parameter ``weight``. - EXPERIMENTAL and subject to change. Gated by the - ``NVTE_GROUPED_LINEAR_SINGLE_PARAM`` environment variable: if the env var - is not set this argument is forced to ``False`` with a warning. + EXPERIMENTAL and subject to change. Requires the native grouped-tensor path + for the current device, dtype, and quantization recipe. delay_wgrad_compute : bool, default = ``False`` Whether to delay weight gradient computation single_grouped_bias : bool, default = ``False`` If ``True`` (and ``bias=True``), store all expert biases as one ``GroupedTensor`` parameter named ``bias`` instead of ``bias0``..``bias{N-1}``. - EXPERIMENTAL and subject to change. Gated by the - ``NVTE_GROUPED_LINEAR_SINGLE_PARAM`` environment variable: if the env var - is not set this argument is forced to ``False`` with a warning. + EXPERIMENTAL and subject to change. Requires the native grouped-tensor path + for the current device, dtype, and quantization recipe. scale_bias : bool, default = ``False`` If ``True`` (and ``bias=True``), expects a probability tensor as an additional extra input and adds ``bias * scales`` instead of ``bias`` @@ -220,9 +218,7 @@ def __init__( self.num_groups: int = num_groups self.in_features: int = in_features self.out_features: int = out_features - single_grouped_weight, single_grouped_bias = resolve_grouped_linear_single_param_flags( - single_grouped_weight, single_grouped_bias - ) + warn_if_single_grouped_parameters(single_grouped_weight, single_grouped_bias) self.single_grouped_weight: bool = single_grouped_weight self.single_grouped_bias: bool = single_grouped_bias self.use_bias: bool = bias diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index 730c9d26161..370772f719f 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -479,9 +479,7 @@ def _pack_grouped_linear_bias_for_cudnn(linear_op: GroupedLinear) -> Optional[to @functools.lru_cache(maxsize=1) def _grouped_gemm_dsrelu_backward_supported() -> bool: """Whether the cuDNN FE grouped GEMM dSReLU backward wrapper is available.""" - if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - return False - if get_device_compute_capability()[0] != 10: + if not torch.cuda.is_available() or get_device_compute_capability()[0] != 10: return False if not _cudnn_frontend_supports_grouped_gemm_srelu(): return False @@ -932,6 +930,29 @@ def validate_grouped_mlp_dims(fc1, activation_op, fc2) -> None: ) +def _is_grouped_mlp_fusion_candidate( + ops: list[FusibleOperation], + recipe: Optional[Recipe], + activation_op_types: tuple[type[FusibleOperation]], +) -> bool: + """Check the recipe and operation pattern before probing CUDA or cuDNN.""" + if len(ops) < 3 or recipe is None or not (recipe.mxfp8() or recipe.nvfp4()): + return False + # NVFP4 graph-safe grouped quantize currently requires RHT. + if recipe.nvfp4() and recipe.disable_rht: + return False + # The fused MXFP8 backward reinterprets grad-output storage as E4M3. It + # cannot consume E5M2 gradients from Format.HYBRID. NVFP4 has separate formats. + if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: + return False + return any( + isinstance(fc1, GroupedLinear) + and isinstance(activation, activation_op_types) + and isinstance(fc2, GroupedLinear) + for fc1, activation, fc2 in zip(ops, ops[1:], ops[2:]) + ) + + def fuse_grouped_mlp_ops( ops: list[FusibleOperation], *, @@ -956,18 +977,9 @@ def fuse_grouped_mlp_ops( list of FusibleOperation Updated operations with matched triples replaced by fused ops. """ - if not fused_op_cls.is_supported(): + if not _is_grouped_mlp_fusion_candidate(ops, recipe, activation_op_types): return ops - if recipe is None or not (recipe.mxfp8() or recipe.nvfp4()): - return ops - # NVFP4 fused grouped MLP uses graph-safe grouped quantize, which currently requires RHT. - if recipe.nvfp4() and recipe.disable_rht: - return ops - # The fused MXFP8 backward reinterprets the grad output's storage as E4M3, so an E5M2 - # backward format would have its gradients misread rather than converted. This declines - # MXFP8 with Format.HYBRID. fp8_format does not describe NVFP4 gradients, so NVFP4 is - # excluded from the check rather than relying on its value. - if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: + if not fused_op_cls.is_supported(): return ops out = [] @@ -1068,9 +1080,7 @@ def grouped_gemm_wgrad_kernel(cls) -> Optional[Callable]: @functools.lru_cache(maxsize=None) def is_supported(cls) -> bool: """Whether this fused operation is supported on the current system.""" - if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - return False - if get_device_compute_capability()[0] != 10: + if not torch.cuda.is_available() or get_device_compute_capability()[0] != 10: return False if not _cudnn_frontend_version_supported(): return False @@ -2825,6 +2835,15 @@ def fuse_glu_ops( ) -> list[FusibleOperation]: """Apply joint GroupedLinear + scaled GLU + GroupedLinear fusion.""" + # Registered at import time; defer CUDA and optional dependency checks until + # a block-scaled pipeline can actually use this fusion. + if not _is_grouped_mlp_fusion_candidate( + ops, recipe, (ScaledSwiGLU, ScaledClampedQGeGLU, ScaledSiTUGLU) + ): + return ops + if not torch.cuda.is_available(): + return ops + # Determine supported activations activation_op_types = [] device_arch = get_device_compute_capability() @@ -2856,6 +2875,11 @@ def fuse_unary_activation_ops( ) -> list[FusibleOperation]: """Apply joint GroupedLinear + scaled unary activation + GroupedLinear fusion.""" + if not _is_grouped_mlp_fusion_candidate(ops, recipe, (ScaledSReLU, ScaledTanhSReLU)): + return ops + if not GroupedMLP_CuTeGEMMUnary.is_supported(): + return ops + # Determine supported activations activation_op_types = [ScaledSReLU] if _cudnn_frontend_supports_grouped_gemm_srelu_tanh(): @@ -2869,8 +2893,7 @@ def fuse_unary_activation_ops( ) -# Register joint fusions if available. -if GroupedMLP_CuTeGEMMGLU.is_supported(): - register_forward_backward_fusion(fuse_glu_ops, prepend=True) -if GroupedMLP_CuTeGEMMUnary.is_supported(): - register_forward_backward_fusion(fuse_unary_activation_ops, prepend=True) +# Register without probing CUDA or importing optional cuDNN kernels. Capability +# checks run when the fuser encounters a supported quantization recipe. +register_forward_backward_fusion(fuse_glu_ops, prepend=True) +register_forward_backward_fusion(fuse_unary_activation_ops, prepend=True) diff --git a/transformer_engine/pytorch/utils.py b/transformer_engine/pytorch/utils.py index ca77f7b8bb3..ffeba7b738c 100644 --- a/transformer_engine/pytorch/utils.py +++ b/transformer_engine/pytorch/utils.py @@ -260,25 +260,13 @@ def interleave_glu_tensor(tensor: torch.Tensor, interleave_size: int) -> torch.T return x.reshape(shape) -def resolve_grouped_linear_single_param_flags( +def warn_if_single_grouped_parameters( single_grouped_weight: bool, single_grouped_bias: bool, -) -> Tuple[bool, bool]: - """Gate ``single_grouped_weight`` / ``single_grouped_bias`` on ``NVTE_GROUPED_LINEAR_SINGLE_PARAM``.""" +) -> None: + """Warn when a caller selects the experimental single grouped parameter layout.""" if not (single_grouped_weight or single_grouped_bias): - return single_grouped_weight, single_grouped_bias - - env_enabled = int(os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "0")) > 0 - if not env_enabled: - warnings.warn( - f"GroupedLinear was constructed with single_grouped_weight={single_grouped_weight} " - f"and single_grouped_bias={single_grouped_bias}, but the " - "NVTE_GROUPED_LINEAR_SINGLE_PARAM environment variable is not set. " - "Disabling single grouped weight/bias and falling back to per-expert parameters.", - UserWarning, - stacklevel=3, - ) - return False, False + return warnings.warn( "GroupedLinear is using single_grouped_weight/single_grouped_bias. " @@ -287,7 +275,6 @@ def resolve_grouped_linear_single_param_flags( UserWarning, stacklevel=3, ) - return single_grouped_weight, single_grouped_bias def attention_mask_func( From 91032dbd43b2dcbb3db004c4fbcc0ea9755a4b47 Mon Sep 17 00:00:00 2001 From: zhongboz Date: Thu, 1 Oct 2026 16:08:43 -0700 Subject: [PATCH 2/3] resolve comments Signed-off-by: zhongboz --- docs/envvars.rst | 29 ---- docs/examples/te_mixtral/te_mixtral_mxfp8.py | 13 +- ...torial_accelerate_hf_mixtral_with_te.ipynb | 12 +- docs/examples/te_mixtral/utils.py | 64 -------- tests/pytorch/distributed/run_ep.py | 60 ++++---- tests/pytorch/test_grouped_mlp.py | 137 ------------------ .../pytorch/ops/fused/grouped_mlp.py | 12 +- 7 files changed, 35 insertions(+), 292 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 61d80dd0e3e..46b70bbe46a 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -134,35 +134,6 @@ Runtime Environment Variables These environment variables control the behavior of Transformer Engine during execution. -Deprecated Grouped Linear and MLP Controls -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ - -.. envvar:: NVTE_GROUPED_LINEAR_SINGLE_PARAM - - :Status: **DEPRECATED**. No longer read. - :Historical default: ``0`` (before deprecation). - :Description: Previously permitted grouped-linear weights and biases to use - a single grouped parameter. This variable is now ignored, - including when set to ``0`` or ``1``. - :Migration: Remove this variable from launch scripts. Select the parameter - layout with the ``single_grouped_weight`` and - ``single_grouped_bias`` constructor arguments; both still - default to ``False``. - -.. envvar:: NVTE_CUTEDSL_FUSED_GROUPED_MLP - - :Status: **DEPRECATED**. No longer read. - :Historical default: ``0`` (before deprecation). - :Description: Previously enabled CuTeDSL grouped MLP fusion in the PyTorch - operation fuser. This variable is now ignored, including when - set to ``0`` or ``1``; setting it to ``0`` no longer disables - fusion. - :Migration: Remove this variable from launch scripts. The operation fuser - automatically selects grouped MLP fusion for compatible - operation sequences when the GPU, recipe, and installed - dependencies support it. No environment variable needs to be - set before importing TE to enable this fusion. - General ^^^^^^^ diff --git a/docs/examples/te_mixtral/te_mixtral_mxfp8.py b/docs/examples/te_mixtral/te_mixtral_mxfp8.py index c35effeeac9..b698c78ae06 100644 --- a/docs/examples/te_mixtral/te_mixtral_mxfp8.py +++ b/docs/examples/te_mixtral/te_mixtral_mxfp8.py @@ -7,14 +7,13 @@ MoE FFN is a TE ``Sequential`` of three fusible ops — ``GroupedLinear`` (gate_up), ``ScaledSwiGLU(glu_interleave_size=32)``, ``GroupedLinear`` (down) — that the OperationFuser collapses into the fused -``ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8`` and backward kernels under -MXFP8. HF gate (``w1``) and up (``w3``) weights are row-interleaved in +``GroupedMLP_CuTeGEMMGLU`` forward and backward kernels under MXFP8. +HF gate (``w1``) and up (``w3``) weights are row-interleaved in blocks of 32 to match the GLU interleaved layout that fused kernel reads. -Supported grouped-MLP sequences are fused automatically. -``utils._enable_fused_mxfp8_grouped_mlp()`` provides legacy SM-version / -cudnn-frontend signature compatibility patches. Requires -``nvidia-cudnn-frontend >= 1.23.0`` and SM>=10 (Blackwell B100/B200/B300+). +This example targets the current TE source tree. Supported grouped-MLP +sequences are fused automatically when the GPU, recipe, and installed +dependencies support it. """ from __future__ import annotations @@ -163,7 +162,7 @@ def _init_method(x: torch.Tensor) -> None: device=device, ) # Wrap as TE Sequential to enable forward/backward op fusion - # (ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8 / dswiglu). + # (GroupedMLP_CuTeGEMMGLU). object.__setattr__( self, "_experts_ffn_op", diff --git a/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb b/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb index e757085862b..c0dc0372d0e 100644 --- a/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb +++ b/docs/examples/te_mixtral/tutorial_accelerate_hf_mixtral_with_te.ipynb @@ -22,7 +22,7 @@ "\n", "Mixtral-8x7B has 8 experts and roughly 47B total parameters. In `BF16` the model weights alone consume ~93 GB, and full `AdamW` fine-tuning needs ~370 GB. This tutorial is tested on 8x B300 GPUs with `Expert Parallelism (EP) = 2` and `Data Parallelism (DP) = 4`, so the experts are divided across 2 GPUs and there are 4 replicas. The container used is [pytorch-26.04-py3](https://catalog.ngc.nvidia.com/orgs/nvidia/containers/pytorch?version=26.04-py3). A sequence length of 8192 and a global batch size of 48 are used across the experiments.\n", "\n", - "Install the required Python packages using the following command in a terminal:" + "This example targets the current Transformer Engine source tree. Install Transformer Engine from this checkout, then install the remaining required Python packages using the following command in a terminal:" ] }, { @@ -346,15 +346,7 @@ "experts_ffn = Sequential(GroupedLinear(gate_up), ScaledSwiGLU(), GroupedLinear(down))\n", "```\n", "\n", - "TE's `Sequential` scans the ops and, if the pattern matches, replaces the `GroupedLinear -> ScaledSwiGLU -> GroupedLinear` pattern with a fused operation object: `ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8` for forward and a matching fused backward op. It reduces framework overhead, fuses the SwiGLU/probability-scaling work into the grouped MLP path, and avoids some intermediate materialization.\n", - "\n", - "
\n", - "\n", - "Note\n", - "\n", - "TE automatically selects grouped-MLP fusion when the operation sequence, GPU, recipe, and installed dependencies support it. No environment variable needs to be set before importing TE to enable this fusion.\n", - "\n", - "
" + "TE's `Sequential` scans the ops and, if the pattern matches, replaces the `GroupedLinear -> ScaledSwiGLU -> GroupedLinear` pattern with a fused `GroupedMLP_CuTeGEMMGLU` operation for forward and backward. It reduces framework overhead, fuses the SwiGLU/probability-scaling work into the grouped MLP path, and avoids some intermediate materialization." ] }, { diff --git a/docs/examples/te_mixtral/utils.py b/docs/examples/te_mixtral/utils.py index 8701afcb6b5..d5bd6e210df 100644 --- a/docs/examples/te_mixtral/utils.py +++ b/docs/examples/te_mixtral/utils.py @@ -148,69 +148,6 @@ def init_baseline_model(hyperparams: HyperParameters): return model -def _enable_fused_mxfp8_grouped_mlp() -> None: - """Improvement 3: adapt legacy fused grouped-MLP kernels without recompiling. - - Current TE selects supported grouped-MLP fusions automatically. The - compatibility patches below target older TE implementations. - - We (a) relax the SM-version check from ``!= 10`` to ``>= 10`` so - SM>=11 successors of B300 fire the kernel, and (b) wrap the cudnn-frontend - grouped-GEMM wrappers so the installed TE's ``c_dtype`` kwarg (dropped by - cudnn-frontend 1.23.0) is silently filtered out. - """ - import inspect - import cudnn # type: ignore - from transformer_engine.pytorch.ops.fused import forward_grouped_mlp as _fwd_mod - from transformer_engine.pytorch.ops.fused import backward_grouped_mlp as _bwd_mod - from transformer_engine.pytorch.utils import get_device_compute_capability - - def _make_is_supported(kernel_method_names): - def _is_supported(cls) -> bool: - if get_device_compute_capability()[0] < 10: - return False - try: - for method_name in kernel_method_names: - getattr(cls, method_name)() - except ImportError: - return False - return True - - return _is_supported - - def _make_compat_kernel(real_callable): - accepted = set(inspect.signature(real_callable).parameters) - - def _compat(**kwargs): - for k in list(kwargs): - if k not in accepted: - kwargs.pop(k) - return real_callable(**kwargs) - - return _compat - - def _patch_kernel_method(cls, method_name, wrapper_name): - compat = _make_compat_kernel(getattr(cudnn, wrapper_name)) - - def _kernel_classmethod(_cls): - return compat - - setattr(cls, method_name, classmethod(_kernel_classmethod)) - - fwd_cls = _fwd_mod.ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8 - bwd_cls = _bwd_mod.BackwardGroupedMLP_CuTeGEMMDSwiGLU_MXFP8 - fwd_cls.is_supported = classmethod( - _make_is_supported(("grouped_gemm_glu_kernel", "grouped_gemm_quant_kernel")) - ) - bwd_cls.is_supported = classmethod( - _make_is_supported(("grouped_gemm_dglu_kernel", "grouped_gemm_quant_kernel")) - ) - _patch_kernel_method(fwd_cls, "grouped_gemm_glu_kernel", "grouped_gemm_glu_wrapper_sm100") - _patch_kernel_method(fwd_cls, "grouped_gemm_quant_kernel", "grouped_gemm_quant_wrapper_sm100") - _patch_kernel_method(bwd_cls, "grouped_gemm_dglu_kernel", "grouped_gemm_dglu_wrapper_sm100") - _patch_kernel_method(bwd_cls, "grouped_gemm_quant_kernel", "grouped_gemm_quant_wrapper_sm100") - - def init_te_mixtral_model(hyperparams: HyperParameters): """Load Mixtral with TE-optimised MoE blocks.""" ensure_model_is_downloaded(hyperparams) @@ -220,7 +157,6 @@ def init_te_mixtral_model(hyperparams: HyperParameters): if hyperparams.model_impl == "te_mixtral_mxfp8": if hyperparams.mixed_precision != "mxfp8": raise ValueError("model_impl='te_mixtral_mxfp8' requires mixed_precision='mxfp8'.") - _enable_fused_mxfp8_grouped_mlp() from te_mixtral_mxfp8 import TEMixtralMXFP8ForCausalLM as ForCausalLM from te_mixtral_mxfp8 import replace_params else: diff --git a/tests/pytorch/distributed/run_ep.py b/tests/pytorch/distributed/run_ep.py index 214f1abc3e0..68ecaf0ba48 100644 --- a/tests/pytorch/distributed/run_ep.py +++ b/tests/pytorch/distributed/run_ep.py @@ -1282,40 +1282,32 @@ def _make_megamoe_model( if recipe is not None else nullcontext() ) - previous_single_param = os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM") - os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] = "1" - try: - with init_ctx: - fc1 = te_ops.GroupedLinear( - NUM_LOCAL_EXPERTS, - HIDDEN_DIM, - 2 * 256, - bias=False, - device=self.cfg.device, - dtype=torch.bfloat16, - single_grouped_weight=True, - accumulate_into_main_grad=accumulate_into_main_grad, - delay_wgrad_compute=delay_wgrad_compute, - ) - activation = te_ops.ScaledSwiGLU( - glu_interleave_size=glu_interleave_size, - ) - fc2 = te_ops.GroupedLinear( - NUM_LOCAL_EXPERTS, - 256, - HIDDEN_DIM, - bias=False, - device=self.cfg.device, - dtype=torch.bfloat16, - single_grouped_weight=True, - accumulate_into_main_grad=accumulate_into_main_grad, - delay_wgrad_compute=delay_wgrad_compute, - ) - finally: - if previous_single_param is None: - del os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] - else: - os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] = previous_single_param + with init_ctx: + fc1 = te_ops.GroupedLinear( + NUM_LOCAL_EXPERTS, + HIDDEN_DIM, + 2 * 256, + bias=False, + device=self.cfg.device, + dtype=torch.bfloat16, + single_grouped_weight=True, + accumulate_into_main_grad=accumulate_into_main_grad, + delay_wgrad_compute=delay_wgrad_compute, + ) + activation = te_ops.ScaledSwiGLU( + glu_interleave_size=glu_interleave_size, + ) + fc2 = te_ops.GroupedLinear( + NUM_LOCAL_EXPERTS, + 256, + HIDDEN_DIM, + bias=False, + device=self.cfg.device, + dtype=torch.bfloat16, + single_grouped_weight=True, + accumulate_into_main_grad=accumulate_into_main_grad, + delay_wgrad_compute=delay_wgrad_compute, + ) combine = te_ops.MoeCombine(config, buffer) dispatch.set_extra_output_channel(0, "tokens_per_expert", output_to_caller=False) dispatch.set_extra_output_channel(1, "routing_weights", output_to_caller=False) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index aabe2aed34a..729a57761f3 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -7,7 +7,6 @@ from collections.abc import Iterable import contextlib import functools -import importlib.util import os import math import random @@ -139,142 +138,6 @@ def _reset_rng_states_per_test(): yield -@pytest.fixture -def isolated_grouped_mlp_module(monkeypatch): - """Import the module with a private fusion registry and fresh capability caches.""" - from transformer_engine.pytorch import utils as te_utils - from transformer_engine.pytorch.ops.fuser import OperationFuser - - monkeypatch.delenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP", raising=False) - - def reject_cuda_query(*args, **kwargs): - raise AssertionError("Importing grouped MLP must not query or initialize CUDA") - - spec = importlib.util.spec_from_file_location( - f"{grouped_mlp_module.__package__}._test_grouped_mlp", grouped_mlp_module.__file__ - ) - module = importlib.util.module_from_spec(spec) - with monkeypatch.context() as import_patch: - import_patch.setattr(te_utils, "get_device_compute_capability", reject_cuda_query) - import_patch.setattr(torch.cuda, "is_available", reject_cuda_query) - import_patch.setattr(torch.cuda, "_lazy_init", reject_cuda_query) - import_patch.setattr(OperationFuser, "forward_backward_fusion_functions", []) - spec.loader.exec_module(module) - registered_fusions = tuple(OperationFuser.forward_backward_fusion_functions) - return module, registered_fusions - - -def test_grouped_mlp_registers_without_cuda_query(isolated_grouped_mlp_module) -> None: - """Both default fusion callbacks must be registered before capability is known.""" - module, registered_fusions = isolated_grouped_mlp_module - assert set(registered_fusions) == {module.fuse_glu_ops, module.fuse_unary_activation_ops} - - -@pytest.mark.parametrize( - "pipeline,quantization", - ( - ("unrelated", None), - ("unrelated", "mxfp8"), - ("unrelated", "nvfp4"), - ("glu", "mxfp8_hybrid"), - ("unary", "mxfp8_hybrid"), - ("glu", "nvfp4_no_rht"), - ("unary", "nvfp4_no_rht"), - ), -) -def test_grouped_mlp_ineligible_pipeline_skips_capability_queries( - monkeypatch, isolated_grouped_mlp_module, pipeline: str, quantization: Optional[str] -) -> None: - """Unrelated ops and unsupported recipes must not trigger device or cuDNN probes.""" - from transformer_engine.common.recipe import Format, MXFP8BlockScaling, NVFP4BlockScaling - - module, registered_fusions = isolated_grouped_mlp_module - - def reject_capability_query(*args, **kwargs): - raise AssertionError("Ineligible pipelines must bypass grouped MLP capability checks") - - monkeypatch.setattr(torch.cuda, "is_available", reject_capability_query) - monkeypatch.setattr(module, "get_pkg_version", reject_capability_query) - recipe = None - if quantization in ("mxfp8", "mxfp8_hybrid"): - fp8_format = Format.HYBRID if quantization == "mxfp8_hybrid" else Format.E4M3 - recipe = MXFP8BlockScaling(fp8_format=fp8_format) - elif quantization in ("nvfp4", "nvfp4_no_rht"): - recipe = NVFP4BlockScaling(disable_rht=quantization == "nvfp4_no_rht") - - if pipeline == "unrelated": - ops = [te.ops.ReLU() for _ in range(3)] - else: - is_glu = pipeline == "glu" - kwargs = {"bias": False, "device": "meta", "dtype": torch.bfloat16} - ops = [ - te.ops.GroupedLinear(2, 64, 128 if is_glu else 64, **kwargs), - te.ops.ScaledSwiGLU(glu_interleave_size=32) if is_glu else te.ops.ScaledSReLU(), - te.ops.GroupedLinear(2, 64, 64, **kwargs), - ] - for fuse in registered_fusions: - assert fuse(ops, recipe=recipe) is ops - - -@pytest.mark.parametrize("activation", ("glu", "unary")) -@pytest.mark.parametrize( - "capability", ("supported", "cuda", "architecture", "frontend", "wrapper", "recipe", "shape") -) -def test_grouped_mlp_default_fusion_capabilities( - monkeypatch, isolated_grouped_mlp_module, activation: str, capability: str -) -> None: - """Default fusion still respects hardware, dependency, recipe, and shape support.""" - from transformer_engine.common.recipe import Format, MXFP8BlockScaling - - module, registered_fusions = isolated_grouped_mlp_module - monkeypatch.setattr(torch.cuda, "is_available", lambda: capability != "cuda") - monkeypatch.setattr( - module, - "get_device_compute_capability", - lambda: (9, 0) if capability == "architecture" else (10, 0), - ) - monkeypatch.setattr( - module, "get_pkg_version", lambda _: "1.22.0" if capability == "frontend" else "1.24.0" - ) - - def unused_kernel(): - raise AssertionError("Fusion planning must not execute a kernel") - - fake_cudnn = types.ModuleType("cudnn") - for name in ("glu", "dglu", "srelu", "dsrelu", "quant", "wgrad"): - setattr(fake_cudnn, f"grouped_gemm_{name}_wrapper_sm100", unused_kernel) - if capability == "wrapper": - del fake_cudnn.grouped_gemm_quant_wrapper_sm100 - monkeypatch.setitem(sys.modules, "cudnn", fake_cudnn) - - hidden_size = 96 if capability == "shape" else 64 - if activation == "glu": - activation_op = te.ops.ScaledSwiGLU(glu_interleave_size=32) - fc1_out_features = 128 - fuse = module.fuse_glu_ops - fused_cls = module.GroupedMLP_CuTeGEMMGLU - else: - activation_op = te.ops.ScaledSReLU() - fc1_out_features = 64 - fuse = module.fuse_unary_activation_ops - fused_cls = module.GroupedMLP_CuTeGEMMUnary - assert fuse in registered_fusions - kwargs = {"bias": False, "device": "meta", "dtype": torch.bfloat16} - ops = [ - te.ops.GroupedLinear(2, hidden_size, fc1_out_features, **kwargs), - activation_op, - te.ops.GroupedLinear(2, 64, hidden_size, **kwargs), - ] - recipe = None if capability == "recipe" else MXFP8BlockScaling(fp8_format=Format.E4M3) - fused_ops = fuse(ops, recipe=recipe) - if capability == "supported": - assert len(fused_ops) == 1 - assert isinstance(fused_ops[0], fused_cls) - assert list(fused_ops[0].basic_ops) == ops - else: - assert fused_ops == ops - - @pytest.mark.parametrize( "unsupported_wrapper", ( diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index 370772f719f..673d6aa46fc 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -935,7 +935,7 @@ def _is_grouped_mlp_fusion_candidate( recipe: Optional[Recipe], activation_op_types: tuple[type[FusibleOperation]], ) -> bool: - """Check the recipe and operation pattern before probing CUDA or cuDNN.""" + """Check whether the recipe and operation pattern support grouped MLP fusion.""" if len(ops) < 3 or recipe is None or not (recipe.mxfp8() or recipe.nvfp4()): return False # NVFP4 graph-safe grouped quantize currently requires RHT. @@ -2835,12 +2835,6 @@ def fuse_glu_ops( ) -> list[FusibleOperation]: """Apply joint GroupedLinear + scaled GLU + GroupedLinear fusion.""" - # Registered at import time; defer CUDA and optional dependency checks until - # a block-scaled pipeline can actually use this fusion. - if not _is_grouped_mlp_fusion_candidate( - ops, recipe, (ScaledSwiGLU, ScaledClampedQGeGLU, ScaledSiTUGLU) - ): - return ops if not torch.cuda.is_available(): return ops @@ -2875,8 +2869,6 @@ def fuse_unary_activation_ops( ) -> list[FusibleOperation]: """Apply joint GroupedLinear + scaled unary activation + GroupedLinear fusion.""" - if not _is_grouped_mlp_fusion_candidate(ops, recipe, (ScaledSReLU, ScaledTanhSReLU)): - return ops if not GroupedMLP_CuTeGEMMUnary.is_supported(): return ops @@ -2893,7 +2885,5 @@ def fuse_unary_activation_ops( ) -# Register without probing CUDA or importing optional cuDNN kernels. Capability -# checks run when the fuser encounters a supported quantization recipe. register_forward_backward_fusion(fuse_glu_ops, prepend=True) register_forward_backward_fusion(fuse_unary_activation_ops, prepend=True) From 5974f7b2e5119aaf4948f1d6dec406f8424c96d6 Mon Sep 17 00:00:00 2001 From: Ravi Ghadia Date: Mon, 5 Oct 2026 16:16:12 -0700 Subject: [PATCH 3/3] feat(pytorch): add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK With #3584, TE fuses GroupedLinear + activation + GroupedLinear into the CuTeDSL grouped MLP ops whenever possible and silently falls back to unfused ops otherwise. A fallback can cost a large share of end-to-end performance and memory, and its cause, for example a cuDNN frontend and cutlass-dsl mismatch that makes the kernels fail to import, is invisible. When NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK=1, warn once per fused op, activation, and reason whenever a grouped MLP pattern is left unfused: CUDA unavailable, unsupported GPU, cuDNN frontend too old or kernels failing to import, or an unsupported recipe, activation, or dimensions. Nothing changes when the variable is unset. Co-authored-by: Claude Signed-off-by: Ravi Ghadia --- docs/envvars.rst | 6 + tests/pytorch/test_grouped_mlp.py | 65 ++++++ .../pytorch/ops/fused/grouped_mlp.py | 205 +++++++++++++++--- 3 files changed, 244 insertions(+), 32 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 46b70bbe46a..0c23d9619bf 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -159,6 +159,12 @@ General :Default: ``0`` :Description: Warn when TE falls back to the CUDA C++ kernels instead of dispatching to available CuTeDSL kernels. Useful to verify if the CuTeDSL path is taken. +.. envvar:: NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK + + :Type: ``int`` (0 or 1) + :Default: ``0`` + :Description: Warn when a ``GroupedLinear`` + activation + ``GroupedLinear`` pattern in ``transformer_engine.pytorch.ops`` is not fused into the CuTeDSL grouped MLP kernels and falls back to unfused ops. The warning states the reason, for example an unsupported GPU, recipe, activation, or dimensions, or cuDNN frontend kernels that fail to import. Useful to verify if the fused grouped MLP path is taken. + .. envvar:: NVTE_GROUPED_TENSOR_HANDLE_POOL_SIZE_MB :Type: ``int`` (positive integer) diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index 4973eb39782..03de7005348 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -13,6 +13,7 @@ import sys import types from typing import Optional +import warnings import pytest @@ -175,6 +176,70 @@ def _without_situglu(): _cudnn_frontend_supports_grouped_gemm_situglu.cache_clear() +@pytest.mark.parametrize( + "reason", + ("unsupported_system", "unsupported_recipe", "unsupported_dims", "unsupported_device"), +) +def test_grouped_mlp_fallback_warning(monkeypatch, reason: str) -> None: + """Warn with the cause of a grouped MLP fallback, only when requested.""" + from transformer_engine.common.recipe import Format, MXFP8BlockScaling + + fused_op_cls = grouped_mlp_module.GroupedMLP_CuTeGEMMGLU + monkeypatch.setattr( + fused_op_cls, "is_supported", classmethod(lambda cls: reason != "unsupported_system") + ) + monkeypatch.setattr( + fused_op_cls, "_unsupported_reason", classmethod(lambda cls: "broken cutlass-dsl") + ) + if reason == "unsupported_device": + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(grouped_mlp_module, "get_device_compute_capability", lambda: (9, 0)) + + in_features = 96 if reason == "unsupported_dims" else 128 + fc1 = te.ops.GroupedLinear(2, in_features, 256, bias=False, device="meta") + activation = te.ops.ScaledSwiGLU(glu_interleave_size=32) + fc2 = te.ops.GroupedLinear(2, 128, in_features, bias=False, device="meta") + recipe = None if reason == "unsupported_recipe" else MXFP8BlockScaling(fp8_format=Format.E4M3) + expected = { + "unsupported_system": "broken cutlass-dsl", + "unsupported_recipe": "requires an MXFP8 or NVFP4 recipe", + "unsupported_dims": "Unsupported dims for FC1", + "unsupported_device": "unsupported compute capability 9.0", + }[reason] + + def fuse(ops): + if reason == "unsupported_device": + return grouped_mlp_module.fuse_glu_ops(ops, recipe=recipe) + return grouped_mlp_module.fuse_grouped_mlp_ops( + ops, + recipe=recipe, + fused_op_cls=fused_op_cls, + activation_op_types=(te.ops.ScaledSwiGLU,), + ) + + ops = [fc1, activation, fc2] + grouped_mlp_module._warn_grouped_mlp_fallback.cache_clear() + try: + monkeypatch.delenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK", raising=False) + with warnings.catch_warnings(): + warnings.simplefilter("error") + assert fuse(ops) == ops + + monkeypatch.setenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK", "1") + with pytest.warns( + UserWarning, match=f"ScaledSwiGLU .* GroupedMLP_CuTeGEMMGLU .*{expected}" + ): + assert fuse(ops) == ops + + # Ops without a grouped MLP pattern never warn + grouped_mlp_module._warn_grouped_mlp_fallback.cache_clear() + with warnings.catch_warnings(): + warnings.simplefilter("error") + assert fuse([fc1, fc2]) == [fc1, fc2] + finally: + grouped_mlp_module._warn_grouped_mlp_fallback.cache_clear() + + def _clear_grouped_glu_kernel_caches() -> None: """Clear cached cuDNN wrapper lookups after tests monkeypatch them.""" fused_cls = te.ops.fused.GroupedMLP_CuTeGEMMGLU diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index a3c19432b16..61554c1a64c 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -10,6 +10,7 @@ import functools import inspect import os +import warnings from importlib.metadata import PackageNotFoundError, version as get_pkg_version from typing import Any, Literal, Optional @@ -74,6 +75,14 @@ def _cudnn_frontend_version_at_least(min_version: str) -> bool: return False +def _cudnn_frontend_version_str() -> str: + """Installed cuDNN frontend package version, for diagnostics.""" + try: + return get_pkg_version("nvidia-cudnn-frontend") + except PackageNotFoundError: + return "not installed" + + def _cudnn_frontend_version_supported() -> bool: """Check cuDNN frontend is at least 1.23.0. @@ -959,9 +968,13 @@ def _compute_grad_params( return w_list + bias_list +_GLU_ACTIVATION_OP_TYPES = (ScaledSwiGLU, ScaledSiTUGLU, ScaledClampedQGeGLU) +_UNARY_ACTIVATION_OP_TYPES = (ScaledSReLU, ScaledTanhSReLU) + + def is_glu_activation(activation_op) -> bool: """Whether an activation consumes a GLU-style doubled input.""" - return isinstance(activation_op, (ScaledSwiGLU, ScaledSiTUGLU, ScaledClampedQGeGLU)) + return isinstance(activation_op, _GLU_ACTIVATION_OP_TYPES) def validate_grouped_mlp_dims(fc1, activation_op, fc2) -> None: @@ -978,7 +991,7 @@ def validate_grouped_mlp_dims(fc1, activation_op, fc2) -> None: ) if is_glu_activation(activation_op): expected_fc1_out_features = 2 * fc2.in_features - elif isinstance(activation_op, (ScaledSReLU, ScaledTanhSReLU)): + elif isinstance(activation_op, _UNARY_ACTIVATION_OP_TYPES): expected_fc1_out_features = fc2.in_features else: raise TypeError(f"Unsupported grouped MLP activation ({activation_op.__class__.__name__}).") @@ -997,27 +1010,85 @@ def validate_grouped_mlp_dims(fc1, activation_op, fc2) -> None: ) +def _grouped_mlp_fallback_warnings_enabled() -> bool: + """Whether to warn when a grouped MLP pattern is left unfused. + + Opt-in via ``NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK``. + """ + return int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK", "0")) > 0 + + +@functools.lru_cache(maxsize=None) +def _warn_grouped_mlp_fallback(fused_op_name: str, activation_name: str, reason: str) -> None: + """Warn, once per fused op, activation, and reason, that a grouped MLP is left unfused.""" + warnings.warn( + f"Not fusing GroupedLinear + {activation_name} + GroupedLinear into {fused_op_name}" + f" ({reason}). Falling back to unfused ops, which can be significantly slower and" + " use more memory.", + UserWarning, + ) + + +def _find_grouped_mlp_activation( + ops: Sequence[FusibleOperation], + activation_op_types: tuple[type[FusibleOperation], ...], +) -> Optional[FusibleOperation]: + """Activation of the first GroupedLinear + activation + GroupedLinear pattern, if any.""" + for fc1, activation, fc2 in zip(ops, ops[1:], ops[2:]): + if ( + isinstance(fc1, GroupedLinear) + and isinstance(activation, activation_op_types) + and isinstance(fc2, GroupedLinear) + ): + return activation + return None + + +def _maybe_warn_grouped_mlp_fallback( + ops: Sequence[FusibleOperation], + activation_op_types: tuple[type[FusibleOperation], ...], + fused_op_cls: type[_GroupedMLP_CuTeGEMMBase], + reason: Callable[[], Optional[str]], +) -> None: + """Warn if requested and ``ops`` contain a grouped MLP with one of ``activation_op_types``. + + ``reason`` is only called when a warning is emitted. + """ + if not _grouped_mlp_fallback_warnings_enabled(): + return + activation = _find_grouped_mlp_activation(ops, activation_op_types) + if activation is None: + return + message = reason() + if message is not None: + _warn_grouped_mlp_fallback(fused_op_cls.__name__, type(activation).__name__, message) + + +def _grouped_mlp_recipe_unsupported_reason(recipe: Optional[Recipe]) -> Optional[str]: + """Why the recipe cannot use fused grouped MLP kernels, or ``None`` if it can.""" + if recipe is None: + return "requires an MXFP8 or NVFP4 recipe, but quantization is disabled" + if not (recipe.mxfp8() or recipe.nvfp4()): + return f"requires an MXFP8 or NVFP4 recipe, got {type(recipe).__name__}" + # NVFP4 graph-safe grouped quantize currently requires RHT. + if recipe.nvfp4() and recipe.disable_rht: + return "NVFP4 requires RHT, but the recipe sets disable_rht=True" + # The fused MXFP8 backward reinterprets grad-output storage as E4M3. It + # cannot consume E5M2 gradients from Format.HYBRID. NVFP4 has separate formats. + if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: + return f"MXFP8 requires E4M3 gradients, but the recipe uses {recipe.fp8_format}" + return None + + def _is_grouped_mlp_fusion_candidate( ops: list[FusibleOperation], recipe: Optional[Recipe], activation_op_types: tuple[type[FusibleOperation]], ) -> bool: """Check whether the recipe and operation pattern support grouped MLP fusion.""" - if len(ops) < 3 or recipe is None or not (recipe.mxfp8() or recipe.nvfp4()): - return False - # NVFP4 graph-safe grouped quantize currently requires RHT. - if recipe.nvfp4() and recipe.disable_rht: + if len(ops) < 3 or _grouped_mlp_recipe_unsupported_reason(recipe) is not None: return False - # The fused MXFP8 backward reinterprets grad-output storage as E4M3. It - # cannot consume E5M2 gradients from Format.HYBRID. NVFP4 has separate formats. - if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: - return False - return any( - isinstance(fc1, GroupedLinear) - and isinstance(activation, activation_op_types) - and isinstance(fc2, GroupedLinear) - for fc1, activation, fc2 in zip(ops, ops[1:], ops[2:]) - ) + return _find_grouped_mlp_activation(ops, activation_op_types) is not None def fuse_grouped_mlp_ops( @@ -1045,8 +1116,17 @@ def fuse_grouped_mlp_ops( Updated operations with matched triples replaced by fused ops. """ if not _is_grouped_mlp_fusion_candidate(ops, recipe, activation_op_types): + _maybe_warn_grouped_mlp_fallback( + ops, + activation_op_types, + fused_op_cls, + lambda: _grouped_mlp_recipe_unsupported_reason(recipe), + ) return ops if not fused_op_cls.is_supported(): + _maybe_warn_grouped_mlp_fallback( + ops, activation_op_types, fused_op_cls, fused_op_cls._unsupported_reason + ) return ops # Check for unsupported NVFP4 recipe configs @@ -1056,12 +1136,22 @@ def fuse_grouped_mlp_ops( return ops if recipe.row_scaled_activation or recipe.nvfp4_4over6 != "none": # 4over6 doesn't used fused kernels + _maybe_warn_grouped_mlp_fallback( + ops, + activation_op_types, + fused_op_cls, + lambda: "NVFP4 row-scaled activations and 4over6 are not supported", + ) return ops if recipe.fp8_format == RecipeFormat.UE5M3: # cuDNN has no SReLU support for UE5M3 for now - activation_op_types = tuple( - filter(lambda t: t not in (ScaledSReLU, ScaledTanhSReLU), activation_op_types) + srelu_op_types = tuple( + t for t in activation_op_types if t in _UNARY_ACTIVATION_OP_TYPES + ) + _maybe_warn_grouped_mlp_fallback( + ops, srelu_op_types, fused_op_cls, lambda: "not supported with NVFP4 UE5M3 scales" ) + activation_op_types = tuple(t for t in activation_op_types if t not in srelu_op_types) if not activation_op_types: return ops @@ -1087,11 +1177,22 @@ def fuse_grouped_mlp_ops( ) ): matches_pattern = False + if _grouped_mlp_fallback_warnings_enabled(): + _warn_grouped_mlp_fallback( + fused_op_cls.__name__, + type(window[1]).__name__, + "non-default parameters require nvidia-cudnn-frontend>=1.24.0," + f" got {_cudnn_frontend_version_str()}", + ) else: try: validate_grouped_mlp_dims(window[0], window[1], window[2]) - except (TypeError, ValueError): + except (TypeError, ValueError) as e: matches_pattern = False + if _grouped_mlp_fallback_warnings_enabled(): + _warn_grouped_mlp_fallback( + fused_op_cls.__name__, type(window[1]).__name__, str(e) + ) if matches_pattern: op = fused_op_cls( @@ -1171,21 +1272,29 @@ def grouped_gemm_act_hadamard_quant_kernel(cls) -> Optional[Callable]: return None @classmethod - @functools.lru_cache(maxsize=None) - def is_supported(cls) -> bool: - """Whether this fused operation is supported on the current system.""" - if not torch.cuda.is_available() or get_device_compute_capability()[0] != 10: - return False + def _unsupported_reason(cls) -> Optional[str]: + """Why this fused operation cannot run on the current system, or ``None`` if it can.""" + if not torch.cuda.is_available(): + return "CUDA is not available" + device_arch = get_device_compute_capability() + if device_arch[0] != 10: + return f"requires compute capability 10.x, got {device_arch[0]}.{device_arch[1]}" if not _cudnn_frontend_version_supported(): - return False + return f"requires nvidia-cudnn-frontend>=1.23.0, got {_cudnn_frontend_version_str()}" try: cls.grouped_gemm_activation_kernel() cls.grouped_gemm_dactivation_kernel() cls.grouped_gemm_quant_kernel() cls.grouped_gemm_wgrad_kernel() - except ImportError: - return False - return True + except ImportError as e: + return f"failed to import cuDNN frontend kernels: {e}" + return None + + @classmethod + @functools.lru_cache(maxsize=None) + def is_supported(cls) -> bool: + """Whether this fused operation is supported on the current system.""" + return cls._unsupported_reason() is None def __init__( self, @@ -3158,10 +3267,11 @@ class GroupedMLP_CuTeGEMMUnary(_GroupedMLP_CuTeGEMMBase): """Joint fused op for block-scaled GroupedLinear + scaled unary activation + GroupedLinear.""" @classmethod - @functools.lru_cache(maxsize=None) - def is_supported(cls) -> bool: - """Whether the SReLU fused operation is supported on the current system.""" - return _cudnn_frontend_supports_grouped_gemm_srelu() and super().is_supported() + def _unsupported_reason(cls) -> Optional[str]: + """SReLU kernels additionally require cuDNN frontend >= 1.24.0.""" + if not _cudnn_frontend_supports_grouped_gemm_srelu(): + return f"requires nvidia-cudnn-frontend>=1.24.0, got {_cudnn_frontend_version_str()}" + return super()._unsupported_reason() @classmethod @functools.lru_cache(maxsize=None) @@ -3218,6 +3328,9 @@ def fuse_glu_ops( """Apply joint GroupedLinear + scaled GLU + GroupedLinear fusion.""" if not torch.cuda.is_available(): + _maybe_warn_grouped_mlp_fallback( + ops, _GLU_ACTIVATION_OP_TYPES, GroupedMLP_CuTeGEMMGLU, lambda: "CUDA is not available" + ) return ops # Determine supported activations @@ -3233,7 +3346,22 @@ def fuse_glu_ops( activation_op_types.append(ScaledClampedQGeGLU) else: # Unsupported device arch + _maybe_warn_grouped_mlp_fallback( + ops, + _GLU_ACTIVATION_OP_TYPES, + GroupedMLP_CuTeGEMMGLU, + lambda: f"unsupported compute capability {device_arch[0]}.{device_arch[1]}", + ) return ops + _maybe_warn_grouped_mlp_fallback( + ops, + tuple(t for t in _GLU_ACTIVATION_OP_TYPES if t not in activation_op_types), + GroupedMLP_CuTeGEMMGLU, + lambda: ( + f"not supported on compute capability {device_arch[0]}.{device_arch[1]}" + f" with nvidia-cudnn-frontend {_cudnn_frontend_version_str()}" + ), + ) return fuse_grouped_mlp_ops( ops, @@ -3252,12 +3380,25 @@ def fuse_unary_activation_ops( """Apply joint GroupedLinear + scaled unary activation + GroupedLinear fusion.""" if not GroupedMLP_CuTeGEMMUnary.is_supported(): + _maybe_warn_grouped_mlp_fallback( + ops, + _UNARY_ACTIVATION_OP_TYPES, + GroupedMLP_CuTeGEMMUnary, + GroupedMLP_CuTeGEMMUnary._unsupported_reason, + ) return ops # Determine supported activations activation_op_types = [ScaledSReLU] if _cudnn_frontend_supports_grouped_gemm_srelu_tanh(): activation_op_types.append(ScaledTanhSReLU) + else: + _maybe_warn_grouped_mlp_fallback( + ops, + (ScaledTanhSReLU,), + GroupedMLP_CuTeGEMMUnary, + lambda: f"not supported by nvidia-cudnn-frontend {_cudnn_frontend_version_str()}", + ) return fuse_grouped_mlp_ops( ops,