From 1a16d5018446ccef15a463d8d883c374f8bcacb3 Mon Sep 17 00:00:00 2001 From: Zhongbo Zhu <42691305+zhongbozhu@users.noreply.github.com> Date: Fri, 14 Aug 2026 22:53:55 +0000 Subject: [PATCH 01/10] [Main] Support device-init grouped linear module too with when not using TE opfuser (#6000) Signed-off-by: Zhongbo Zhu <42691305+zhongbozhu@users.noreply.github.com> --- .../core/extensions/transformer_engine.py | 13 + megatron/core/transformer/moe/experts.py | 116 +++- megatron/core/transformer/moe/moe_utils.py | 30 +- .../core/transformer/moe/token_dispatcher.py | 5 +- .../core/transformer/transformer_config.py | 48 +- .../models/test_hybrid_moe_model.py | 1 + .../transformer/moe/test_grouped_mlp.py | 341 +++++++++- ...test_grouped_tensor_dispatcher_numerics.py | 599 ++++++++++++++++++ ...test_moe_single_grouped_weight_numerics.py | 51 +- 9 files changed, 1158 insertions(+), 46 deletions(-) create mode 100644 tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 348847e7399..365faba6331 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -1951,6 +1951,10 @@ def sharded_state_dict( if HAVE_TE and is_te_min_version("1.9.0.dev0"): + _TE_GROUPED_LINEAR_SUPPORTS_GROUPED_TENSOR = ( + "use_grouped_tensor" in inspect.signature(te.pytorch.GroupedLinear.__init__).parameters + ) + class TEGroupedLinear(te.pytorch.GroupedLinear): """ Wrapper for the Transformer-Engine's `GroupedLinear` layer. @@ -2064,6 +2068,14 @@ def __init__( config, "moe_single_grouped_bias", False ) + if _TE_GROUPED_LINEAR_SUPPORTS_GROUPED_TENSOR: + extra_kwargs["use_grouped_tensor"] = config.moe_use_grouped_tensor + elif config.moe_use_grouped_tensor and not config.use_transformer_engine_op_fuser: + raise RuntimeError( + "moe_use_grouped_tensor=True requires a Transformer Engine GroupedLinear " + "that exposes the use_grouped_tensor argument." + ) + self.te_quant_params: Optional[TEQuantizationParams] = None quant_config = get_quant_config_or_none(name, config.quant_recipe) self.finish_init(quant_config) @@ -2565,6 +2577,7 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): ) else: + _TE_GROUPED_LINEAR_SUPPORTS_GROUPED_TENSOR = False TEGroupedLinear = None # type: ignore[assignment, misc] TEColumnParallelGroupedLinear = None # type: ignore[assignment, misc] TERowParallelGroupedLinear = None # type: ignore[assignment, misc] diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index ec9ba66e809..cbd9f10d81c 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -292,25 +292,74 @@ def __init__( self.config.moe_mlp_glu_interleave_size, ) - if self.config.fp8 or self.config.fp4: - assert HAVE_TE, "FP8 and FP4 requires TE." - align_size = 256 if self._with_fused_impl else None + self._use_grouped_tensor = self.config.moe_use_grouped_tensor + if self.config.fp8 or self.config.fp4 or self._use_grouped_tensor: + assert HAVE_TE, "Quantized or TE grouped-tensor GroupedMLP execution requires TE." + align_size = ( + get_align_size_for_quantization(self.config) if self._use_grouped_tensor else None + ) self.quantization_padding = Fp8Padding(self.num_local_experts, align_size=align_size) self.quantization_unpadding = Fp8Unpadding( self.num_local_experts, align_size=align_size ) + @staticmethod + def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, permuted_probs): + """Apply a packed expert bias without reading token counts on the host.""" + # TODO: get rid of the .float() by having fused kernel compute in FP32 + shape = intermediate_parallel.shape + hidden_size = shape[-1] + output_dtype = intermediate_parallel.dtype + flat_output = intermediate_parallel.view(-1, hidden_size).float() + flat_probs = permuted_probs.reshape(-1, 1).float() + + if tokens_per_expert.device != packed_bias.device: + raise ValueError("Packed MoE bias and tokens_per_expert must be on the same device.") + + # Permutation stores tokens contiguously by expert. Repeat bias row e by that expert's + # token count to create one bias row per permuted token: + # + # packed_bias = [bias_e0, bias_e1] + # tokens_per_expert = [ 2, 1] + # bias_per_token = [bias_e0, bias_e0, bias_e1] + # + # output_size avoids a stream synchronization to compute sum(tokens_per_expert). + # Cast before repeating so both forward arithmetic and repeat_interleave's backward + # reduction are computed in FP32. Autograd casts the final parameter gradient once. + bias_per_token = torch.repeat_interleave( + packed_bias.float(), tokens_per_expert, dim=0, output_size=flat_output.size(0) + ) + return (flat_output + bias_per_token * flat_probs).view(shape).to(output_dtype) + @staticmethod def _apply_bias(intermediate_parallel, bias_parallel, tokens_per_expert, permuted_probs): if bias_parallel is None: return intermediate_parallel + + # CUDA-graph-safe packed path. With single_grouped_bias=True, TE returns one packed + # GroupedTensor [num_experts, hidden_size]. The grouped-tensor backend also provides + # tokens_per_expert as a tensor on the same device. + if isinstance(bias_parallel, torch.Tensor) and isinstance(tokens_per_expert, torch.Tensor): + return TEGroupedMLP._apply_packed_bias( + intermediate_parallel, bias_parallel, tokens_per_expert, permuted_probs + ) + + # Eager-only CPU-metadata path. The legacy contract returns List[Tensor[hidden_size]], + # and torch.split plus the Python zip below require concrete host token counts. A packed + # bias paired with Python counts also uses this compatibility path. Converting a tensor + # with .tolist() synchronizes and copies device data to the host, so this path must never + # be included in a CUDA graph. + if isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = tokens_per_expert.tolist() + shape = intermediate_parallel.shape + flat_output = intermediate_parallel.view(-1, shape[-1]) return ( torch.cat( [ t + b * p for t, b, p in zip( - torch.split(intermediate_parallel.view(-1, shape[-1]), tokens_per_expert), + torch.split(flat_output, tokens_per_expert), bias_parallel, torch.split(permuted_probs, tokens_per_expert), ) @@ -659,9 +708,21 @@ def _fused_forward( # Apply padding if needed unpadded_tokens_per_expert = None + # Some dispatchers have already padded each expert's token segment before the tokens + # reach this module: + # * router padding changes the routing map before dispatch; + # * HybridEP/NCCL-EP pad as part of their fused dispatch/permute operation; + # * DeepEP can pad in the fused local permutation after communication. + # Padding those tensors again would insert a second set of dummy tokens and make + # tokens_per_expert disagree with the already-permuted token buffer, so skip the local + # Fp8Padding fallback in those cases. if skip_routed_expert_padding(self.config): pass - elif self.config.fp8 or self.config.fp4: + # Regular AllToAll normally reaches this branch because its permutation does not insert + # the padding needed by the fused grouped-MLP contract. FP8/FP4 require recipe-specific + # alignment, while the TE operation-fuser grouped-tensor path currently uses 256-token + # expert segments. + elif self.config.fp8 or self.config.fp4 or self._use_grouped_tensor: tokens_per_expert = tokens_per_expert.tolist() unpadded_tokens_per_expert = tokens_per_expert permuted_local_hidden_states, tokens_per_expert = self.quantization_padding( @@ -672,7 +733,22 @@ def _fused_forward( ) permuted_probs = permuted_probs.squeeze(-1) tokens_per_expert = torch.tensor( - tokens_per_expert, dtype=torch.int, device=permuted_probs.device + tokens_per_expert, dtype=torch.int64, device=permuted_probs.device + ) + + if self._use_grouped_tensor: + if not isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = torch.tensor( + tokens_per_expert, dtype=torch.int64, device=permuted_local_hidden_states.device + ) + else: + tokens_per_expert = tokens_per_expert.to( + device=permuted_local_hidden_states.device, dtype=torch.int64, non_blocking=True + ) + else: + raise RuntimeError( + "The Transformer Engine operation-fuser MoE path requires " + "moe_use_grouped_tensor=True." ) # if the number of tokens is 0, pad the hidden states to 256 @@ -761,11 +837,23 @@ def forward( # Apply padding if needed unpadded_tokens_per_expert = None - tokens_per_expert: list[int] = tokens_per_expert.tolist() permuted_probs = permuted_probs.unsqueeze(-1) + # The token buffer may already contain per-expert padding when padding was performed + # before expert compute: + # * router padding modified the routing map before dispatch; + # * HybridEP/NCCL-EP fused padding into dispatch/permute; + # * DeepEP fused padding into its post-communication local permutation. + # In those cases tokens_per_expert already describes the padded expert segments. Running + # Fp8Padding again would change the segment lengths without matching the existing token + # layout, so this module must leave both tensors unchanged. if skip_routed_expert_padding(self.config): pass - elif self.config.fp8 or self.config.fp4: + # Regular AllToAll normally supplies unpadded expert segments and therefore uses this + # explicit fallback. FP8/FP4 need their recipe-specific alignment. MCore currently also + # applies its common aligned-segment contract to the GroupedTensor backend so quantized + # grouped execution receives supported shapes + elif self.config.fp8 or self.config.fp4 or self._use_grouped_tensor: + tokens_per_expert = tokens_per_expert.tolist() unpadded_tokens_per_expert = tokens_per_expert permuted_local_hidden_states, tokens_per_expert = self.quantization_padding( permuted_local_hidden_states, tokens_per_expert @@ -774,6 +862,18 @@ def forward( permuted_probs, unpadded_tokens_per_expert ) + if self._use_grouped_tensor: + if not isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = torch.tensor( + tokens_per_expert, dtype=torch.int64, device=permuted_local_hidden_states.device + ) + else: + tokens_per_expert = tokens_per_expert.to( + device=permuted_local_hidden_states.device, dtype=torch.int64, non_blocking=True + ) + elif isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = tokens_per_expert.tolist() + if self.config.moe_apply_probs_on_input: assert ( self.config.moe_router_topk == 1 diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index c8e197d2a3b..c8a7e2e3ce5 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -1354,8 +1354,11 @@ def forward( inp = inp.view(-1, inp_shape[-1]) if te_general_gemm is not None and router_dtype != torch.float64: - output = te_general_gemm(weight, inp, router_dtype, layout="TN", bias=bias) - output = output[0] + # cuBLASLt's non-FP8 bias epilogue expects bias and output to have the same + # dtype. Router parameters may be BF16 while router logits are FP32, so cast the + # small bias vector before passing it to TE. + gemm_bias = bias.to(router_dtype) if bias is not None else None + output = te_general_gemm(weight, inp, router_dtype, layout="TN", bias=gemm_bias)[0] elif bias is None: output = torch.mm(inp.to(router_dtype), weight.to(router_dtype).t()) else: @@ -1433,22 +1436,33 @@ def get_align_size_for_quantization(config: TransformerConfig) -> int: Returns: int: The alignment size for quantization. """ - # CUTLASS kernel for grouped GEMM assumes 256 alignment. - if config.use_transformer_engine_op_fuser: + # TE's grouped-tensor and fused grouped-MLP kernels require 256-token alignment. + if config.use_transformer_engine_op_fuser or config.moe_use_grouped_tensor: return 256 if config.fp8: return get_fp8_align_size(config.fp8_recipe) if config.fp4: return get_fp4_align_size(config.fp4_recipe) - # Only FP8 or FP4 requires padding. Defaults to 0. + # Legacy high-precision grouped GEMM does not require padding. Defaults to 0. return 0 +def _deepep_permute_pads_grouped_tensor_input(config: TransformerConfig) -> bool: + """Whether DeepEP fused permutation pads input for TE grouped-tensor GEMM.""" + return ( + config.moe_use_grouped_tensor + and config.moe_token_dispatcher_type == "flex" + and config.moe_flex_dispatcher_backend == "deepep" + and config.moe_permute_fusion + and fused_permute_and_pad_with_probs is not None + ) + + def skip_routed_expert_padding(config: TransformerConfig) -> bool: """Whether the expert module should skip quantization padding. - Returns True when padding is already applied by the router or the - HybridEP / NCCL-EP dispatcher. + Returns True when padding is already applied by the router, the HybridEP / NCCL-EP + dispatcher, or DeepEP's fused permutation kernel. """ if config.moe_router_padding_for_quantization: return True @@ -1457,6 +1471,8 @@ def skip_routed_expert_padding(config: TransformerConfig) -> bool: "ncclep", ): return True + if _deepep_permute_pads_grouped_tensor_input(config): + return True return False diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 1bfead220cd..de7367bd421 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1198,8 +1198,9 @@ def dispatch( "HybridEP only supports float32 probs, please set --moe-router-dtype=fp32" ) self.token_probs = self.token_probs.float() # downcast or upcast - if self.config.fp8 or self.config.fp4: - self.pad_multiple = get_align_size_for_quantization(self.config) + align_size = get_align_size_for_quantization(self.config) + if align_size > 0: + self.pad_multiple = align_size if self._padded_num_tokens is not None and hidden_states.shape[0] < self._padded_num_tokens: pad_rows = self._padded_num_tokens - hidden_states.shape[0] hidden_states = torch.cat( diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index e8cd33aedb8..f578ece5ed5 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -918,21 +918,29 @@ class TransformerConfig(ModelParallelConfig): use for debugging.""" moe_grouped_gemm: bool = False - """When there are multiple experts per rank, compress multiple local (potentially small) gemms - in a single kernel launch to improve the utilization and performance by leveraging the Grouped - GEMM feature introduced since CUTLASS 2.8 (https://github.com/fanshiqing/grouped_gemm). + """Use grouped GEMM to execute multiple local MoE experts together. + + The concrete implementation is selected by Transformer Engine. Set + ``moe_use_grouped_tensor=True`` to use its CUDA-graph-safe GroupedTensor path. + """ + + moe_use_grouped_tensor: bool = False + """Use Transformer Engine's native GroupedTensor path for grouped MoE GEMMs. + + This path uses padded expert segments and CUDA split metadata so it can be captured in CUDA + graphs. Enabling the Transformer Engine operation fuser also enables this option. """ moe_single_grouped_weight: bool = False """When using TE GroupedLinear for MoE experts, store expert weights as a single grouped parameter via Transformer Engine's `GroupedTensor`. Requires ``moe_grouped_gemm=True`` and - ``use_transformer_engine_op_fuser=True``. + ``moe_use_grouped_tensor=True``. """ moe_single_grouped_bias: bool = False """When using TE GroupedLinear for MoE experts, store expert biases as a single grouped - parameter via Transformer Engine's `GroupedTensor`. Requires ``moe_grouped_gemm=True`` - and ``add_bias_linear=True``.""" + parameter via Transformer Engine's `GroupedTensor`. Requires ``moe_grouped_gemm=True``, + ``moe_use_grouped_tensor=True``, and ``add_bias_linear=True``.""" moe_aux_loss_coeff: Union[float, List[float]] = 0.0 """Scaling coefficient for the aux loss. A starting value of 1e-2 is recommended. @@ -1601,6 +1609,12 @@ def __post_init__(self): self.experimental_attention_variant ) + if self.use_transformer_engine_op_fuser and self.moe_grouped_gemm: + self.moe_use_grouped_tensor = True + + if self.moe_use_grouped_tensor and not self.moe_grouped_gemm: + raise ValueError("moe_use_grouped_tensor=True requires moe_grouped_gemm=True.") + if self.cp_partition_mode not in ("zigzag", "contiguous"): raise ValueError(f"Unsupported cp_partition_mode: {self.cp_partition_mode}") @@ -2067,15 +2081,18 @@ def __post_init__(self): "moe_single_grouped_weight is currently supported with high-precision " "primary weights, fp8_recipe='mxfp8', or fp4_recipe='nvfp4'." ) - if not self.use_transformer_engine_op_fuser: + if self.fp4 and not self.fp4_param: raise ValueError( - "moe_single_grouped_weight requires " - "use_transformer_engine_op_fuser=True. The non-op-fuser TE GroupedLinear " - "path splits the grouped parameter into per-expert tensors and does not " - "support single-grouped-weight training." + "moe_single_grouped_weight with FP4 compute requires fp4_param=True " + "(--fp4-param-gather). Without FP4 parameter gather, Transformer Engine " + "uses a split-quantize fallback that is being deprecated." ) + if not self.moe_use_grouped_tensor: + raise ValueError("moe_single_grouped_weight requires moe_use_grouped_tensor=True.") if self.moe_single_grouped_bias and not self.add_bias_linear: raise ValueError("moe_single_grouped_bias requires add_bias_linear=True.") + if self.moe_single_grouped_bias and not self.moe_use_grouped_tensor: + raise ValueError("moe_single_grouped_bias requires moe_use_grouped_tensor=True.") if self.moe_enable_deepep: if self.moe_token_dispatcher_type != "flex": @@ -2103,6 +2120,12 @@ def __post_init__(self): "moe_flex_dispatcher_backend='ncclep' requires " "moe_token_dispatcher_type='flex'." ) + if self.moe_use_grouped_tensor and not self.use_transformer_engine_op_fuser: + raise ValueError( + "moe_use_grouped_tensor=True without use_transformer_engine_op_fuser is " + "not yet supported with the NCCL-EP dispatcher. Use the TE op-fuser path " + "or select the alltoall, DeepEP, or HybridEP dispatcher." + ) # moe_deepep_num_sms / moe_hybridep_num_sms are deprecated and unified into # moe_flex_dispatcher_num_sms. If either is set, route it (an explicit @@ -2189,10 +2212,11 @@ def __post_init__(self): if ( self.moe_flex_dispatcher_backend == "hybridep" and not self.use_transformer_engine_op_fuser + and not self.moe_use_grouped_tensor ): raise ValueError( "moe_expert_rank_capacity_factor with the 'hybridep' backend requires " - "use_transformer_engine_op_fuser to be enabled." + "use_transformer_engine_op_fuser=True or moe_use_grouped_tensor=True." ) if self.cpu_offloading and ( diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index edf7608cbce..4840db03630 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -201,6 +201,7 @@ "moe_flex_dispatcher_num_sms": None, "moe_grad_scale_func": None, "moe_grouped_gemm": True, + "moe_use_grouped_tensor": False, "moe_hybridep_num_sms": None, "moe_hybridep_num_sms_preprocessing": 108, "moe_hybridep_num_blocks_permute": None, diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index 09e26605a98..e56ea17066c 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -31,11 +31,41 @@ def test_op_fuser_transformer_config_args_are_exposed(): _add_network_size_args(parser) args = parser.parse_args( - ["--use-transformer-engine-op-fuser", "--moe-mlp-glu-interleave-size", "16"] + [ + "--use-transformer-engine-op-fuser", + "--moe-mlp-glu-interleave-size", + "16", + "--moe-use-grouped-tensor", + ] ) assert args.use_transformer_engine_op_fuser is True assert args.moe_mlp_glu_interleave_size == 16 + assert args.moe_use_grouped_tensor is True + + +def test_op_fuser_enables_grouped_tensor(): + config = TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=4, + num_moe_experts=2, + moe_grouped_gemm=True, + use_transformer_engine_op_fuser=True, + ) + + assert config.moe_use_grouped_tensor is True + + +def test_grouped_tensor_requires_grouped_gemm(): + with pytest.raises(ValueError, match="requires moe_grouped_gemm=True"): + TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=4, + num_moe_experts=2, + moe_use_grouped_tensor=True, + ) def test_remove_glu_interleaving_restores_contiguous_gate_and_linear_halves(): @@ -170,9 +200,13 @@ def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): moe_router_padding_for_quantization=False, moe_token_dispatcher_type=None, moe_flex_dispatcher_backend=None, + moe_use_grouped_tensor=True, moe_paged_stash=False, delay_offload_until_cuda_graph=False, ) + module._use_grouped_tensor = True + module.quantization_padding = lambda tensor, token_counts: (tensor, token_counts) + module.quantization_unpadding = lambda tensor, token_counts: tensor module._fused_ops = None fused_ops = FakeFusedOps() module._make_fused_ops = lambda: fused_ops @@ -185,9 +219,9 @@ def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): torch.testing.assert_close(output, torch.ones_like(hidden_states)) assert module._fused_ops[0] is fused_ops assert fused_ops.args[0] is hidden_states - assert fused_ops.args[1] is tokens_per_expert - assert fused_ops.args[2] is probs - assert fused_ops.args[3] is tokens_per_expert + torch.testing.assert_close(fused_ops.args[1], tokens_per_expert) + torch.testing.assert_close(fused_ops.args[2], probs) + torch.testing.assert_close(fused_ops.args[3], tokens_per_expert) def test_apply_bias_returns_input_unchanged_when_bias_is_none(): @@ -215,6 +249,20 @@ def test_apply_bias_combines_per_expert_bias_and_probs(): assert output.dtype == intermediate.dtype +def test_apply_bias_combines_packed_grouped_bias_and_accumulates_gradient(): + intermediate = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) + packed_bias = torch.tensor([[10.0, 20.0], [100.0, 200.0]], requires_grad=True) + tokens_per_expert = torch.tensor([2, 1], dtype=torch.int64) + permuted_probs = torch.tensor([0.25, 0.5, 1.5]) + expected = torch.tensor([[3.5, 7.0], [8.0, 14.0], [155.0, 306.0]]) + + output = TEGroupedMLP._apply_bias(intermediate, packed_bias, tokens_per_expert, permuted_probs) + output.sum().backward() + + torch.testing.assert_close(output, expected) + torch.testing.assert_close(packed_bias.grad, torch.tensor([[0.75, 0.75], [1.5, 1.5]])) + + def test_make_fused_impl_pre_forward_hook_dispatches_submodule_hooks(): module = TEGroupedMLP.__new__(TEGroupedMLP) torch.nn.Module.__init__(module) @@ -1089,6 +1137,291 @@ def test_gpu_make_fused_ops_constructs_with_real_te(self): experts.linear_fc2, f"weight{idx}" ) + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.internal + @pytest.mark.parametrize("single_grouped_bias", (False, True)) + def test_gpu_fused_path_scales_fc2_bias(self, monkeypatch, single_grouped_bias): + """FC2 bias and its gradients must use the per-token router probability.""" + try: + from transformer_engine.pytorch.ops import GroupedLinear + except ImportError: + pytest.skip("TE op fuser API not available") + import inspect + + if "scale_bias" not in inspect.signature(GroupedLinear.__init__).parameters: + pytest.skip("Installed TE op fuser GroupedLinear lacks `scale_bias` support") + if single_grouped_bias: + monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + + tf_config = TransformerConfig( + num_layers=1, + hidden_size=self.hidden_size, + num_attention_heads=4, + num_moe_experts=self.num_experts, + use_cpu_initialization=False, + add_bias_linear=True, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=False, + bias_dropout_fusion=False, + bf16=True, + params_dtype=torch.bfloat16, + moe_router_load_balancing_type="sinkhorn", + moe_router_topk=1, + moe_grouped_gemm=True, + use_transformer_engine_op_fuser=True, + moe_single_grouped_bias=single_grouped_bias, + ) + _set_random_seed(seed_=123, data_parallel_random_init=False) + submodules = get_submodules( + get_gpt_layer_with_transformer_engine_submodules( + self.num_experts, moe_grouped_gemm=True + ).mlp + ) + layer = MoELayer(tf_config, submodules) + layer = Float16Module(layer.config, layer).module + layer.cuda() + experts = layer.experts + assert isinstance(experts, TEGroupedMLP) + + with torch.no_grad(): + for linear in (experts.linear_fc1, experts.linear_fc2): + for expert_idx in range(self.num_experts): + getattr(linear, f"weight{expert_idx}").zero_() + if not single_grouped_bias: + getattr(linear, f"bias{expert_idx}").zero_() + if single_grouped_bias: + linear.bias.rowwise_data.zero_() + if single_grouped_bias: + packed_fc2_bias = experts.linear_fc2.bias.rowwise_data.view( + self.num_experts, self.hidden_size + ) + packed_fc2_bias[0].fill_(2.0) + packed_fc2_bias[1].fill_(4.0) + else: + experts.linear_fc2.bias0.fill_(2.0) + experts.linear_fc2.bias1.fill_(4.0) + + hidden_states = torch.zeros( + 3, self.hidden_size, dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + tokens_per_expert = torch.tensor([2, 1], dtype=torch.int32, device="cuda") + probs = torch.tensor( + [0.25, 0.5, 0.125], dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + + output, _ = experts(hidden_states, tokens_per_expert, probs) + expected_output = torch.cat( + ( + probs[:2, None] * torch.full_like(output[:2], 2.0), + probs[2:, None] * torch.full_like(output[2:], 4.0), + ) + ) + torch.testing.assert_close(output, expected_output) + + output.sum().backward() + expected_prob_grad = ( + torch.tensor([2.0, 2.0, 4.0], dtype=torch.bfloat16, device="cuda") * self.hidden_size + ) + torch.testing.assert_close(probs.grad, expected_prob_grad) + if single_grouped_bias: + assert experts.linear_fc2.bias.grad is not None + packed_dbias = experts.linear_fc2.bias.grad.view(self.num_experts, self.hidden_size) + torch.testing.assert_close( + packed_dbias[0], torch.ones_like(packed_dbias[0]) * probs[:2].detach().sum() + ) + torch.testing.assert_close( + packed_dbias[1], torch.ones_like(packed_dbias[1]) * probs[2:].detach().sum() + ) + assert experts._fused_ops[0][2].bias is experts.linear_fc2.bias + else: + torch.testing.assert_close( + experts.linear_fc2.bias0.grad, + torch.ones_like(experts.linear_fc2.bias0) * probs[:2].detach().sum(), + ) + torch.testing.assert_close( + experts.linear_fc2.bias1.grad, + torch.ones_like(experts.linear_fc2.bias1) * probs[2:].detach().sum(), + ) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.internal + @pytest.mark.parametrize("use_op_fuser", (False, True), ids=("module", "op-fuser")) + @pytest.mark.parametrize( + "single_grouped_weight,single_grouped_bias", + ((True, False), (False, True), (True, True)), + ids=("single-weight", "single-bias", "single-weight-and-bias"), + ) + def test_gpu_single_grouped_parent_gradient_parity( + self, monkeypatch, use_op_fuser, single_grouped_weight, single_grouped_bias + ): + """Single grouped parents must receive the same gradients as discrete parameters. + + This is an integration test over the real MCore TEGroupedMLP wrapper. It covers both + gradient ownership mechanisms: + + * the module path returns the packed FC2 bias to ``_apply_packed_bias``, where normal + PyTorch autograd must update the registered grouped parent; + * the op-fuser path reattaches the same wrapper parameters to TE op shells, whose custom + backward must return packed wgrad/dbias in the corresponding parent slots. + + Comparing gradients directly is intentional. A forward or dgrad-only check would not + catch a disconnected grouped parent that the optimizer can never update. + """ + try: + from transformer_engine.pytorch.module import GroupedLinear as ModuleGroupedLinear + from transformer_engine.pytorch.ops import GroupedLinear as OpGroupedLinear + except ImportError: + pytest.skip("Required TE GroupedLinear APIs are not available") + import inspect + + module_parameters = inspect.signature(ModuleGroupedLinear.__init__).parameters + op_parameters = inspect.signature(OpGroupedLinear.__init__).parameters + if ( + "use_grouped_tensor" not in module_parameters + or "single_grouped_bias" not in module_parameters + or "single_grouped_bias" not in op_parameters + ): + pytest.skip("Installed TE lacks native single grouped bias support") + + monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") + + def build_experts(single_weight, single_bias): + config = TransformerConfig( + num_layers=1, + hidden_size=self.hidden_size, + num_attention_heads=4, + num_moe_experts=self.num_experts, + use_cpu_initialization=False, + add_bias_linear=True, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=False, + bias_dropout_fusion=False, + bf16=True, + params_dtype=torch.bfloat16, + moe_router_load_balancing_type="sinkhorn", + moe_router_topk=1, + moe_grouped_gemm=True, + moe_use_grouped_tensor=True, + use_transformer_engine_op_fuser=use_op_fuser, + moe_single_grouped_weight=single_weight, + moe_single_grouped_bias=single_bias, + ) + submodules = get_submodules( + get_gpt_layer_with_transformer_engine_submodules( + self.num_experts, moe_grouped_gemm=True + ).mlp + ) + assert isinstance(submodules, MoESubmodules) + layer = MoELayer(config, submodules) + layer = Float16Module(layer.config, layer).module + layer.cuda() + assert isinstance(layer.experts, TEGroupedMLP) + return layer.experts + + def copy_linear_params(reference_linear, target_linear): + reference_weights = torch.stack( + [ + getattr(reference_linear, f"weight{idx}").detach() + for idx in range(self.num_experts) + ] + ) + reference_biases = torch.stack( + [ + getattr(reference_linear, f"bias{idx}").detach() + for idx in range(self.num_experts) + ] + ) + with torch.no_grad(): + if target_linear.single_grouped_weight: + target_linear.weight.rowwise_data.view_as(reference_weights).copy_( + reference_weights + ) + else: + for idx in range(self.num_experts): + getattr(target_linear, f"weight{idx}").copy_(reference_weights[idx]) + + if target_linear.single_grouped_bias: + target_linear.bias.rowwise_data.view_as(reference_biases).copy_( + reference_biases + ) + else: + for idx in range(self.num_experts): + getattr(target_linear, f"bias{idx}").copy_(reference_biases[idx]) + + def packed_grad(linear, name): + if getattr(linear, f"single_grouped_{name}"): + grad = getattr(linear, name).grad + assert grad is not None, f"Grouped {name} parent did not receive a gradient" + return grad.float() + grads = [getattr(linear, f"{name}{idx}").grad for idx in range(self.num_experts)] + assert all(grad is not None for grad in grads) + return torch.stack(grads).float() + + torch.manual_seed(1234) + reference = build_experts(False, False) + torch.manual_seed(5678) + target = build_experts(single_grouped_weight, single_grouped_bias) + copy_linear_params(reference.linear_fc1, target.linear_fc1) + copy_linear_params(reference.linear_fc2, target.linear_fc2) + + tokens_per_expert = torch.tensor([256, 256], dtype=torch.int64, device="cuda") + num_tokens = int(tokens_per_expert.sum().item()) + base_input = 0.1 * torch.randn( + num_tokens, self.hidden_size, dtype=torch.bfloat16, device="cuda" + ) + base_probs = torch.rand(num_tokens, dtype=torch.bfloat16, device="cuda") + grad_output = 0.1 * torch.randn( + num_tokens, self.hidden_size, dtype=torch.bfloat16, device="cuda" + ) + + reference_input = base_input.detach().clone().requires_grad_(True) + reference_probs = base_probs.detach().clone().requires_grad_(True) + reference_output, _ = reference(reference_input, tokens_per_expert, reference_probs) + reference_output.backward(grad_output) + + target_input = base_input.detach().clone().requires_grad_(True) + target_probs = base_probs.detach().clone().requires_grad_(True) + target_output, _ = target(target_input, tokens_per_expert, target_probs) + target_output.backward(grad_output) + + tolerances = {"rtol": 1e-2, "atol": 1e-2} + torch.testing.assert_close(target_output, reference_output, **tolerances) + torch.testing.assert_close(target_input.grad, reference_input.grad, **tolerances) + torch.testing.assert_close(target_probs.grad, reference_probs.grad, **tolerances) + + for target_linear, reference_linear in ( + (target.linear_fc1, reference.linear_fc1), + (target.linear_fc2, reference.linear_fc2), + ): + torch.testing.assert_close( + packed_grad(target_linear, "weight"), + packed_grad(reference_linear, "weight"), + **tolerances, + ) + torch.testing.assert_close( + packed_grad(target_linear, "bias"), + packed_grad(reference_linear, "bias"), + **tolerances, + ) + + if use_op_fuser: + ops = target._fused_ops[0] + fc1_op = ops[0] + fc2_op = ops[2] + # The fused-op shells must register MCore's original parent parameters, not + # detached copies or member views, so autograd and the optimizer update the same objects. + if single_grouped_weight: + assert fc1_op.weight is target.linear_fc1.weight + assert fc2_op.weight is target.linear_fc2.weight + if single_grouped_bias: + assert fc1_op.bias is target.linear_fc1.bias + assert fc2_op.bias is target.linear_fc2.bias + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.internal def test_gpu_fused_path_loss_decreases(self): diff --git a/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py b/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py new file mode 100644 index 00000000000..0e60e64c809 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py @@ -0,0 +1,599 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Distributed MoE coverage for the TE grouped-tensor expert path. + +The dispatchers place expert-token padding at different boundaries: + +* All-to-All returns unpadded expert segments and TEGroupedMLP pads them before FC1. +* DeepEP communicates first, then its local fused permutation pads expert segments. +* HybridEP fuses communication, permutation, and expert-segment padding. +* NCCL-EP returns aligned expert segments from fused dispatch, like HybridEP. Its non-op-fuser + grouped-tensor integration is not enabled yet, so its parity and lifecycle cases remain skipped. + +The numerical tests compare each grouped-tensor configuration with the old discrete-parameter, +CPU-split path on the same dispatcher. The lifecycle tests inspect the actual expert-compute +boundary to ensure padding rows are zero and the dispatcher-specific inverse removes them. +""" + +import inspect +import os +from typing import Dict + +import pytest +import torch +import torch.nn.functional as F + +from megatron.core import config as mcore_config +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) +from megatron.core.transformer.module import Float16Module +from megatron.core.transformer.moe.fused_a2a import ( + HAVE_DEEP_EP, + HAVE_HYBRIDEP, + reset_hybrid_ep_buffer, +) +from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules +from megatron.core.transformer.moe.moe_utils import fused_permute_and_pad_with_probs +from megatron.core.transformer.spec_utils import get_submodules +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.initialize import _set_random_seed +from tests.unit_tests.test_utilities import Utils + +pytestmark = [ + pytest.mark.internal, + pytest.mark.launch_on_gb200, + pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available"), +] + + +_ALIGN_SIZE = 256 +_HIDDEN_SIZE = 256 +_MOE_FFN_HIDDEN_SIZE = 256 +_NUM_LOCAL_EXPERTS = 2 +_NUM_LOCAL_TOKENS = 128 +_TOLERANCES = {"rtol": 1e-2, "atol": 1e-2} +_NCCL_EP_GROUPED_TENSOR_UNSUPPORTED_REASON = ( + "NCCL-EP support for the non-op-fuser grouped-tensor expert path is not implemented yet" +) + +_PARAMETER_LAYOUTS = ( + pytest.param(False, False, False, id="discrete-weight-no-bias"), + pytest.param(True, False, False, id="single-weight-no-bias"), + pytest.param(False, True, False, id="discrete-weight-discrete-bias"), + pytest.param(True, True, False, id="single-weight-discrete-bias"), + pytest.param(False, True, True, id="discrete-weight-single-bias"), + pytest.param(True, True, True, id="single-weight-single-bias"), +) + + +def _require_test_environment(dispatcher: str) -> int: + """Validate runtime support and return the EP world size used by the test.""" + world_size = torch.distributed.get_world_size() + if world_size < 2: + pytest.skip("DeepEP/HybridEP parity requires at least two distributed ranks") + if dispatcher == "deepep" and not HAVE_DEEP_EP: + pytest.skip("DeepEP is not available") + if dispatcher == "hybridep" and not HAVE_HYBRIDEP: + pytest.skip("HybridEP is not available") + if dispatcher == "deepep" and fused_permute_and_pad_with_probs is None: + pytest.skip("DeepEP grouped-tensor padding requires TE fused permute-and-pad") + + try: + from transformer_engine.pytorch.module import GroupedLinear + except ImportError: + pytest.skip("Transformer Engine GroupedLinear is not available") + parameters = inspect.signature(GroupedLinear.__init__).parameters + if "use_grouped_tensor" not in parameters or "single_grouped_bias" not in parameters: + pytest.skip("Installed TE lacks native grouped-tensor parameter support") + return world_size + + +def _dispatcher_options(dispatcher: str) -> Dict[str, object]: + """Return the production padding configuration for one dispatcher.""" + if dispatcher == "alltoall": + return { + "moe_token_dispatcher_type": "alltoall", + "moe_flex_dispatcher_backend": None, + "moe_permute_fusion": True, + } + if dispatcher == "deepep": + # DeepEP communicates first. Its fused local permutation then groups and pads tokens. + return { + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "deepep", + "moe_permute_fusion": True, + } + if dispatcher == "hybridep": + # HybridEP owns permutation and padding inside its fused dispatch/combine kernels. + return { + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "hybridep", + "moe_permute_fusion": True, + } + if dispatcher == "ncclep": + # NCCL-EP dispatch packs aligned expert segments, but the module grouped-tensor path is + # intentionally config-rejected until its end-to-end numerics and padding are validated. + return { + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "ncclep", + "moe_permute_fusion": True, + "moe_expert_rank_capacity_factor": 8.0, + # The lifecycle assertions inspect the dynamically narrowed expert buffer. Static + # NCCL-EP intentionally exposes the full receive-capacity buffer instead. + "moe_ncclep_static_shape": False, + } + raise ValueError(f"Unknown dispatcher {dispatcher!r}") + + +def _build_moe_layer( + dispatcher: str, + *, + ep_size: int, + use_grouped_tensor: bool, + single_grouped_weight: bool, + use_bias: bool, + single_grouped_bias: bool, +) -> MoELayer: + """Build a small real TE MoE layer without using the TE operation fuser.""" + options = _dispatcher_options(dispatcher) + transformer_config = TransformerConfig( + num_layers=1, + hidden_size=_HIDDEN_SIZE, + num_attention_heads=8, + num_moe_experts=ep_size * _NUM_LOCAL_EXPERTS, + moe_ffn_hidden_size=_MOE_FFN_HIDDEN_SIZE, + use_cpu_initialization=False, + add_bias_linear=use_bias, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=False, + bias_dropout_fusion=False, + bf16=True, + params_dtype=torch.bfloat16, + moe_router_load_balancing_type="none", + moe_router_topk=2, + moe_aux_loss_coeff=0.0, + moe_router_dtype="fp32", + moe_grouped_gemm=True, + moe_use_grouped_tensor=use_grouped_tensor, + moe_single_grouped_weight=single_grouped_weight, + moe_single_grouped_bias=single_grouped_bias, + use_transformer_engine_op_fuser=False, + tensor_model_parallel_size=1, + expert_model_parallel_size=ep_size, + sequence_parallel=False, + **options, + ) + submodules = get_submodules( + get_gpt_layer_with_transformer_engine_submodules( + num_experts=transformer_config.num_moe_experts, moe_grouped_gemm=True + ).mlp + ) + assert isinstance(submodules, MoESubmodules) + layer = MoELayer(transformer_config, submodules) + layer = Float16Module(layer.config, layer).module + layer.cuda() + layer.set_layer_number(0) + return layer + + +def _copy_linear_parameters(reference, target) -> None: + """Copy discrete expert parameters into either a discrete or packed target layout.""" + for parameter_name in ("weight", "bias"): + if parameter_name == "bias" and not reference.use_bias: + continue + reference_parts = torch.stack( + [ + getattr(reference, f"{parameter_name}{idx}").detach() + for idx in range(reference.num_gemms) + ] + ) + target_is_grouped = getattr(target, f"single_grouped_{parameter_name}") + if target_is_grouped: + grouped_parameter = getattr(target, parameter_name) + grouped_parameter.rowwise_data.view_as(reference_parts).copy_(reference_parts) + else: + for idx, part in enumerate(reference_parts): + getattr(target, f"{parameter_name}{idx}").copy_(part) + + +@torch.no_grad() +def _copy_layer_parameters(reference: MoELayer, target: MoELayer) -> None: + """Give reference and target identical router and expert parameters.""" + target_parameters = dict(target.named_parameters()) + for name, parameter in reference.named_parameters(): + if not name.startswith("experts."): + target_parameters[name].copy_(parameter) + _copy_linear_parameters(reference.experts.linear_fc1, target.experts.linear_fc1) + _copy_linear_parameters(reference.experts.linear_fc2, target.experts.linear_fc2) + + +def _canonical_gradient(linear, parameter_name: str) -> torch.Tensor: + """Return expert gradients as one [experts, ...] FP32 tensor for either layout.""" + if getattr(linear, f"single_grouped_{parameter_name}"): + gradient = getattr(linear, parameter_name).grad + assert gradient is not None, f"Grouped {parameter_name} parent has no gradient" + return gradient.reshape(linear.num_gemms, -1).float() + + gradients = [getattr(linear, f"{parameter_name}{idx}").grad for idx in range(linear.num_gemms)] + assert all(gradient is not None for gradient in gradients) + return torch.stack([gradient.reshape(-1) for gradient in gradients]).float() + + +def _run_forward_backward( + layer: MoELayer, base_input: torch.Tensor, grad_output: torch.Tensor +) -> Dict[str, torch.Tensor]: + """Run one MoE step and collect values sensitive to dispatch and parameter layout.""" + layer.zero_grad(set_to_none=True) + hidden_states = base_input.detach().clone().requires_grad_(True) + output, _ = layer(hidden_states) + output.backward(grad_output) + + result = { + "output": output.detach(), + "input_grad": hidden_states.grad.detach(), + "router_grad": layer.router.weight.grad.detach(), + "fc1_weight_grad": _canonical_gradient(layer.experts.linear_fc1, "weight"), + "fc2_weight_grad": _canonical_gradient(layer.experts.linear_fc2, "weight"), + } + if layer.config.add_bias_linear: + result["fc1_bias_grad"] = _canonical_gradient(layer.experts.linear_fc1, "bias") + result["fc2_bias_grad"] = _canonical_gradient(layer.experts.linear_fc2, "bias") + return result + + +def _run_numerical_parity_case( + dispatcher: str, *, single_grouped_weight: bool, use_bias: bool, single_grouped_bias: bool +) -> None: + """Compare grouped-tensor execution with the old path on the same dispatcher.""" + ep_size = _require_test_environment(dispatcher) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, expert_model_parallel_size=ep_size + ) + mcore_config.ENABLE_EXPERIMENTAL = True + + _set_random_seed(seed_=1234, data_parallel_random_init=False) + reference = _build_moe_layer( + dispatcher, + ep_size=ep_size, + use_grouped_tensor=False, + single_grouped_weight=False, + use_bias=use_bias, + single_grouped_bias=False, + ) + target = _build_moe_layer( + dispatcher, + ep_size=ep_size, + use_grouped_tensor=True, + single_grouped_weight=single_grouped_weight, + use_bias=use_bias, + single_grouped_bias=single_grouped_bias, + ) + _copy_layer_parameters(reference, target) + + base_input = torch.randn( + _NUM_LOCAL_TOKENS, 1, _HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda" + ) + grad_output = torch.randn_like(base_input) + + reference_result = _run_forward_backward(reference, base_input, grad_output) + target_result = _run_forward_backward(target, base_input, grad_output) + + assert target.experts._use_grouped_tensor + assert not reference.experts._use_grouped_tensor + assert reference_result.keys() == target_result.keys() + for name in reference_result: + torch.testing.assert_close( + target_result[name], + reference_result[name], + msg=lambda message, value=name: f"{dispatcher} {value} mismatch: {message}", + **_TOLERANCES, + ) + + +def _make_padding_mask( + real_tokens_per_expert: torch.Tensor, padded_tokens_per_expert: torch.Tensor +) -> torch.Tensor: + """Build a mask for the padding suffix in every expert-major segment. + + For example, using an alignment of four for readability: + + ``real_tokens_per_expert = tensor([3, 2])`` + ``padded_tokens_per_expert = tensor([4, 4])`` + + The packed expert-major rows are ``[e0 real x3][e0 pad x1][e1 real x2][e1 pad x2]``, + so this function returns ``[F, F, F, T, F, F, T, T]``. Indexing the packed hidden states or + probabilities with that mask selects only the synthetic rows that must contain exact zeros. + """ + masks = [] + for real_count, padded_count in zip( + real_tokens_per_expert.cpu().tolist(), padded_tokens_per_expert.cpu().tolist() + ): + # Each expert contributes a False prefix for real rows followed by a True padding suffix. + masks.append(torch.zeros(real_count, dtype=torch.bool, device="cuda")) + masks.append(torch.ones(padded_count - real_count, dtype=torch.bool, device="cuda")) + return torch.cat(masks) + + +def _infer_real_tokens_per_expert( + padded_probs: torch.Tensor, padded_tokens_per_expert: torch.Tensor +) -> torch.Tensor: + """Infer real expert counts from the zero suffix in every padded probability segment. + + This dropless test uses FP32 softmax router probabilities, so every routed token has a + nonzero probability. All padding implementations use exact zeros and append them after the + real rows in each expert-major segment. This gives one dispatcher-independent source of truth + without inspecting backend-specific routing metadata. + """ + real_counts = [] + offset = 0 + for padded_count in padded_tokens_per_expert.cpu().tolist(): + segment = padded_probs[offset : offset + padded_count] + nonzero_rows = segment != 0 + real_count = int(nonzero_rows.sum().item()) + + # Padding must be one contiguous suffix. Interspersed zero rows would preserve the total + # nonzero count while still violating the expert-major layout expected by grouped GEMM. + assert torch.all(nonzero_rows[:real_count]) + assert not torch.any(nonzero_rows[real_count:]) + real_counts.append(real_count) + offset += padded_count + + assert offset == padded_probs.numel() + return torch.tensor(real_counts, dtype=torch.int64, device=padded_probs.device) + + +def _install_padding_probes(layer: MoELayer, monkeypatch): + """Observe the tensors crossing each padding boundary without changing execution. + + ``register_forward_pre_hook`` runs immediately before the selected module's ``forward``. + The hook receives the module and the tuple of positional arguments that forward is about to + consume. Returning ``None`` leaves those arguments unchanged, so these hooks are read-only + probes rather than replacements for any production operation. + + There are two relevant boundaries. ``layer.experts`` sees what the token dispatcher hands to + TEGroupedMLP, while ``linear_fc1`` sees the final padded tensors and CUDA expert counts that + TE's grouped-tensor GEMM actually consumes. They are different for AllToAll, where TEGroupedMLP + owns padding, but identical for DeepEP and HybridEP, whose fused dispatch paths already pad. + """ + captured = {"padding_calls": 0, "unpadding_calls": 0} + + def capture_dispatcher_input(_module, args): + # TEGroupedMLP.forward(hidden, tokens_per_expert, permuted_probs) is about to run. + # DeepEP and HybridEP have already padded at this boundary, so preserve their router + # probabilities for the generic real-row inference below. Detach so the test does not + # retain the graph, and clone in case downstream computation reuses the input storage. + captured["dispatcher_probs"] = args[2].detach().clone() + + def capture_fc1_input(_module, args): + # GroupedLinear.forward(hidden, m_splits, ...) is about to run. This is the authoritative + # view of the rows and device-side split tensor presented to the grouped GEMM. + captured["padded_hidden"] = args[0].detach().clone() + captured["padded_counts"] = args[1].detach().clone() + + # Pre-hooks observe module inputs before either module can transform them. The returned hook + # handles need not be retained because each test owns this layer and executes one forward. + layer.experts.register_forward_pre_hook(capture_dispatcher_input) + layer.experts.linear_fc1.register_forward_pre_hook(capture_fc1_input) + + def capture_padding(_module, args, output): + # quantization_padding is called once for hidden states and once for router probabilities + # in the AllToAll path. A forward hook is used here because the padded tensor is its output. + captured["padding_calls"] += 1 + padded_tensor = output[0] + # Probability padding receives [tokens, 1], whereas hidden-state padding receives + # [tokens, hidden_size]. Preserve the padded probabilities for an exact-zero assertion. + if args[0].shape[-1] == 1: + captured["padded_probs"] = padded_tensor.detach().clone().reshape(-1) + + def capture_unpadding(_module, _args, output): + # AllToAll uses TEGroupedMLP's explicit unpadding module after expert compute. Recording + # the call proves that its locally inserted padding reaches the matching removal path. + captured["unpadding_calls"] += 1 + + layer.experts.quantization_padding.register_forward_hook(capture_padding) + layer.experts.quantization_unpadding.register_forward_hook(capture_unpadding) + + # All dispatchers share this final restoration API. + original_combine_postprocess = layer.token_dispatcher.combine_postprocess + + def capture_combine_postprocess(hidden_states, *args, **kwargs): + output = original_combine_postprocess(hidden_states, *args, **kwargs) + captured["restored_shape"] = output.shape + return output + + monkeypatch.setattr(layer.token_dispatcher, "combine_postprocess", capture_combine_postprocess) + + return captured + + +def _run_padding_lifecycle_case(dispatcher: str, monkeypatch) -> None: + """Verify exact zero padding, 256 alignment, no double-padding, and unpadding. + + This test follows one real MoE forward through dispatch, expert compute, and restoration. It + independently reconstructs the number of real tokens assigned to each local expert, then + compares that metadata with the padded CUDA ``m_splits`` observed directly at FC1. Numerical + parity with the legacy backend is tested separately; this case focuses on padding ownership + and the structural contract required by TE's grouped-tensor kernels. + """ + # EP spans the whole torchrun world so every tested dispatcher performs real communication. + # TP remains one to keep the independently reconstructed local-expert counts unambiguous. + ep_size = _require_test_environment(dispatcher) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, expert_model_parallel_size=ep_size + ) + mcore_config.ENABLE_EXPERIMENTAL = True + + # Use the strictest parameter layout. If packed weight and packed bias reach the native + # grouped-tensor path correctly, discrete parameter layouts use the same padding lifecycle. + _set_random_seed(seed_=1357, data_parallel_random_init=False) + layer = _build_moe_layer( + dispatcher, + ep_size=ep_size, + use_grouped_tensor=True, + single_grouped_weight=True, + use_bias=True, + single_grouped_bias=True, + ) + # Install observers before the forward so they capture dispatcher output, FC1 input, and the + # common dispatcher boundary that returns to the original token layout. + captured = _install_padding_probes(layer, monkeypatch) + + torch.manual_seed(9753) + hidden_states = torch.randn( + _NUM_LOCAL_TOKENS, 1, _HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + output, _ = layer(hidden_states) + + padded_counts = captured["padded_counts"] + + # TE's grouped-tensor API requires device-resident int64 splits. Every expert segment must be + # represented by its m_split, and the physical FC1 input must contain their total. + assert padded_counts.device.type == "cuda" + assert padded_counts.dtype == torch.int64 + assert torch.all(padded_counts % _ALIGN_SIZE == 0) + assert captured["padded_hidden"].shape[0] == padded_counts.sum().item() + + if dispatcher == "alltoall": + # All-to-All returns real expert rows; TEGroupedMLP pads hidden states and probs itself. + assert captured["padding_calls"] == 2 + assert captured["unpadding_calls"] == 1 + else: + # DeepEP and HybridEP already return padded rows. TEGroupedMLP must not pad them again. + assert captured["padding_calls"] == 0 + assert captured["unpadding_calls"] == 0 + # For fused dispatchers, probabilities are already padded when they enter TEGroupedMLP, + # so the experts pre-hook is the correct observation point for the zero check below. + captured["padded_probs"] = captured["dispatcher_probs"].reshape(-1) + + # Router probabilities provide a common representation across all dispatchers: real routed + # rows are nonzero and padded rows are an exact-zero suffix. + real_counts = _infer_real_tokens_per_expert(captured["padded_probs"], padded_counts) + expected_padded_counts = ((real_counts + _ALIGN_SIZE - 1) // _ALIGN_SIZE) * _ALIGN_SIZE + torch.testing.assert_close(padded_counts, expected_padded_counts, rtol=0, atol=0) + + # Expert-major layout is [expert 0 real][expert 0 pad][expert 1 real][expert 1 pad]... + # Build that exact mask and require both hidden states and routing probabilities to use + # numerical zero for every synthetic row. Merely allocating the right shape is insufficient. + padding_mask = _make_padding_mask(real_counts, padded_counts) + # An expert may already have a 256-aligned token count and legitimately need no padding. + # Boolean indexing with an empty mask is valid; otherwise these checks inspect every pad row. + assert not torch.any(captured["padded_hidden"][padding_mask]) + assert not torch.any(captured["padded_probs"][padding_mask]) + + # Regardless of where a backend removes padding, the common dispatcher postprocess contract + # must restore exactly the shape that entered this MoE layer. + assert captured["restored_shape"] == hidden_states.shape + + # The public MoE contract is unchanged by internal alignment. Run backward as a final check + # that padding/unpadding preserved a connected, finite autograd path to the original input. + assert output.shape == hidden_states.shape + output.float().square().mean().backward() + assert hidden_states.grad is not None + assert torch.isfinite(hidden_states.grad).all() + + +class TestGroupedTensorDispatcherNumerics: + """Distributed numerical and padding coverage for grouped-tensor dispatchers.""" + + def setup_method(self, method): + if not torch.distributed.is_available() or Utils.world_size < 2: + pytest.skip("Distributed dispatcher tests must be launched with torchrun") + self._old_single_param_env = os.environ.get("NVTE_GROUPED_LINEAR_SINGLE_PARAM") + self._previous_experimental = mcore_config.ENABLE_EXPERIMENTAL + os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] = "1" + Utils.initialize_distributed() + + def teardown_method(self, method): + try: + mcore_config.ENABLE_EXPERIMENTAL = self._previous_experimental + reset_hybrid_ep_buffer() + Utils.destroy_model_parallel() + finally: + if self._old_single_param_env is None: + os.environ.pop("NVTE_GROUPED_LINEAR_SINGLE_PARAM", None) + else: + os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] = self._old_single_param_env + + @pytest.mark.parametrize( + "single_grouped_weight,use_bias,single_grouped_bias", _PARAMETER_LAYOUTS + ) + @pytest.mark.timeout(180) + def test_alltoall_grouped_tensor_moe_parity( + self, single_grouped_weight, use_bias, single_grouped_bias + ): + """All-to-All grouped-tensor MoE forward/backward matches its legacy expert path.""" + _run_numerical_parity_case( + "alltoall", + single_grouped_weight=single_grouped_weight, + use_bias=use_bias, + single_grouped_bias=single_grouped_bias, + ) + + @pytest.mark.parametrize( + "single_grouped_weight,use_bias,single_grouped_bias", _PARAMETER_LAYOUTS + ) + @pytest.mark.timeout(180) + def test_deepep_grouped_tensor_moe_parity( + self, single_grouped_weight, use_bias, single_grouped_bias + ): + """DeepEP grouped-tensor MoE forward/backward matches its legacy expert path.""" + _run_numerical_parity_case( + "deepep", + single_grouped_weight=single_grouped_weight, + use_bias=use_bias, + single_grouped_bias=single_grouped_bias, + ) + + @pytest.mark.parametrize( + "single_grouped_weight,use_bias,single_grouped_bias", _PARAMETER_LAYOUTS + ) + @pytest.mark.timeout(180) + def test_hybridep_grouped_tensor_moe_parity( + self, single_grouped_weight, use_bias, single_grouped_bias + ): + """HybridEP grouped-tensor MoE forward/backward matches its legacy expert path.""" + _run_numerical_parity_case( + "hybridep", + single_grouped_weight=single_grouped_weight, + use_bias=use_bias, + single_grouped_bias=single_grouped_bias, + ) + + @pytest.mark.skip(reason=_NCCL_EP_GROUPED_TENSOR_UNSUPPORTED_REASON) + @pytest.mark.parametrize( + "single_grouped_weight,use_bias,single_grouped_bias", _PARAMETER_LAYOUTS + ) + @pytest.mark.timeout(180) + def test_ncclep_grouped_tensor_moe_parity( + self, single_grouped_weight, use_bias, single_grouped_bias + ): + """NCCL-EP grouped-tensor parity coverage reserved for future enablement.""" + _run_numerical_parity_case( + "ncclep", + single_grouped_weight=single_grouped_weight, + use_bias=use_bias, + single_grouped_bias=single_grouped_bias, + ) + + @pytest.mark.timeout(180) + def test_alltoall_grouped_tensor_padding_lifecycle(self, monkeypatch): + """All-to-All explicitly pads in TEGroupedMLP and removes it before combine.""" + _run_padding_lifecycle_case("alltoall", monkeypatch) + + @pytest.mark.timeout(180) + def test_deepep_grouped_tensor_padding_lifecycle(self, monkeypatch): + """DeepEP fused local permutation pads, and local unpermute removes those rows.""" + _run_padding_lifecycle_case("deepep", monkeypatch) + + @pytest.mark.timeout(180) + def test_hybridep_grouped_tensor_padding_lifecycle(self, monkeypatch): + """HybridEP fused dispatch pads, and fused combine returns the original token shape.""" + _run_padding_lifecycle_case("hybridep", monkeypatch) + + @pytest.mark.skip(reason=_NCCL_EP_GROUPED_TENSOR_UNSUPPORTED_REASON) + @pytest.mark.timeout(180) + def test_ncclep_grouped_tensor_padding_lifecycle(self, monkeypatch): + """NCCL-EP padding lifecycle coverage reserved for future enablement.""" + _run_padding_lifecycle_case("ncclep", monkeypatch) diff --git a/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py b/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py index f5a4afaeffc..711405a5f8a 100644 --- a/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py +++ b/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py @@ -54,8 +54,12 @@ _TE_GROUPED_LINEAR_SUPPORTS_SINGLE_PARAM = ( "single_grouped_weight" in inspect.signature(TEGroupedLinear.__init__).parameters ) + _TE_GROUPED_LINEAR_SUPPORTS_USE_GROUPED_TENSOR = ( + "use_grouped_tensor" in inspect.signature(TEGroupedLinear.__init__).parameters + ) except (ImportError, AttributeError): _TE_GROUPED_LINEAR_SUPPORTS_SINGLE_PARAM = False + _TE_GROUPED_LINEAR_SUPPORTS_USE_GROUPED_TENSOR = False pytestmark = [ pytest.mark.internal, @@ -67,6 +71,10 @@ not _TE_GROUPED_LINEAR_SUPPORTS_SINGLE_PARAM, reason="Installed TE GroupedLinear does not expose single_grouped_weight", ), + pytest.mark.skipif( + not _TE_GROUPED_LINEAR_SUPPORTS_USE_GROUPED_TENSOR, + reason="Installed TE GroupedLinear does not expose use_grouped_tensor", + ), ] @@ -193,6 +201,7 @@ def create_test_args( args.num_experts = 2 args.moe_layer_freq = 1 args.moe_grouped_gemm = True + args.moe_use_grouped_tensor = True args.moe_single_grouped_weight = single_weight args.moe_token_dispatcher_type = "alltoall" args.moe_router_topk = 1 @@ -682,7 +691,22 @@ def test_mxfp8_single_weight_torch_dist_checkpoint_matches_discrete_baseline( self.assert_all_ranks_passed(local_passed, local_error) - @pytest.mark.parametrize("precision", ["bf16", "mxfp8", "nvfp4"]) + @pytest.mark.parametrize( + "precision", + [ + "bf16", + "mxfp8", + pytest.param( + "nvfp4", + marks=pytest.mark.skip( + reason=( + "NVFP4 single grouped weights are not supported by the " + "TransformerEngine native grouped-tensor path yet." + ) + ), + ), + ], + ) @pytest.mark.parametrize("gradient_accumulation_fusion", [False, True]) def test_single_grouped_weight_parity_with_primary_param_gather( self, precision, gradient_accumulation_fusion @@ -710,17 +734,18 @@ def test_single_grouped_weight_parity_without_primary_param_gather( use_transformer_engine_op_fuser=True, ) - def test_single_grouped_weight_parity_module_grouped_linear(self): - """Single grouped weights require the TE op-fuser execution path.""" - args = self.create_test_args( - precision="bf16", - primary_param_gather=False, - single_weight=True, - gradient_accumulation_fusion=False, + @pytest.mark.parametrize( + "precision,primary_param_gather", [("bf16", False), ("mxfp8", False), ("mxfp8", True)] + ) + @pytest.mark.parametrize("gradient_accumulation_fusion", [False, True]) + def test_single_grouped_weight_parity_module_grouped_linear( + self, precision, primary_param_gather, gradient_accumulation_fusion + ): + """Compare native TE GroupedLinear single and discrete parameter layouts.""" + _skip_if_unsupported(precision) + self.run_parity_case( + precision=precision, + primary_param_gather=primary_param_gather, + gradient_accumulation_fusion=gradient_accumulation_fusion, use_transformer_engine_op_fuser=False, ) - with pytest.raises( - ValueError, - match="moe_single_grouped_weight requires use_transformer_engine_op_fuser=True", - ): - core_transformer_config_from_args(args) From 4347f363f6b9e46a9b2a35e6c1520ea65903270f Mon Sep 17 00:00:00 2001 From: hongbinl Date: Mon, 24 Aug 2026 18:15:54 -0700 Subject: [PATCH 02/10] feat: support paged stash with GroupedTensor experts Signed-off-by: hongbinl --- docs/user-guide/features/paged_stash.md | 26 +- megatron/core/fusions/fused_bias_geglu.py | 15 +- megatron/core/fusions/fused_bias_swiglu.py | 16 +- .../fusions/fused_weighted_squared_relu.py | 12 +- megatron/core/transformer/moe/experts.py | 232 +++++++++++------- .../transformer/moe/test_grouped_mlp.py | 77 ++++++ .../transformer/moe/test_paged_stashing.py | 78 +++++- 7 files changed, 355 insertions(+), 101 deletions(-) diff --git a/docs/user-guide/features/paged_stash.md b/docs/user-guide/features/paged_stash.md index b5b97144905..56b7db1233d 100644 --- a/docs/user-guide/features/paged_stash.md +++ b/docs/user-guide/features/paged_stash.md @@ -13,7 +13,7 @@ **Paged stash** = **sync-free** expert execution + **paged stashing** (packing routed-expert activations for backward into paged buffers). -**Sync-free:** `--moe-flex-dispatcher-backend hybridep`, `--use-transformer-engine-op-fuser`, and `--moe-expert-rank-capacity-factor` pre-size dispatch and fused grouped expert buffers from a user-controlled capacity, avoiding a per-step device query / realloc loop for buffer sizing. +**Sync-free:** `--moe-flex-dispatcher-backend hybridep` and `--moe-expert-rank-capacity-factor` pre-size dispatch and grouped expert buffers from a user-controlled capacity, avoiding a per-step device query / realloc loop for buffer sizing. Expert compute can use either `--use-transformer-engine-op-fuser` or the device-initiated Transformer Engine GroupedTensor path (`--moe-grouped-gemm --moe-use-grouped-tensor`). **Paged stashing:** `--moe-paged-stash` stores those activations in paged CUDA buffers (optional pinned host spill). It helps save activation memory; sync-free still works without it, at the cost of higher activation memory use. @@ -21,20 +21,38 @@ Whenever `moe_expert_rank_capacity_factor` is set, a **runner** wraps forward-ba ## Prerequisites -HybridEP + TE fused grouped experts are required whenever `moe_expert_rank_capacity_factor` is set. With `moe_paged_stash` enabled: capacity factor must be set; no `cpu_offloading`; `offload_modules` must not include `expert_fc1`, `moe_act`, or `fused_group_mlp`. The runner is active whenever capacity factor is set (even without `--moe-paged-stash`) for over-budget reruns; stash overflow is checked only when paged stashing is on. +HybridEP and TE grouped experts are required whenever `moe_expert_rank_capacity_factor` is set. The non-op-fuser path requires a Transformer Engine version whose GroupedLinear marks saved GroupedTensor activation buffers for paged stashing. With `moe_paged_stash` enabled: capacity factor must be set; no `cpu_offloading`; `offload_modules` must not include `expert_fc1`, `moe_act`, or `fused_group_mlp`. The runner is active whenever capacity factor is set (even without `--moe-paged-stash`) for over-budget reruns; stash overflow is checked only when paged stashing is on. ## Configuration ```bash -# Sync-free +# Common static-budget configuration +--moe-token-dispatcher-type flex --moe-flex-dispatcher-backend hybridep ---use-transformer-engine-op-fuser --moe-expert-rank-capacity-factor # Paged stashing (to avoid memory waste due to fragmentation) --moe-paged-stash + +# Choose one expert-compute path: + +# A. TE operation fuser (used by the full-iteration CUDA graph + CuTe DSL route) +--use-transformer-engine-op-fuser + +# B. Device-initiated GroupedLinear, without the operation fuser +--moe-grouped-gemm +--moe-use-grouped-tensor ``` +Path B removes host-device synchronization from grouped GEMM split metadata, but FC1, activation, +and FC2 remain separate launches. Without a full-iteration CUDA graph it can therefore retain +significant CPU launch overhead even though the expert path is host-device sync-free. + +In this context, sync-free refers to the steady-state expert data path. The initial paged-stash +capture performs host reads, and the runner reads reduced overflow/over-budget state at the end of +a pass to decide whether to rerun; it does not imply that the complete iteration has no CPU-GPU +synchronization at all. + ## Tuning (paged stashing only) ```bash diff --git a/megatron/core/fusions/fused_bias_geglu.py b/megatron/core/fusions/fused_bias_geglu.py index 7a7fbe7f9ec..9c158badf86 100644 --- a/megatron/core/fusions/fused_bias_geglu.py +++ b/megatron/core/fusions/fused_bias_geglu.py @@ -1,9 +1,17 @@ -# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import torch from megatron.core.jit import jit_fuser + +def _propagate_paged_stash_marker(source, target): + """Preserve Megatron's dynamic-activation marker across view/cast operations.""" + if hasattr(source, "grouped_tensor_scale_inv"): + setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + return target + + ###### BIAS GELU FUSION/ NO AUTOGRAD ################ # 1/sqrt(2*pi)-> 0.3989423 # 1/sqrt(2) -> 0.70710678 @@ -324,6 +332,7 @@ def forward( torch.Tensor: Output tensor of shape [N, H] after weighted Quick-GEGLU. """ input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input + _propagate_paged_stash_marker(input, input_for_backward) ctx.save_for_backward(input_for_backward, weights, linear_offset) ctx.ori_input_dtype = input.dtype ctx.fp8_input_store = fp8_input_store @@ -374,6 +383,7 @@ def forward( """ # Optionally store the input in FP8 for memory savings. input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input + _propagate_paged_stash_marker(input, input_for_backward) # Save tensors for backward. ctx.save_for_backward(input_for_backward, bias, weights, linear_offset) @@ -420,6 +430,7 @@ def weighted_bias_quick_geglu_impl( output: [num_selected_experts * seq_len, hidden_size] """ ori_shape = input.shape + paged_stash_source = input assert len(ori_shape) in [2, 3] if clamp_value is not None: x_glu, x_linear = input.chunk(2, -1) @@ -430,7 +441,7 @@ def weighted_bias_quick_geglu_impl( ), -1, ) - input = input.view(-1, ori_shape[-1]) + input = _propagate_paged_stash_marker(paged_stash_source, input.view(-1, ori_shape[-1])) linear_offset = torch.tensor(linear_offset, dtype=input.dtype, device=input.device) if bias is not None: output = WeightedBiasQuickGeGLUFunction.apply( diff --git a/megatron/core/fusions/fused_bias_swiglu.py b/megatron/core/fusions/fused_bias_swiglu.py index 4cd678be816..939f25864b6 100644 --- a/megatron/core/fusions/fused_bias_swiglu.py +++ b/megatron/core/fusions/fused_bias_swiglu.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # pylint: disable=missing-function-docstring, missing-class-docstring @@ -12,6 +12,13 @@ ###### BIAS SWIGLU FUSION/ NO AUTOGRAD ################ +def _propagate_paged_stash_marker(source, target): + """Preserve Megatron's dynamic-activation marker across view/cast operations.""" + if hasattr(source, "grouped_tensor_scale_inv"): + setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + return target + + @jit_fuser def swiglu(y): """Performs SwiGLU (Swish-Gated Linear Unit) activation function. @@ -178,6 +185,7 @@ def forward(ctx, input, bias, fp8_input_store, cpu_offload_input, clamp_value): torch.Tensor: Result of applying bias addition followed by SwiGLU activation. """ input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input + _propagate_paged_stash_marker(input, input_for_backward) if cpu_offload_input: input_for_backward.activation_offloading = True bias.activation_offloading = True @@ -236,6 +244,7 @@ def forward(ctx, input, fp8_input_store, cpu_offload_input, clamp_value): torch.Tensor: Result of applying SwiGLU activation. """ input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input + _propagate_paged_stash_marker(input, input_for_backward) if cpu_offload_input: input_for_backward.activation_offloading = True ctx.save_for_backward(input_for_backward) @@ -275,6 +284,7 @@ class WeightedSwiGLUFunction(torch.autograd.Function): @staticmethod def forward(ctx, input, weights, fp8_input_store, clamp_value): input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input + _propagate_paged_stash_marker(input, input_for_backward) ctx.save_for_backward(input_for_backward, weights) ctx.ori_input_dtype = input.dtype ctx.fp8_input_store = fp8_input_store @@ -322,7 +332,7 @@ def bias_swiglu_impl(input, bias, fp8_input_store=False, cpu_offload_input=False """ ori_shape = input.shape assert len(ori_shape) in [2, 3] - input = input.view(-1, ori_shape[-1]) + input = _propagate_paged_stash_marker(input, input.view(-1, ori_shape[-1])) if bias is not None: output = BiasSwiGLUFunction.apply( input, bias, fp8_input_store, cpu_offload_input, clamp_value @@ -339,7 +349,7 @@ def weighted_bias_swiglu_impl(input, bias, weights, fp8_input_store=False, clamp """ ori_shape = input.shape assert len(ori_shape) in [2, 3] - input = input.view(-1, ori_shape[-1]) + input = _propagate_paged_stash_marker(input, input.view(-1, ori_shape[-1])) if bias is not None: raise NotImplementedError("Bias is not supported for weighted swiglu fusion") else: diff --git a/megatron/core/fusions/fused_weighted_squared_relu.py b/megatron/core/fusions/fused_weighted_squared_relu.py index 02dabc14c3b..484e1edef46 100644 --- a/megatron/core/fusions/fused_weighted_squared_relu.py +++ b/megatron/core/fusions/fused_weighted_squared_relu.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import torch import torch.nn.functional as F @@ -7,6 +7,14 @@ from megatron.core.jit import jit_fuser from megatron.core.utils import nvtx_decorator + +def _propagate_paged_stash_marker(source, target): + """Preserve Megatron's dynamic-activation marker across view operations.""" + if hasattr(source, "grouped_tensor_scale_inv"): + setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + return target + + ###################### WEIGHTED SQUARED ReLU FUSION ###################### @@ -103,7 +111,7 @@ def weighted_squared_relu_impl(input: torch.Tensor, weights: torch.Tensor) -> to """ ori_shape = input.shape assert len(ori_shape) in [2, 3] - input = input.view(-1, ori_shape[-1]) + input = _propagate_paged_stash_marker(input, input.view(-1, ori_shape[-1])) output = WeightedSquaredReLUFunction.apply(input, weights) diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index cbd9f10d81c..b0556fc04d9 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -312,6 +312,12 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe output_dtype = intermediate_parallel.dtype flat_output = intermediate_parallel.view(-1, hidden_size).float() flat_probs = permuted_probs.reshape(-1, 1).float() + paged_stash_marked = hasattr(intermediate_parallel, "grouped_tensor_scale_inv") or hasattr( + permuted_probs, "grouped_tensor_scale_inv" + ) + if paged_stash_marked: + setattr(flat_output, "grouped_tensor_scale_inv", False) + setattr(flat_probs, "grouped_tensor_scale_inv", False) if tokens_per_expert.device != packed_bias.device: raise ValueError("Packed MoE bias and tokens_per_expert must be on the same device.") @@ -329,6 +335,8 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe bias_per_token = torch.repeat_interleave( packed_bias.float(), tokens_per_expert, dim=0, output_size=flat_output.size(0) ) + if paged_stash_marked: + setattr(bias_per_token, "grouped_tensor_scale_inv", False) return (flat_output + bias_per_token * flat_probs).view(shape).to(output_dtype) @staticmethod @@ -689,6 +697,43 @@ def forward_post_hook(_module, _inputs, output): return forward_post_hook + def _start_paged_stash_group( + self, permuted_local_hidden_states: torch.Tensor, tokens_per_expert: torch.Tensor + ) -> tuple[torch.Tensor, object]: + """Start the grouped-MLP paged-stash scope when it is enabled.""" + if not self.config.moe_paged_stash: + return permuted_local_hidden_states, nullcontext() + + permuted_local_hidden_states = paged_stash_group_start(permuted_local_hidden_states) + max_num_tokens = permuted_local_hidden_states.shape[0] + # Average/expected tokens is a pre-padding estimate used by paged stashing heuristics. + # moe_expert_rank_capacity_factor is required when moe_paged_stash is enabled. + cap_factor = self.config.moe_expert_rank_capacity_factor + avg_num_tokens = ( + int(max_num_tokens // cap_factor) if cap_factor is not None and cap_factor > 0 else None + ) + stash_context = get_paged_stash_context( + name="grouped_mlp", + max_num_tokens=max_num_tokens, + num_tokens_tensor=tokens_per_expert.sum(), + avg_num_tokens=avg_num_tokens, + ) + return permuted_local_hidden_states, stash_context + + def _commit_paged_stash_group(self, output: torch.Tensor) -> torch.Tensor: + """Commit the grouped-MLP paged-stash scope when it is enabled.""" + if self.config.moe_paged_stash: + output = paged_stash_group_commit(output, name="grouped_mlp") + return output + + def _mark_paged_stash_tensors(self, *tensors: Optional[torch.Tensor]) -> None: + """Mark dynamic unfused activations for the paged-stash saved-tensor hook.""" + if not self.config.moe_paged_stash: + return + for tensor in tensors: + if tensor is not None: + setattr(tensor, "grouped_tensor_scale_inv", False) + def _fused_forward( self, permuted_local_hidden_states: torch.Tensor, @@ -752,25 +797,9 @@ def _fused_forward( ) # if the number of tokens is 0, pad the hidden states to 256 - if self.config.moe_paged_stash: - permuted_local_hidden_states = paged_stash_group_start(permuted_local_hidden_states) - max_num_tokens = permuted_local_hidden_states.shape[0] - # Average/expected tokens is a pre-padding estimate used by paged stashing heuristics. - # moe_expert_rank_capacity_factor is required when moe_paged_stash is enabled. - cap_factor = self.config.moe_expert_rank_capacity_factor - avg_num_tokens = ( - int(max_num_tokens // cap_factor) - if cap_factor is not None and cap_factor > 0 - else None - ) - stash_context = get_paged_stash_context( - name="grouped_mlp", - max_num_tokens=max_num_tokens, - num_tokens_tensor=tokens_per_expert.sum(), - avg_num_tokens=avg_num_tokens, - ) - else: - stash_context = nullcontext() + permuted_local_hidden_states, stash_context = self._start_paged_stash_group( + permuted_local_hidden_states, tokens_per_expert + ) fine_grained_activation_offloading = getattr(self, "offload_fused_group_mlp", False) offload_name = "fused_group_mlp" fused_group_mlp_manager = off_interface( @@ -796,9 +825,7 @@ def _fused_forward( # Remove padding if needed if unpadded_tokens_per_expert is not None: output = self.quantization_unpadding(output, unpadded_tokens_per_expert) - if self.config.moe_paged_stash: - output = paged_stash_group_commit(output, name="grouped_mlp") - return output + return self._commit_paged_stash_group(output) @staticmethod def _remove_glu_interleaving(x: torch.Tensor, interleave_size: int) -> torch.Tensor: @@ -809,70 +836,14 @@ def _remove_glu_interleaving(x: torch.Tensor, interleave_size: int) -> torch.Ten x = x.view(shape) return x - def forward( + def _unfused_forward( self, permuted_local_hidden_states: torch.Tensor, - tokens_per_expert: torch.Tensor, + tokens_per_expert: torch.Tensor | list[int], permuted_probs: torch.Tensor, - ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - """Forward of TEGroupedMLP - - Args: - permuted_local_hidden_states (torch.Tensor): The permuted input hidden states of the - local experts. - tokens_per_expert (torch.Tensor): The number of tokens per expert. - permuted_probs (torch.Tensor): The permuted probs of each token produced by the router. - - Return: - output (torch.Tensor): The output of the local experts. - """ - - # Call fused impl if enabled - if self._with_fused_impl: - output = self._fused_forward( - permuted_local_hidden_states, tokens_per_expert, permuted_probs - ) - output_bias = None - return output, output_bias - - # Apply padding if needed - unpadded_tokens_per_expert = None - permuted_probs = permuted_probs.unsqueeze(-1) - # The token buffer may already contain per-expert padding when padding was performed - # before expert compute: - # * router padding modified the routing map before dispatch; - # * HybridEP/NCCL-EP fused padding into dispatch/permute; - # * DeepEP fused padding into its post-communication local permutation. - # In those cases tokens_per_expert already describes the padded expert segments. Running - # Fp8Padding again would change the segment lengths without matching the existing token - # layout, so this module must leave both tensors unchanged. - if skip_routed_expert_padding(self.config): - pass - # Regular AllToAll normally supplies unpadded expert segments and therefore uses this - # explicit fallback. FP8/FP4 need their recipe-specific alignment. MCore currently also - # applies its common aligned-segment contract to the GroupedTensor backend so quantized - # grouped execution receives supported shapes - elif self.config.fp8 or self.config.fp4 or self._use_grouped_tensor: - tokens_per_expert = tokens_per_expert.tolist() - unpadded_tokens_per_expert = tokens_per_expert - permuted_local_hidden_states, tokens_per_expert = self.quantization_padding( - permuted_local_hidden_states, tokens_per_expert - ) - permuted_probs, _ = self.quantization_padding( - permuted_probs, unpadded_tokens_per_expert - ) - - if self._use_grouped_tensor: - if not isinstance(tokens_per_expert, torch.Tensor): - tokens_per_expert = torch.tensor( - tokens_per_expert, dtype=torch.int64, device=permuted_local_hidden_states.device - ) - else: - tokens_per_expert = tokens_per_expert.to( - device=permuted_local_hidden_states.device, dtype=torch.int64, non_blocking=True - ) - elif isinstance(tokens_per_expert, torch.Tensor): - tokens_per_expert = tokens_per_expert.tolist() + ) -> torch.Tensor: + """Run FC1, activation, and FC2 without the TE operation fuser.""" + self._mark_paged_stash_tensors(permuted_local_hidden_states, permuted_probs) if self.config.moe_apply_probs_on_input: assert ( @@ -883,6 +854,7 @@ def forward( permuted_local_hidden_states = permuted_local_hidden_states.to(original_dtype) # Probs already applied, so reset to 1. permuted_probs = torch.ones_like(permuted_probs) + self._mark_paged_stash_tensors(permuted_local_hidden_states, permuted_probs) expert_fc1_manager = off_interface( self.offload_expert_fc1, permuted_local_hidden_states, "expert_fc1" @@ -896,10 +868,12 @@ def forward( forced_released_tensors=[permuted_local_hidden_states], delay_offload=self.config.delay_offload_until_cuda_graph, ) + self._mark_paged_stash_tensors(fc1_output) moe_act_manager = off_interface(self.offload_moe_act, fc1_output, "moe_act") def bias_act_func(intermediate_parallel, bias_parallel, permuted_probs): + self._mark_paged_stash_tensors(intermediate_parallel, permuted_probs) # Whether activation function is interleaved GLU with_glu_interleaving = ( @@ -914,9 +888,11 @@ def bias_act_func(intermediate_parallel, bias_parallel, permuted_probs): intermediate_parallel = self._remove_glu_interleaving( intermediate_parallel, self.config.moe_mlp_glu_interleave_size ) + self._mark_paged_stash_tensors(intermediate_parallel) intermediate_parallel = self.activation_func(intermediate_parallel) if permuted_probs is not None: original_dtype = intermediate_parallel.dtype + self._mark_paged_stash_tensors(intermediate_parallel, permuted_probs) intermediate_parallel = intermediate_parallel * permuted_probs intermediate_parallel = intermediate_parallel.to(original_dtype) elif self.config.bias_activation_fusion and not with_glu_interleaving: @@ -960,22 +936,26 @@ def glu(x): x, self.config.moe_mlp_glu_interleave_size ) x_glu, x_linear = torch.chunk(x, 2, dim=-1) + self._mark_paged_stash_tensors(x_glu, x_linear) if (val := self.config.activation_func_clamp_value) is not None: x_glu = x_glu.clamp(min=None, max=val) x_linear = x_linear.clamp(min=-val, max=val) - return self.config.activation_func(x_glu) * ( - x_linear + self.config.glu_linear_offset - ) + self._mark_paged_stash_tensors(x_glu, x_linear) + x_glu = self.config.activation_func(x_glu) + self._mark_paged_stash_tensors(x_glu) + x_linear = x_linear + self.config.glu_linear_offset + self._mark_paged_stash_tensors(x_linear) + return x_glu * x_linear intermediate_parallel = glu(intermediate_parallel) else: intermediate_parallel = self.activation_func(intermediate_parallel) original_dtype = intermediate_parallel.dtype + self._mark_paged_stash_tensors(intermediate_parallel, permuted_probs) intermediate_parallel = intermediate_parallel * permuted_probs intermediate_parallel = intermediate_parallel.to(original_dtype) return intermediate_parallel - moe_act_manager = off_interface(self.offload_moe_act, fc1_output, "moe_act") if self.activation_recompute: self.activation_checkpoint = tensor_parallel.CheckpointWithoutOutput() with moe_act_manager as fc1_output: @@ -985,7 +965,7 @@ def glu(x): else: with moe_act_manager as fc1_output: bias_act_output = bias_act_func(fc1_output, bias_parallel, permuted_probs) - + self._mark_paged_stash_tensors(bias_act_output) output, output_bias = apply_module(self.linear_fc2)(bias_act_output, tokens_per_expert) if self.activation_recompute: self.activation_checkpoint.discard_output_and_register_recompute(output) @@ -997,12 +977,86 @@ def glu(x): forced_released_tensors=[fc1_output], delay_offload=self.config.delay_offload_until_cuda_graph, ) - output = self._apply_bias(output, output_bias, tokens_per_expert, permuted_probs) + self._mark_paged_stash_tensors(output, permuted_probs) + return self._apply_bias(output, output_bias, tokens_per_expert, permuted_probs) + + def forward( + self, + permuted_local_hidden_states: torch.Tensor, + tokens_per_expert: torch.Tensor, + permuted_probs: torch.Tensor, + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Forward of TEGroupedMLP + + Args: + permuted_local_hidden_states (torch.Tensor): The permuted input hidden states of the + local experts. + tokens_per_expert (torch.Tensor): The number of tokens per expert. + permuted_probs (torch.Tensor): The permuted probs of each token produced by the router. + + Return: + output (torch.Tensor): The output of the local experts. + """ + # Call fused impl if enabled + if self._with_fused_impl: + output = self._fused_forward( + permuted_local_hidden_states, tokens_per_expert, permuted_probs + ) + output_bias = None + return output, output_bias + + # Apply padding if needed + unpadded_tokens_per_expert = None + permuted_probs = permuted_probs.unsqueeze(-1) + # The token buffer may already contain per-expert padding when padding was performed + # before expert compute: + # * router padding modified the routing map before dispatch; + # * HybridEP/NCCL-EP fused padding into dispatch/permute; + # * DeepEP fused padding into its post-communication local permutation. + # In those cases tokens_per_expert already describes the padded expert segments. Running + # Fp8Padding again would change the segment lengths without matching the existing token + # layout, so this module must leave both tensors unchanged. + if skip_routed_expert_padding(self.config): + pass + # Regular AllToAll normally supplies unpadded expert segments and therefore uses this + # explicit fallback. FP8/FP4 need their recipe-specific alignment. MCore currently also + # applies its common aligned-segment contract to the GroupedTensor backend so quantized + # grouped execution receives supported shapes + elif self.config.fp8 or self.config.fp4 or self._use_grouped_tensor: + tokens_per_expert = tokens_per_expert.tolist() + unpadded_tokens_per_expert = tokens_per_expert + permuted_local_hidden_states, tokens_per_expert = self.quantization_padding( + permuted_local_hidden_states, tokens_per_expert + ) + permuted_probs, _ = self.quantization_padding( + permuted_probs, unpadded_tokens_per_expert + ) + + if self._use_grouped_tensor: + if not isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = torch.tensor( + tokens_per_expert, dtype=torch.int64, device=permuted_local_hidden_states.device + ) + else: + tokens_per_expert = tokens_per_expert.to( + device=permuted_local_hidden_states.device, dtype=torch.int64, non_blocking=True + ) + elif isinstance(tokens_per_expert, torch.Tensor): + tokens_per_expert = tokens_per_expert.tolist() + + permuted_local_hidden_states, stash_context = self._start_paged_stash_group( + permuted_local_hidden_states, tokens_per_expert + ) + with stash_context: + output = self._unfused_forward( + permuted_local_hidden_states, tokens_per_expert, permuted_probs + ) # upad and concat the output if unpadded_tokens_per_expert is not None: output = self.quantization_unpadding(output, unpadded_tokens_per_expert) + output = self._commit_paged_stash_group(output) output_bias = None return output, output_bias diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index e56ea17066c..a4e743fff29 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -68,6 +68,26 @@ def test_grouped_tensor_requires_grouped_gemm(): ) +def test_paged_stash_allows_non_fused_grouped_tensor_hybridep(): + config = TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=4, + num_moe_experts=2, + moe_grouped_gemm=True, + moe_use_grouped_tensor=True, + moe_token_dispatcher_type="flex", + moe_flex_dispatcher_backend="hybridep", + moe_expert_rank_capacity_factor=1.5, + moe_paged_stash=True, + use_transformer_engine_op_fuser=False, + ) + + assert config.moe_paged_stash is True + assert config.moe_use_grouped_tensor is True + assert config.use_transformer_engine_op_fuser is False + + def test_remove_glu_interleaving_restores_contiguous_gate_and_linear_halves(): interleaved = torch.tensor([[1, 2, 5, 6, 3, 4, 7, 8], [11, 12, 15, 16, 13, 14, 17, 18]]) expected = torch.tensor([[1, 2, 3, 4, 5, 6, 7, 8], [11, 12, 13, 14, 15, 16, 17, 18]]) @@ -224,6 +244,63 @@ def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): torch.testing.assert_close(fused_ops.args[3], tokens_per_expert) +def test_non_fused_forward_wraps_compute_in_paged_stash_scope(monkeypatch): + events = [] + + class FakeStashContext: + def __enter__(self): + events.append("enter") + + def __exit__(self, exc_type, exc_value, traceback): + events.append("exit") + + module = TEGroupedMLP.__new__(TEGroupedMLP) + module.config = SimpleNamespace( + fp8=False, fp4=False, moe_paged_stash=True, moe_expert_rank_capacity_factor=1.5 + ) + module._with_fused_impl = False + module._use_grouped_tensor = True + + monkeypatch.setattr(experts_module, "skip_routed_expert_padding", lambda _config: True) + + def group_start(hidden_states): + events.append("start") + return hidden_states + + def get_context(**kwargs): + events.append("context") + assert kwargs["name"] == "grouped_mlp" + assert kwargs["max_num_tokens"] == 2 + torch.testing.assert_close(kwargs["num_tokens_tensor"], torch.tensor(2)) + assert kwargs["avg_num_tokens"] == 1 + return FakeStashContext() + + def group_commit(output, *, name): + events.append("commit") + assert name == "grouped_mlp" + return output + + monkeypatch.setattr(experts_module, "paged_stash_group_start", group_start) + monkeypatch.setattr(experts_module, "get_paged_stash_context", get_context) + monkeypatch.setattr(experts_module, "paged_stash_group_commit", group_commit) + + def unfused_forward(hidden_states, tokens_per_expert, permuted_probs): + events.append("compute") + assert isinstance(tokens_per_expert, torch.Tensor) + return hidden_states + permuted_probs + + module._unfused_forward = unfused_forward + + hidden_states = torch.zeros(2, 4) + tokens_per_expert = torch.tensor([1, 1]) + probs = torch.ones(2) + output, output_bias = module.forward(hidden_states, tokens_per_expert, probs) + + torch.testing.assert_close(output, torch.ones_like(hidden_states)) + assert output_bias is None + assert events == ["start", "context", "enter", "compute", "exit", "commit"] + + def test_apply_bias_returns_input_unchanged_when_bias_is_none(): intermediate = torch.arange(6, dtype=torch.float32).view(3, 2) diff --git a/tests/unit_tests/transformer/moe/test_paged_stashing.py b/tests/unit_tests/transformer/moe/test_paged_stashing.py index 967a8313d8b..b8c3a89a374 100644 --- a/tests/unit_tests/transformer/moe/test_paged_stashing.py +++ b/tests/unit_tests/transformer/moe/test_paged_stashing.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import pytest import torch @@ -11,6 +11,7 @@ from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.moe.moe_utils import get_align_size_for_quantization from megatron.core.transformer.moe.paged_stash import ( + PagedStashManager, check_paged_stash_overflow, paged_stash_init_chunk_handler, paged_stash_reset, @@ -119,6 +120,7 @@ def __init__( moe_flex_dispatcher_backend=kwargs.get("moe_flex_dispatcher_backend", None), moe_ncclep_static_shape=kwargs.get("moe_ncclep_static_shape", False), moe_grouped_gemm=kwargs.get("moe_grouped_gemm", False), + moe_use_grouped_tensor=kwargs.get("moe_use_grouped_tensor", False), moe_paged_stash=kwargs.get("moe_paged_stash", False), moe_expert_rank_capacity_factor=kwargs.get("moe_expert_rank_capacity_factor", None), moe_router_padding_for_fp8=kwargs.get("moe_router_padding_for_fp8", True), @@ -129,6 +131,8 @@ def __init__( ), gated_linear_unit=kwargs.get("gated_linear_unit", False), activation_func=kwargs.get("activation_func", F.gelu), + bias_activation_fusion=kwargs.get("bias_activation_fusion", False), + activation_func_fp8_input_store=kwargs.get("activation_func_fp8_input_store", False), moe_router_force_biased=kwargs.get("moe_router_force_biased", None), moe_paged_stash_buffer_size_factor_cuda=0.5, moe_paged_stash_buffer_size_factor_cpu=1.5, @@ -225,6 +229,78 @@ def _is_mxfp8_supported() -> bool: ) +@pytest.mark.skipif(not _is_mxfp8_supported(), reason=_MXFP8_SKIP_REASON) +@pytest.mark.skipif(not is_hybrid_ep_available(), reason="Hybrid EP are not available") +class TestPagedStashingGroupedTensor: + """Paged stashing with device-initiated GroupedLinear and no TE operation fuser.""" + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.internal + def test_forward_backward_without_op_fuser(self): + config.ENABLE_EXPERIMENTAL = True + + container = MoEModelTestContainer( + tp_size=1, + ep_size=4, + pp_size=1, + num_moe_experts=8, + num_layers=4, + moe_router_topk=2, + moe_router_load_balancing_type="aux_loss", + moe_token_dispatcher_type="flex", + moe_permute_fusion=True, + hidden_size=1024, + moe_flex_dispatcher_backend="hybridep", + test_dtype=torch.bfloat16, + moe_grouped_gemm=True, + moe_use_grouped_tensor=True, + moe_paged_stash=True, + moe_expert_rank_capacity_factor=1.5, + moe_paged_stash_buffer_size_factor_cuda=2.0, + moe_paged_stash_buffer_size_factor_cpu=0.0, + use_transformer_engine_op_fuser=False, + moe_router_padding_for_quantization=True, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=True, + ) + + assert container.config.use_transformer_engine_op_fuser is False + assert container.config.moe_use_grouped_tensor is True + + hidden_states = torch.randn((1024, 1, container.config.hidden_size), dtype=torch.bfloat16) + + # Capture the activation layout and token maxima. + paged_stash_reset(True, config=container.config) + paged_stash_init_chunk_handler(1, 0) + output_ref, hidden_states_grad_ref, _, _ = _forward_backward_all_layers( + container, hidden_states + ) + + stash_manager = PagedStashManager.get_instance() + assert ( + stash_manager.max_tokens_across_vp_stages + ), "No dynamic GroupedLinear/activation tensors were captured for paged stashing" + assert any( + dtype == torch.bfloat16 + for dtype, _hidden_size in stash_manager.max_tokens_across_vp_stages + ), "The fused activation's BF16 saved tensors were not captured for paged stashing" + container.zero_grad() + + # Allocate the stash buffers from the capture and exercise the real stash/reload path. + paged_stash_reset(True, config=container.config) + paged_stash_init_chunk_handler(1, 0) + output, hidden_states_grad, _, _ = _forward_backward_all_layers(container, hidden_states) + + overflow = check_paged_stash_overflow() + assert overflow.any().item() == 0 + torch.testing.assert_close(output, output_ref, atol=1e-4, rtol=1e-4) + torch.testing.assert_close(hidden_states_grad, hidden_states_grad_ref, atol=1e-4, rtol=1e-4) + + @pytest.mark.skipif(not _is_mxfp8_supported(), reason=_MXFP8_SKIP_REASON) @pytest.mark.skipif( not _te_grouped_mlp_op_fuser_environment_supported(), From d631ce6ffd7b037a9ac39b391619ac408999d9c2 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Mon, 24 Aug 2026 20:15:24 -0700 Subject: [PATCH 03/10] refactor: use TE paged stash markers Signed-off-by: hongbinl --- docs/user-guide/features/paged_stash.md | 7 ++++ megatron/core/fusions/fused_bias_geglu.py | 10 ++++-- megatron/core/fusions/fused_bias_swiglu.py | 10 ++++-- .../fusions/fused_weighted_squared_relu.py | 10 ++++-- megatron/core/transformer/moe/experts.py | 36 +++++++++++-------- .../transformer/moe/test_grouped_mlp.py | 16 +++++++++ 6 files changed, 69 insertions(+), 20 deletions(-) diff --git a/docs/user-guide/features/paged_stash.md b/docs/user-guide/features/paged_stash.md index 56b7db1233d..fb67d714f3d 100644 --- a/docs/user-guide/features/paged_stash.md +++ b/docs/user-guide/features/paged_stash.md @@ -47,6 +47,13 @@ HybridEP and TE grouped experts are required whenever `moe_expert_rank_capacity_ Path B removes host-device synchronization from grouped GEMM split metadata, but FC1, activation, and FC2 remain separate launches. Without a full-iteration CUDA graph it can therefore retain significant CPU launch overhead even though the expert path is host-device sync-free. +The legacy multi-stream cuBLAS GroupedLinear path is not supported because it materializes split +metadata on the host; paged stashing would not make that expert path sync-free. + +Paged stash identifies dynamic saved activations through Transformer Engine's +`mark_grouped_tensor` utility. The non-op-fuser integration marks these tensors explicitly rather +than inferring dynamic shapes from warmup iterations: shape sampling can misclassify a dynamic +tensor as static and is therefore not a safe correctness contract. In this context, sync-free refers to the steady-state expert data path. The initial paged-stash capture performs host reads, and the runner reads reduced overflow/over-budget state at the end of diff --git a/megatron/core/fusions/fused_bias_geglu.py b/megatron/core/fusions/fused_bias_geglu.py index 9c158badf86..f794f90a05f 100644 --- a/megatron/core/fusions/fused_bias_geglu.py +++ b/megatron/core/fusions/fused_bias_geglu.py @@ -6,9 +6,15 @@ def _propagate_paged_stash_marker(source, target): - """Preserve Megatron's dynamic-activation marker across view/cast operations.""" + """Preserve TE's dynamic-activation marker across view/cast operations.""" if hasattr(source, "grouped_tensor_scale_inv"): - setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + try: + from transformer_engine.pytorch.utils import mark_grouped_tensor + except ImportError as exc: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) from exc + mark_grouped_tensor(target) return target diff --git a/megatron/core/fusions/fused_bias_swiglu.py b/megatron/core/fusions/fused_bias_swiglu.py index 939f25864b6..833289f37ae 100644 --- a/megatron/core/fusions/fused_bias_swiglu.py +++ b/megatron/core/fusions/fused_bias_swiglu.py @@ -13,9 +13,15 @@ def _propagate_paged_stash_marker(source, target): - """Preserve Megatron's dynamic-activation marker across view/cast operations.""" + """Preserve TE's dynamic-activation marker across view/cast operations.""" if hasattr(source, "grouped_tensor_scale_inv"): - setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + try: + from transformer_engine.pytorch.utils import mark_grouped_tensor + except ImportError as exc: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) from exc + mark_grouped_tensor(target) return target diff --git a/megatron/core/fusions/fused_weighted_squared_relu.py b/megatron/core/fusions/fused_weighted_squared_relu.py index 484e1edef46..bb2091d6e22 100644 --- a/megatron/core/fusions/fused_weighted_squared_relu.py +++ b/megatron/core/fusions/fused_weighted_squared_relu.py @@ -9,9 +9,15 @@ def _propagate_paged_stash_marker(source, target): - """Preserve Megatron's dynamic-activation marker across view operations.""" + """Preserve TE's dynamic-activation marker across view operations.""" if hasattr(source, "grouped_tensor_scale_inv"): - setattr(target, "grouped_tensor_scale_inv", source.grouped_tensor_scale_inv) + try: + from transformer_engine.pytorch.utils import mark_grouped_tensor + except ImportError as exc: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) from exc + mark_grouped_tensor(target) return target diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index b0556fc04d9..14909979c25 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -61,9 +61,14 @@ import transformer_engine as te from megatron.core.extensions.transformer_engine import Fp8Padding, Fp8Unpadding + try: + from transformer_engine.pytorch.utils import mark_grouped_tensor as _te_mark_grouped_tensor + except ImportError: + _te_mark_grouped_tensor = None else: te = None # type: ignore[assignment, misc] Fp8Padding, Fp8Unpadding = None, None + _te_mark_grouped_tensor = None try: import flashinfer.fused_moe as fused_moe @@ -316,8 +321,13 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe permuted_probs, "grouped_tensor_scale_inv" ) if paged_stash_marked: - setattr(flat_output, "grouped_tensor_scale_inv", False) - setattr(flat_probs, "grouped_tensor_scale_inv", False) + if _te_mark_grouped_tensor is None: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) + # The multiply below saves these two token-shaped operands. The additive output + # operand is not saved by autograd and does not need a marker. + _te_mark_grouped_tensor(flat_probs) if tokens_per_expert.device != packed_bias.device: raise ValueError("Packed MoE bias and tokens_per_expert must be on the same device.") @@ -336,7 +346,7 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe packed_bias.float(), tokens_per_expert, dim=0, output_size=flat_output.size(0) ) if paged_stash_marked: - setattr(bias_per_token, "grouped_tensor_scale_inv", False) + _te_mark_grouped_tensor(bias_per_token) return (flat_output + bias_per_token * flat_probs).view(shape).to(output_dtype) @staticmethod @@ -730,9 +740,11 @@ def _mark_paged_stash_tensors(self, *tensors: Optional[torch.Tensor]) -> None: """Mark dynamic unfused activations for the paged-stash saved-tensor hook.""" if not self.config.moe_paged_stash: return - for tensor in tensors: - if tensor is not None: - setattr(tensor, "grouped_tensor_scale_inv", False) + if _te_mark_grouped_tensor is None: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) + _te_mark_grouped_tensor(*tensors) def _fused_forward( self, @@ -843,18 +855,17 @@ def _unfused_forward( permuted_probs: torch.Tensor, ) -> torch.Tensor: """Run FC1, activation, and FC2 without the TE operation fuser.""" - self._mark_paged_stash_tensors(permuted_local_hidden_states, permuted_probs) - if self.config.moe_apply_probs_on_input: assert ( self.config.moe_router_topk == 1 ), "`moe_apply_probs_on_input` only works with `moe_router_topk`=1." + # MulBackward saves both operands before GroupedLinear sees the scaled input. + self._mark_paged_stash_tensors(permuted_local_hidden_states, permuted_probs) original_dtype = permuted_local_hidden_states.dtype permuted_local_hidden_states = permuted_probs * permuted_local_hidden_states permuted_local_hidden_states = permuted_local_hidden_states.to(original_dtype) # Probs already applied, so reset to 1. permuted_probs = torch.ones_like(permuted_probs) - self._mark_paged_stash_tensors(permuted_local_hidden_states, permuted_probs) expert_fc1_manager = off_interface( self.offload_expert_fc1, permuted_local_hidden_states, "expert_fc1" @@ -868,7 +879,6 @@ def _unfused_forward( forced_released_tensors=[permuted_local_hidden_states], delay_offload=self.config.delay_offload_until_cuda_graph, ) - self._mark_paged_stash_tensors(fc1_output) moe_act_manager = off_interface(self.offload_moe_act, fc1_output, "moe_act") @@ -942,9 +952,9 @@ def glu(x): x_linear = x_linear.clamp(min=-val, max=val) self._mark_paged_stash_tensors(x_glu, x_linear) x_glu = self.config.activation_func(x_glu) - self._mark_paged_stash_tensors(x_glu) x_linear = x_linear + self.config.glu_linear_offset - self._mark_paged_stash_tensors(x_linear) + # MulBackward saves both newly-created operands. + self._mark_paged_stash_tensors(x_glu, x_linear) return x_glu * x_linear intermediate_parallel = glu(intermediate_parallel) @@ -965,7 +975,6 @@ def glu(x): else: with moe_act_manager as fc1_output: bias_act_output = bias_act_func(fc1_output, bias_parallel, permuted_probs) - self._mark_paged_stash_tensors(bias_act_output) output, output_bias = apply_module(self.linear_fc2)(bias_act_output, tokens_per_expert) if self.activation_recompute: self.activation_checkpoint.discard_output_and_register_recompute(output) @@ -977,7 +986,6 @@ def glu(x): forced_released_tensors=[fc1_output], delay_offload=self.config.delay_offload_until_cuda_graph, ) - self._mark_paged_stash_tensors(output, permuted_probs) return self._apply_bias(output, output_bias, tokens_per_expert, permuted_probs) def forward( diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index a4e743fff29..5d6c6674c74 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -88,6 +88,22 @@ def test_paged_stash_allows_non_fused_grouped_tensor_hybridep(): assert config.use_transformer_engine_op_fuser is False +def test_paged_stash_marking_delegates_to_transformer_engine(monkeypatch): + marked = [] + module = TEGroupedMLP.__new__(TEGroupedMLP) + module.config = SimpleNamespace(moe_paged_stash=True) + tensors = (torch.zeros(2, 4), torch.ones(2, 1)) + + monkeypatch.setattr( + experts_module, "_te_mark_grouped_tensor", lambda *args: marked.append(args) + ) + module._mark_paged_stash_tensors(*tensors) + + assert len(marked) == 1 + assert marked[0][0] is tensors[0] + assert marked[0][1] is tensors[1] + + def test_remove_glu_interleaving_restores_contiguous_gate_and_linear_halves(): interleaved = torch.tensor([[1, 2, 5, 6, 3, 4, 7, 8], [11, 12, 15, 16, 13, 14, 17, 18]]) expected = torch.tensor([[1, 2, 3, 4, 5, 6, 7, 8], [11, 12, 13, 14, 15, 16, 17, 18]]) From 5c7e4798cc30f41059dd4de6e963053f9ba575b1 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Mon, 24 Aug 2026 22:26:22 -0700 Subject: [PATCH 04/10] refactor: require fused activations for paged stash Signed-off-by: hongbinl --- docs/user-guide/features/paged_stash.md | 3 +- megatron/core/transformer/moe/experts.py | 7 ---- .../core/transformer/transformer_config.py | 18 +++++++++ .../transformer/moe/test_grouped_mlp.py | 39 +++++++++++++++++++ 4 files changed, 59 insertions(+), 8 deletions(-) diff --git a/docs/user-guide/features/paged_stash.md b/docs/user-guide/features/paged_stash.md index fb67d714f3d..ea5c9666ce8 100644 --- a/docs/user-guide/features/paged_stash.md +++ b/docs/user-guide/features/paged_stash.md @@ -21,7 +21,7 @@ Whenever `moe_expert_rank_capacity_factor` is set, a **runner** wraps forward-ba ## Prerequisites -HybridEP and TE grouped experts are required whenever `moe_expert_rank_capacity_factor` is set. The non-op-fuser path requires a Transformer Engine version whose GroupedLinear marks saved GroupedTensor activation buffers for paged stashing. With `moe_paged_stash` enabled: capacity factor must be set; no `cpu_offloading`; `offload_modules` must not include `expert_fc1`, `moe_act`, or `fused_group_mlp`. The runner is active whenever capacity factor is set (even without `--moe-paged-stash`) for over-budget reruns; stash overflow is checked only when paged stashing is on. +HybridEP and TE grouped experts are required whenever `moe_expert_rank_capacity_factor` is set. The non-op-fuser path requires a Transformer Engine version whose GroupedLinear marks saved GroupedTensor activation buffers for paged stashing. It currently supports only fused SwiGLU or QuickGeGLU (`bias_activation_fusion=True`) without GLU interleaving; restricting the activation contract keeps dynamic-tensor marking at the fused autograd boundaries. With `moe_paged_stash` enabled: capacity factor must be set; no `cpu_offloading`; `offload_modules` must not include `expert_fc1`, `moe_act`, or `fused_group_mlp`. The runner is active whenever capacity factor is set (even without `--moe-paged-stash`) for over-budget reruns; stash overflow is checked only when paged stashing is on. ## Configuration @@ -42,6 +42,7 @@ HybridEP and TE grouped experts are required whenever `moe_expert_rank_capacity_ # B. Device-initiated GroupedLinear, without the operation fuser --moe-grouped-gemm --moe-use-grouped-tensor +# Keep the default fused SwiGLU activation; do not pass --no-bias-swiglu-fusion. ``` Path B removes host-device synchronization from grouped GEMM split metadata, but FC1, activation, diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 14909979c25..4e3fd1bf474 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -898,11 +898,9 @@ def bias_act_func(intermediate_parallel, bias_parallel, permuted_probs): intermediate_parallel = self._remove_glu_interleaving( intermediate_parallel, self.config.moe_mlp_glu_interleave_size ) - self._mark_paged_stash_tensors(intermediate_parallel) intermediate_parallel = self.activation_func(intermediate_parallel) if permuted_probs is not None: original_dtype = intermediate_parallel.dtype - self._mark_paged_stash_tensors(intermediate_parallel, permuted_probs) intermediate_parallel = intermediate_parallel * permuted_probs intermediate_parallel = intermediate_parallel.to(original_dtype) elif self.config.bias_activation_fusion and not with_glu_interleaving: @@ -946,22 +944,17 @@ def glu(x): x, self.config.moe_mlp_glu_interleave_size ) x_glu, x_linear = torch.chunk(x, 2, dim=-1) - self._mark_paged_stash_tensors(x_glu, x_linear) if (val := self.config.activation_func_clamp_value) is not None: x_glu = x_glu.clamp(min=None, max=val) x_linear = x_linear.clamp(min=-val, max=val) - self._mark_paged_stash_tensors(x_glu, x_linear) x_glu = self.config.activation_func(x_glu) x_linear = x_linear + self.config.glu_linear_offset - # MulBackward saves both newly-created operands. - self._mark_paged_stash_tensors(x_glu, x_linear) return x_glu * x_linear intermediate_parallel = glu(intermediate_parallel) else: intermediate_parallel = self.activation_func(intermediate_parallel) original_dtype = intermediate_parallel.dtype - self._mark_paged_stash_tensors(intermediate_parallel, permuted_probs) intermediate_parallel = intermediate_parallel * permuted_probs intermediate_parallel = intermediate_parallel.to(original_dtype) return intermediate_parallel diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index f578ece5ed5..6a3515afb7b 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -2561,6 +2561,24 @@ def __post_init__(self): f"(paged stash covers those activations). " f"Remove: {moe_offload_conflict}" ) + if not self.use_transformer_engine_op_fuser: + if not self.moe_use_grouped_tensor: + raise ValueError( + "moe_paged_stash without use_transformer_engine_op_fuser requires " + "moe_use_grouped_tensor=True." + ) + if ( + not self.bias_activation_fusion + or not self.gated_linear_unit + or self.activation_func not in (F.silu, quick_gelu) + or self.moe_mlp_glu_interleave_size is not None + ): + raise ValueError( + "moe_paged_stash with the non-op-fuser GroupedTensor path requires " + "fused SwiGLU or QuickGeGLU: set bias_activation_fusion=True, " + "gated_linear_unit=True, activation_func to silu or quick_gelu, and " + "moe_mlp_glu_interleave_size=None." + ) if ( self.num_layers_in_first_pipeline_stage is not None diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index 5d6c6674c74..158239bd2ce 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -81,6 +81,9 @@ def test_paged_stash_allows_non_fused_grouped_tensor_hybridep(): moe_expert_rank_capacity_factor=1.5, moe_paged_stash=True, use_transformer_engine_op_fuser=False, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=True, ) assert config.moe_paged_stash is True @@ -88,6 +91,42 @@ def test_paged_stash_allows_non_fused_grouped_tensor_hybridep(): assert config.use_transformer_engine_op_fuser is False +@pytest.mark.parametrize( + "invalid_fused_activation_config", + [ + pytest.param({"bias_activation_fusion": False}, id="fusion-disabled"), + pytest.param({"gated_linear_unit": False}, id="not-gated"), + pytest.param({"activation_func": F.gelu}, id="unsupported-activation"), + pytest.param({"moe_mlp_glu_interleave_size": 16}, id="glu-interleaved"), + ], +) +def test_non_fused_grouped_tensor_paged_stash_requires_fused_bias_activation( + invalid_fused_activation_config, +): + kwargs = { + "num_layers": 1, + "hidden_size": 128, + "num_attention_heads": 4, + "num_moe_experts": 2, + "moe_grouped_gemm": True, + "moe_use_grouped_tensor": True, + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "hybridep", + "moe_expert_rank_capacity_factor": 1.5, + "moe_paged_stash": True, + "use_transformer_engine_op_fuser": False, + "gated_linear_unit": True, + "activation_func": F.silu, + "bias_activation_fusion": True, + } + kwargs.update(invalid_fused_activation_config) + + with pytest.raises( + ValueError, match="non-op-fuser GroupedTensor path requires fused SwiGLU or QuickGeGLU" + ): + TransformerConfig(**kwargs) + + def test_paged_stash_marking_delegates_to_transformer_engine(monkeypatch): marked = [] module = TEGroupedMLP.__new__(TEGroupedMLP) From 2b2c2616b17976dba6be342c3b3a63eb59bd6f37 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Mon, 24 Aug 2026 23:39:40 -0700 Subject: [PATCH 05/10] refactor: minimize paged stash integration diff Signed-off-by: hongbinl --- megatron/core/transformer/moe/experts.py | 80 +++++++++++++----------- 1 file changed, 43 insertions(+), 37 deletions(-) diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 4e3fd1bf474..80ea6ab8d73 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -707,35 +707,6 @@ def forward_post_hook(_module, _inputs, output): return forward_post_hook - def _start_paged_stash_group( - self, permuted_local_hidden_states: torch.Tensor, tokens_per_expert: torch.Tensor - ) -> tuple[torch.Tensor, object]: - """Start the grouped-MLP paged-stash scope when it is enabled.""" - if not self.config.moe_paged_stash: - return permuted_local_hidden_states, nullcontext() - - permuted_local_hidden_states = paged_stash_group_start(permuted_local_hidden_states) - max_num_tokens = permuted_local_hidden_states.shape[0] - # Average/expected tokens is a pre-padding estimate used by paged stashing heuristics. - # moe_expert_rank_capacity_factor is required when moe_paged_stash is enabled. - cap_factor = self.config.moe_expert_rank_capacity_factor - avg_num_tokens = ( - int(max_num_tokens // cap_factor) if cap_factor is not None and cap_factor > 0 else None - ) - stash_context = get_paged_stash_context( - name="grouped_mlp", - max_num_tokens=max_num_tokens, - num_tokens_tensor=tokens_per_expert.sum(), - avg_num_tokens=avg_num_tokens, - ) - return permuted_local_hidden_states, stash_context - - def _commit_paged_stash_group(self, output: torch.Tensor) -> torch.Tensor: - """Commit the grouped-MLP paged-stash scope when it is enabled.""" - if self.config.moe_paged_stash: - output = paged_stash_group_commit(output, name="grouped_mlp") - return output - def _mark_paged_stash_tensors(self, *tensors: Optional[torch.Tensor]) -> None: """Mark dynamic unfused activations for the paged-stash saved-tensor hook.""" if not self.config.moe_paged_stash: @@ -809,9 +780,25 @@ def _fused_forward( ) # if the number of tokens is 0, pad the hidden states to 256 - permuted_local_hidden_states, stash_context = self._start_paged_stash_group( - permuted_local_hidden_states, tokens_per_expert - ) + if self.config.moe_paged_stash: + permuted_local_hidden_states = paged_stash_group_start(permuted_local_hidden_states) + max_num_tokens = permuted_local_hidden_states.shape[0] + # Average/expected tokens is a pre-padding estimate used by paged stashing heuristics. + # moe_expert_rank_capacity_factor is required when moe_paged_stash is enabled. + cap_factor = self.config.moe_expert_rank_capacity_factor + avg_num_tokens = ( + int(max_num_tokens // cap_factor) + if cap_factor is not None and cap_factor > 0 + else None + ) + stash_context = get_paged_stash_context( + name="grouped_mlp", + max_num_tokens=max_num_tokens, + num_tokens_tensor=tokens_per_expert.sum(), + avg_num_tokens=avg_num_tokens, + ) + else: + stash_context = nullcontext() fine_grained_activation_offloading = getattr(self, "offload_fused_group_mlp", False) offload_name = "fused_group_mlp" fused_group_mlp_manager = off_interface( @@ -837,7 +824,9 @@ def _fused_forward( # Remove padding if needed if unpadded_tokens_per_expert is not None: output = self.quantization_unpadding(output, unpadded_tokens_per_expert) - return self._commit_paged_stash_group(output) + if self.config.moe_paged_stash: + output = paged_stash_group_commit(output, name="grouped_mlp") + return output @staticmethod def _remove_glu_interleaving(x: torch.Tensor, interleave_size: int) -> torch.Tensor: @@ -1046,9 +1035,25 @@ def forward( elif isinstance(tokens_per_expert, torch.Tensor): tokens_per_expert = tokens_per_expert.tolist() - permuted_local_hidden_states, stash_context = self._start_paged_stash_group( - permuted_local_hidden_states, tokens_per_expert - ) + if self.config.moe_paged_stash: + permuted_local_hidden_states = paged_stash_group_start(permuted_local_hidden_states) + max_num_tokens = permuted_local_hidden_states.shape[0] + # Average/expected tokens is a pre-padding estimate used by paged stashing heuristics. + # moe_expert_rank_capacity_factor is required when moe_paged_stash is enabled. + cap_factor = self.config.moe_expert_rank_capacity_factor + avg_num_tokens = ( + int(max_num_tokens // cap_factor) + if cap_factor is not None and cap_factor > 0 + else None + ) + stash_context = get_paged_stash_context( + name="grouped_mlp", + max_num_tokens=max_num_tokens, + num_tokens_tensor=tokens_per_expert.sum(), + avg_num_tokens=avg_num_tokens, + ) + else: + stash_context = nullcontext() with stash_context: output = self._unfused_forward( permuted_local_hidden_states, tokens_per_expert, permuted_probs @@ -1057,7 +1062,8 @@ def forward( if unpadded_tokens_per_expert is not None: output = self.quantization_unpadding(output, unpadded_tokens_per_expert) - output = self._commit_paged_stash_group(output) + if self.config.moe_paged_stash: + output = paged_stash_group_commit(output, name="grouped_mlp") output_bias = None return output, output_bias From e2dbbd7ca76bf0c3848624808a532d8d370c4368 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 25 Aug 2026 01:08:35 -0700 Subject: [PATCH 06/10] refactor: narrow paged stash marker propagation Signed-off-by: hongbinl --- megatron/core/fusions/fused_bias_swiglu.py | 4 +--- .../fusions/fused_weighted_squared_relu.py | 18 ++---------------- megatron/core/transformer/moe/experts.py | 6 +++--- 3 files changed, 6 insertions(+), 22 deletions(-) diff --git a/megatron/core/fusions/fused_bias_swiglu.py b/megatron/core/fusions/fused_bias_swiglu.py index 833289f37ae..105d3585a9c 100644 --- a/megatron/core/fusions/fused_bias_swiglu.py +++ b/megatron/core/fusions/fused_bias_swiglu.py @@ -191,7 +191,6 @@ def forward(ctx, input, bias, fp8_input_store, cpu_offload_input, clamp_value): torch.Tensor: Result of applying bias addition followed by SwiGLU activation. """ input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input - _propagate_paged_stash_marker(input, input_for_backward) if cpu_offload_input: input_for_backward.activation_offloading = True bias.activation_offloading = True @@ -250,7 +249,6 @@ def forward(ctx, input, fp8_input_store, cpu_offload_input, clamp_value): torch.Tensor: Result of applying SwiGLU activation. """ input_for_backward = input.to(torch.float8_e4m3fn) if fp8_input_store else input - _propagate_paged_stash_marker(input, input_for_backward) if cpu_offload_input: input_for_backward.activation_offloading = True ctx.save_for_backward(input_for_backward) @@ -338,7 +336,7 @@ def bias_swiglu_impl(input, bias, fp8_input_store=False, cpu_offload_input=False """ ori_shape = input.shape assert len(ori_shape) in [2, 3] - input = _propagate_paged_stash_marker(input, input.view(-1, ori_shape[-1])) + input = input.view(-1, ori_shape[-1]) if bias is not None: output = BiasSwiGLUFunction.apply( input, bias, fp8_input_store, cpu_offload_input, clamp_value diff --git a/megatron/core/fusions/fused_weighted_squared_relu.py b/megatron/core/fusions/fused_weighted_squared_relu.py index bb2091d6e22..02dabc14c3b 100644 --- a/megatron/core/fusions/fused_weighted_squared_relu.py +++ b/megatron/core/fusions/fused_weighted_squared_relu.py @@ -1,4 +1,4 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. import torch import torch.nn.functional as F @@ -7,20 +7,6 @@ from megatron.core.jit import jit_fuser from megatron.core.utils import nvtx_decorator - -def _propagate_paged_stash_marker(source, target): - """Preserve TE's dynamic-activation marker across view operations.""" - if hasattr(source, "grouped_tensor_scale_inv"): - try: - from transformer_engine.pytorch.utils import mark_grouped_tensor - except ImportError as exc: - raise RuntimeError( - "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." - ) from exc - mark_grouped_tensor(target) - return target - - ###################### WEIGHTED SQUARED ReLU FUSION ###################### @@ -117,7 +103,7 @@ def weighted_squared_relu_impl(input: torch.Tensor, weights: torch.Tensor) -> to """ ori_shape = input.shape assert len(ori_shape) in [2, 3] - input = _propagate_paged_stash_marker(input, input.view(-1, ori_shape[-1])) + input = input.view(-1, ori_shape[-1]) output = WeightedSquaredReLUFunction.apply(input, weights) diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 80ea6ab8d73..1fe396b33f2 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -936,9 +936,9 @@ def glu(x): if (val := self.config.activation_func_clamp_value) is not None: x_glu = x_glu.clamp(min=None, max=val) x_linear = x_linear.clamp(min=-val, max=val) - x_glu = self.config.activation_func(x_glu) - x_linear = x_linear + self.config.glu_linear_offset - return x_glu * x_linear + return self.config.activation_func(x_glu) * ( + x_linear + self.config.glu_linear_offset + ) intermediate_parallel = glu(intermediate_parallel) else: From a3227074a08f5ba1623d28667f2fafa625783295 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 25 Aug 2026 03:10:28 -0700 Subject: [PATCH 07/10] refactor: route paged stash markers through TE adapter Signed-off-by: hongbinl --- .../core/extensions/transformer_engine.py | 15 +++++++++++ megatron/core/fusions/fused_bias_geglu.py | 9 +++---- megatron/core/fusions/fused_bias_swiglu.py | 9 +++---- megatron/core/transformer/moe/experts.py | 21 +++------------- .../transformer/moe/test_grouped_mlp.py | 25 ++++++++++++++++--- 5 files changed, 47 insertions(+), 32 deletions(-) diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 365faba6331..bb76d10fb2d 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -72,8 +72,14 @@ import transformer_engine as te from transformer_engine.pytorch.fp8 import FP8GlobalStateManager, fp8_autocast, fp8_model_init + try: + from transformer_engine.pytorch.utils import mark_grouped_tensor as _te_mark_grouped_tensor + except ImportError: + _te_mark_grouped_tensor = None + HAVE_TE = True except ImportError: + _te_mark_grouped_tensor = None if TYPE_CHECKING: # For type checking, treat transformer_engine as always available. import transformer_engine as te @@ -89,6 +95,15 @@ _TE_CONFIG_TYPE_KEY = "transformer_engine_config_type" +def mark_grouped_tensor(*tensors: Any) -> None: + """Mark dynamic grouped tensors through the Transformer Engine compatibility boundary.""" + if _te_mark_grouped_tensor is None: + raise RuntimeError( + "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." + ) + _te_mark_grouped_tensor(*tensors) + + class TransformerEngineConfigType(enum.Enum): """Configuration object types in config dictionary""" diff --git a/megatron/core/fusions/fused_bias_geglu.py b/megatron/core/fusions/fused_bias_geglu.py index f794f90a05f..ff698d01919 100644 --- a/megatron/core/fusions/fused_bias_geglu.py +++ b/megatron/core/fusions/fused_bias_geglu.py @@ -8,12 +8,9 @@ def _propagate_paged_stash_marker(source, target): """Preserve TE's dynamic-activation marker across view/cast operations.""" if hasattr(source, "grouped_tensor_scale_inv"): - try: - from transformer_engine.pytorch.utils import mark_grouped_tensor - except ImportError as exc: - raise RuntimeError( - "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." - ) from exc + # Lazy import avoids the transformer_engine extension -> MLP -> fusion import cycle. + from megatron.core.extensions.transformer_engine import mark_grouped_tensor + mark_grouped_tensor(target) return target diff --git a/megatron/core/fusions/fused_bias_swiglu.py b/megatron/core/fusions/fused_bias_swiglu.py index 105d3585a9c..4b6af2a271f 100644 --- a/megatron/core/fusions/fused_bias_swiglu.py +++ b/megatron/core/fusions/fused_bias_swiglu.py @@ -15,12 +15,9 @@ def _propagate_paged_stash_marker(source, target): """Preserve TE's dynamic-activation marker across view/cast operations.""" if hasattr(source, "grouped_tensor_scale_inv"): - try: - from transformer_engine.pytorch.utils import mark_grouped_tensor - except ImportError as exc: - raise RuntimeError( - "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." - ) from exc + # Lazy import avoids the transformer_engine extension -> MLP -> fusion import cycle. + from megatron.core.extensions.transformer_engine import mark_grouped_tensor + mark_grouped_tensor(target) return target diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 1fe396b33f2..e08ac576b82 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -19,7 +19,7 @@ from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp4Recipe, Fp8Recipe -from megatron.core.extensions.transformer_engine import HAVE_TE +from megatron.core.extensions.transformer_engine import HAVE_TE, mark_grouped_tensor from megatron.core.fusions.fused_bias_geglu import quick_gelu, weighted_bias_quick_geglu_impl from megatron.core.fusions.fused_bias_swiglu import weighted_bias_swiglu_impl from megatron.core.fusions.fused_weighted_squared_relu import weighted_squared_relu_impl @@ -61,14 +61,9 @@ import transformer_engine as te from megatron.core.extensions.transformer_engine import Fp8Padding, Fp8Unpadding - try: - from transformer_engine.pytorch.utils import mark_grouped_tensor as _te_mark_grouped_tensor - except ImportError: - _te_mark_grouped_tensor = None else: te = None # type: ignore[assignment, misc] Fp8Padding, Fp8Unpadding = None, None - _te_mark_grouped_tensor = None try: import flashinfer.fused_moe as fused_moe @@ -321,13 +316,9 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe permuted_probs, "grouped_tensor_scale_inv" ) if paged_stash_marked: - if _te_mark_grouped_tensor is None: - raise RuntimeError( - "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." - ) # The multiply below saves these two token-shaped operands. The additive output # operand is not saved by autograd and does not need a marker. - _te_mark_grouped_tensor(flat_probs) + mark_grouped_tensor(flat_probs) if tokens_per_expert.device != packed_bias.device: raise ValueError("Packed MoE bias and tokens_per_expert must be on the same device.") @@ -346,7 +337,7 @@ def _apply_packed_bias(intermediate_parallel, packed_bias, tokens_per_expert, pe packed_bias.float(), tokens_per_expert, dim=0, output_size=flat_output.size(0) ) if paged_stash_marked: - _te_mark_grouped_tensor(bias_per_token) + mark_grouped_tensor(bias_per_token) return (flat_output + bias_per_token * flat_probs).view(shape).to(output_dtype) @staticmethod @@ -711,11 +702,7 @@ def _mark_paged_stash_tensors(self, *tensors: Optional[torch.Tensor]) -> None: """Mark dynamic unfused activations for the paged-stash saved-tensor hook.""" if not self.config.moe_paged_stash: return - if _te_mark_grouped_tensor is None: - raise RuntimeError( - "Paged stashing requires Transformer Engine's mark_grouped_tensor utility." - ) - _te_mark_grouped_tensor(*tensors) + mark_grouped_tensor(*tensors) def _fused_forward( self, diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index 158239bd2ce..759ec48b078 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -10,6 +10,8 @@ import megatron.core.transformer.moe.experts as experts_module from megatron.core.activations import squared_relu +from megatron.core.extensions import transformer_engine as te_ext +from megatron.core.fusions import fused_bias_geglu, fused_bias_swiglu from megatron.core.fusions.fused_bias_geglu import quick_gelu from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_local_submodules, @@ -133,9 +135,7 @@ def test_paged_stash_marking_delegates_to_transformer_engine(monkeypatch): module.config = SimpleNamespace(moe_paged_stash=True) tensors = (torch.zeros(2, 4), torch.ones(2, 1)) - monkeypatch.setattr( - experts_module, "_te_mark_grouped_tensor", lambda *args: marked.append(args) - ) + monkeypatch.setattr(te_ext, "_te_mark_grouped_tensor", lambda *args: marked.append(args)) module._mark_paged_stash_tensors(*tensors) assert len(marked) == 1 @@ -143,6 +143,25 @@ def test_paged_stash_marking_delegates_to_transformer_engine(monkeypatch): assert marked[0][1] is tensors[1] +@pytest.mark.parametrize( + "propagate_marker", + ( + fused_bias_geglu._propagate_paged_stash_marker, + fused_bias_swiglu._propagate_paged_stash_marker, + ), +) +def test_fused_activation_marker_propagation_uses_te_adapter(monkeypatch, propagate_marker): + marked = [] + source = torch.zeros(2, 4) + target = torch.ones(2, 4) + source.grouped_tensor_scale_inv = False + + monkeypatch.setattr(te_ext, "_te_mark_grouped_tensor", lambda *args: marked.append(args)) + + assert propagate_marker(source, target) is target + assert marked == [(target,)] + + def test_remove_glu_interleaving_restores_contiguous_gate_and_linear_halves(): interleaved = torch.tensor([[1, 2, 5, 6, 3, 4, 7, 8], [11, 12, 15, 16, 13, 14, 17, 18]]) expected = torch.tensor([[1, 2, 3, 4, 5, 6, 7, 8], [11, 12, 13, 14, 15, 16, 17, 18]]) From 34d3765230ca9e526dd05beb25ae72396332f1c3 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 25 Aug 2026 04:12:52 -0700 Subject: [PATCH 08/10] docs: trim paged stash implementation details Signed-off-by: hongbinl --- docs/user-guide/features/paged_stash.md | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/docs/user-guide/features/paged_stash.md b/docs/user-guide/features/paged_stash.md index ea5c9666ce8..e5aeefa4bea 100644 --- a/docs/user-guide/features/paged_stash.md +++ b/docs/user-guide/features/paged_stash.md @@ -51,16 +51,6 @@ significant CPU launch overhead even though the expert path is host-device sync- The legacy multi-stream cuBLAS GroupedLinear path is not supported because it materializes split metadata on the host; paged stashing would not make that expert path sync-free. -Paged stash identifies dynamic saved activations through Transformer Engine's -`mark_grouped_tensor` utility. The non-op-fuser integration marks these tensors explicitly rather -than inferring dynamic shapes from warmup iterations: shape sampling can misclassify a dynamic -tensor as static and is therefore not a safe correctness contract. - -In this context, sync-free refers to the steady-state expert data path. The initial paged-stash -capture performs host reads, and the runner reads reduced overflow/over-budget state at the end of -a pass to decide whether to rerun; it does not imply that the complete iteration has no CPU-GPU -synchronization at all. - ## Tuning (paged stashing only) ```bash From f38e49f99ccaacead19e83d7d0716161422f42de Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 25 Aug 2026 04:26:27 -0700 Subject: [PATCH 09/10] chore: align dev copyright headers Signed-off-by: hongbinl --- megatron/core/transformer/moe/moe_utils.py | 2 +- .../transformer/moe/test_grouped_tensor_dispatcher_numerics.py | 2 +- .../transformer/moe/test_moe_single_grouped_weight_numerics.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index c8a7e2e3ce5..5e493c4d947 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import functools import math from dataclasses import dataclass diff --git a/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py b/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py index 0e60e64c809..da56a3d790f 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py +++ b/tests/unit_tests/transformer/moe/test_grouped_tensor_dispatcher_numerics.py @@ -1,4 +1,4 @@ -# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. """Distributed MoE coverage for the TE grouped-tensor expert path. diff --git a/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py b/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py index 711405a5f8a..8b05807faaf 100644 --- a/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py +++ b/tests/unit_tests/transformer/moe/test_moe_single_grouped_weight_numerics.py @@ -1,4 +1,4 @@ -# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import gc import inspect From e569fa38b73b4ad13efd0ef16791edf665716e67 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 25 Aug 2026 06:06:40 -0700 Subject: [PATCH 10/10] test: gate grouped-tensor paged stash on TE support Signed-off-by: hongbinl --- .../transformer/moe/test_paged_stashing.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/tests/unit_tests/transformer/moe/test_paged_stashing.py b/tests/unit_tests/transformer/moe/test_paged_stashing.py index b8c3a89a374..4c054ffb205 100644 --- a/tests/unit_tests/transformer/moe/test_paged_stashing.py +++ b/tests/unit_tests/transformer/moe/test_paged_stashing.py @@ -1,5 +1,7 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import inspect + import pytest import torch import torch.nn.functional as F @@ -211,6 +213,17 @@ def _te_grouped_mlp_op_fuser_environment_supported() -> bool: return is_te_min_version("2.14.0") +def _te_grouped_tensor_environment_supported() -> bool: + """Return whether TE GroupedLinear exposes the device-initiated grouped-tensor API.""" + if not HAVE_TE: + return False + try: + from transformer_engine.pytorch import GroupedLinear + except ImportError: + return False + return "use_grouped_tensor" in inspect.signature(GroupedLinear.__init__).parameters + + _TE_GROUPED_MLP_OP_FUSER_SKIP_REASON = ( "TEGroupedMLP op fuser (tests use use_transformer_engine_op_fuser=True) requires TE>=2.14 " "with GroupedLinear/ScaledSwiGLU ops" @@ -230,6 +243,10 @@ def _is_mxfp8_supported() -> bool: @pytest.mark.skipif(not _is_mxfp8_supported(), reason=_MXFP8_SKIP_REASON) +@pytest.mark.skipif( + not _te_grouped_tensor_environment_supported(), + reason="Installed TE GroupedLinear does not expose use_grouped_tensor", +) @pytest.mark.skipif(not is_hybrid_ep_available(), reason="Hybrid EP are not available") class TestPagedStashingGroupedTensor: """Paged stashing with device-initiated GroupedLinear and no TE operation fuser."""