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..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/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..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. -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 -``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 8463c387b4c..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", - "`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", - "\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 559bab885ee..d5bd6e210df 100644 --- a/docs/examples/te_mixtral/utils.py +++ b/docs/examples/te_mixtral/utils.py @@ -148,76 +148,6 @@ def init_baseline_model(hyperparams: HyperParameters): return model -def _enable_fused_mxfp8_grouped_mlp() -> None: - """Improvement 3: enable the fused ``ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8`` and - backward kernel in the installed TE 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``. - - We also (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 - 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 int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - return False - 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) @@ -227,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/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/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_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 f1ca21a4b2f..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 @@ -459,10 +524,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, @@ -486,9 +549,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, @@ -535,13 +597,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) @@ -937,13 +992,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 quantization in nvfp4_variant_names and dtype != torch.bfloat16: @@ -1223,13 +1271,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: @@ -1713,8 +1754,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, @@ -2188,8 +2227,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": @@ -2694,8 +2731,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") @@ -2827,8 +2862,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 5c0dbf3ab28..11188f2eaf7 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 ( @@ -1681,15 +1681,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 @@ -1764,9 +1762,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: @@ -1980,8 +1976,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 508895505fd..30396a15917 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, @@ -170,17 +170,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`` @@ -221,9 +219,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 e30b98eea3f..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. @@ -532,9 +541,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 @@ -961,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: @@ -980,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__}).") @@ -999,6 +1010,87 @@ 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 _grouped_mlp_recipe_unsupported_reason(recipe) is not None: + return False + return _find_grouped_mlp_activation(ops, activation_op_types) is not None + + def fuse_grouped_mlp_ops( ops: list[FusibleOperation], *, @@ -1023,21 +1115,18 @@ def fuse_grouped_mlp_ops( list of FusibleOperation Updated operations with matched triples replaced by fused ops. """ - if not fused_op_cls.is_supported(): - return ops - - # Fused kernels are only supported for MXFP8 and NVFP4 - if recipe is None: - return ops - if recipe.custom(): - # Check if custom recipe explicitly enables fusion - if not getattr(recipe, "enable_cutedsl_fused_grouped_mlp", False): - return ops - elif not (recipe.mxfp8() or recipe.nvfp4()): + 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 - - # MXFP8 kernel assumes E4M3 data, so reject hybrid E4M3/E5M2 data - if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: + 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 @@ -1047,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 @@ -1078,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( @@ -1162,23 +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 int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - return False - if 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, @@ -3151,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) @@ -3210,6 +3327,12 @@ def fuse_glu_ops( ) -> list[FusibleOperation]: """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 activation_op_types = [] device_arch = get_device_compute_capability() @@ -3223,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, @@ -3241,10 +3379,26 @@ def fuse_unary_activation_ops( ) -> list[FusibleOperation]: """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, @@ -3254,8 +3408,5 @@ 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_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(