diff --git a/docs/user-guide/features/paged_stash.md b/docs/user-guide/features/paged_stash.md index b5b97144905..e5aeefa4bea 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,36 @@ 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. 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 ```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 +# 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, +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. + ## Tuning (paged stashing only) ```bash diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 51b63983007..24d2d81a509 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from __future__ import annotations import copy @@ -70,8 +70,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 @@ -88,6 +94,15 @@ _EXPERT_PARAMETER_NAME_PATTERN = re.compile(r"(weight|bias)\d*") +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) + + def _set_expert_parameter_attributes( module: torch.nn.Module, parallel_mode: Optional[str], use_expert_pgs: bool ) -> None: diff --git a/megatron/core/fusions/fused_bias_geglu.py b/megatron/core/fusions/fused_bias_geglu.py index 7a7fbe7f9ec..ff698d01919 100644 --- a/megatron/core/fusions/fused_bias_geglu.py +++ b/megatron/core/fusions/fused_bias_geglu.py @@ -1,9 +1,20 @@ -# 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 TE's dynamic-activation marker across view/cast operations.""" + if hasattr(source, "grouped_tensor_scale_inv"): + # 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 + + ###### BIAS GELU FUSION/ NO AUTOGRAD ################ # 1/sqrt(2*pi)-> 0.3989423 # 1/sqrt(2) -> 0.70710678 @@ -324,6 +335,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 +386,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 +433,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 +444,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 8a32d90b871..d3ee01feed5 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 @@ -14,6 +14,16 @@ ###### BIAS SWIGLU FUSION/ NO AUTOGRAD ################ +def _propagate_paged_stash_marker(source, target): + """Preserve TE's dynamic-activation marker across view/cast operations.""" + if hasattr(source, "grouped_tensor_scale_inv"): + # 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 + + @jit_fuser def swiglu(y): """Performs SwiGLU (Swish-Gated Linear Unit) activation function. @@ -389,6 +399,7 @@ def forward( ctx, input, weights, fp8_input_store, clamp_value, gate_clamp_scale, linear_clamp_scale ): 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 @@ -505,7 +516,7 @@ def weighted_bias_swiglu_impl( assert len(ori_shape) in [2, 3] assert gate_clamp_scale is not None or linear_clamp_scale is None assert gate_clamp_scale is None or clamp_value is None - 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/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 1ae2b3b032e..500242eb27b 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 @@ -43,6 +43,7 @@ ) from megatron.core.transformer.moe.paged_stash import ( get_paged_stash_context, + mark_paged_stash_recompute_managed, paged_stash_group_commit, paged_stash_group_start, ) @@ -328,6 +329,13 @@ 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: + # The multiply below saves these two token-shaped operands. The additive output + # operand is not saved by autograd and does not need a marker. + 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.") @@ -345,6 +353,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: + mark_grouped_tensor(bias_per_token) return (flat_output + bias_per_token * flat_probs).view(shape).to(output_dtype) @staticmethod @@ -707,6 +717,12 @@ def _ensure_main_grad_for_fused_impl(self) -> None: self._ensure_main_grad(self.linear_fc1) self._ensure_main_grad(self.linear_fc2) + 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 + mark_grouped_tensor(*tensors) + def _fused_forward( self, permuted_local_hidden_states: torch.Tensor, @@ -846,88 +862,19 @@ 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, - output_buffer: Optional[torch.Tensor] = None, - grad_input_buffer: Optional[torch.Tensor] = None, - ) -> 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. - output_buffer (torch.Tensor, optional): Preallocated buffer to write the fc2 output into - (NCCL-EP zero-copy fwd combine); only the fused op-fuser path supports it. - grad_input_buffer (torch.Tensor, optional): Preallocated buffer to write the fc1 dgrad - into (NCCL-EP zero-copy bwd dispatch); only the fused op-fuser path supports it. - - 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_buffer, - grad_input_buffer, - ) - output_bias = None - return output, output_bias - assert ( - output_buffer is None and grad_input_buffer is None - ), "output_buffer/grad_input_buffer require the TE op-fuser (fused) path" - - # 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.""" 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) @@ -950,6 +897,7 @@ def forward( 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 = ( @@ -1048,6 +996,8 @@ def glu(x): bias_act_output = self.activation_checkpoint.checkpoint( bias_act_func, fc1_output, bias_parallel, permuted_probs ) + if self.config.moe_paged_stash: + mark_paged_stash_recompute_managed(bias_act_output) else: with moe_act_manager as fc1_output: bias_act_output = bias_act_func(fc1_output, bias_parallel, permuted_probs) @@ -1062,12 +1012,116 @@ 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) + 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, + output_buffer: Optional[torch.Tensor] = None, + grad_input_buffer: Optional[torch.Tensor] = None, + ) -> 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. + output_buffer (torch.Tensor, optional): Preallocated buffer to write the fc2 output into + (NCCL-EP zero-copy fwd combine); only the fused op-fuser path supports it. + grad_input_buffer (torch.Tensor, optional): Preallocated buffer to write the fc1 dgrad + into (NCCL-EP zero-copy bwd dispatch); only the fused op-fuser path supports it. + + 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_buffer, + grad_input_buffer, + ) + output_bias = None + return output, output_bias + assert ( + output_buffer is None and grad_input_buffer is None + ), "output_buffer/grad_input_buffer require the TE op-fuser (fused) path" + + # 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() + + 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 + ) # upad and concat the output 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") output_bias = None return output, output_bias diff --git a/megatron/core/transformer/moe/paged_stash.py b/megatron/core/transformer/moe/paged_stash.py index a3a3fff76c1..e213f1c0bb9 100644 --- a/megatron/core/transformer/moe/paged_stash.py +++ b/megatron/core/transformer/moe/paged_stash.py @@ -24,6 +24,12 @@ _MAX_RERUN_ATTEMPTS = 2 SCALE_INV_BLOCK_SIZE = 32 +_RECOMPUTE_MANAGED_TENSOR_ATTR = '_mcore_paged_stash_recompute_managed' + + +def mark_paged_stash_recompute_managed(tensor: torch.Tensor) -> None: + """Mark a tensor whose storage is released and restored by activation recomputation.""" + setattr(tensor, _RECOMPUTE_MANAGED_TENSOR_ATTR, True) class PagedStashBuffer: @@ -686,6 +692,12 @@ def on_save_for_backward(self, tensor: torch.Tensor) -> Any: Hook called when autograd saves a tensor for backward pass. Returns a tag to identify the tensor later. """ + # CheckpointWithoutOutput intentionally releases the storage of its output after the + # following module has saved it, then restores that same StorageImpl during backward + # recomputation. Paged stashing must not take ownership of such a tensor in between. + if getattr(tensor, _RECOMPUTE_MANAGED_TENSOR_ATTR, False): + return tensor + # Handle 0-dim tensors (torch.Size([])) - they have no size(0) if ( self.max_num_tokens is None diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index c3cfe45865d..197ff765ac9 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -2395,6 +2395,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 91a908f5393..21993ed29f2 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -1,4 +1,4 @@ -# Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import argparse import sys @@ -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, @@ -68,6 +70,98 @@ 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, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=True, + ) + + assert config.moe_paged_stash is True + assert config.moe_use_grouped_tensor is True + 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) + module.config = SimpleNamespace(moe_paged_stash=True) + tensors = (torch.zeros(2, 4), torch.ones(2, 1)) + + monkeypatch.setattr(te_ext, "_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] + + +@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_clamped_swiglu_allows_te_op_fuser(): config = TransformerConfig( num_layers=1, @@ -246,6 +340,63 @@ def __call__(self, *args): assert len(fused_ops.args) == 4 +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 aca0dae077b..e9a6395c393 100644 --- a/tests/unit_tests/transformer/moe/test_paged_stashing.py +++ b/tests/unit_tests/transformer/moe/test_paged_stashing.py @@ -1,4 +1,6 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import inspect import pytest import torch @@ -11,7 +13,9 @@ 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, + mark_paged_stash_recompute_managed, paged_stash_init_chunk_handler, paged_stash_reset, ) @@ -122,6 +126,7 @@ def __init__( moe_dispatch_fwd_dtype=kwargs.get("moe_dispatch_fwd_dtype", 'bf16'), moe_combine_bwd_dtype=kwargs.get("moe_combine_bwd_dtype", 'bf16'), 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), @@ -132,6 +137,10 @@ 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), + recompute_granularity=kwargs.get("recompute_granularity", None), + recompute_modules=kwargs.get("recompute_modules", None), moe_router_force_biased=kwargs.get("moe_router_force_biased", None), # Shrinking the CUDA factor and zeroing the CPU one (no host-spill fallback) is how # a test forces a paged-stash overflow: the pool is provisioned once at the @@ -243,6 +252,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" @@ -261,6 +281,95 @@ def _is_mxfp8_supported() -> bool: ) +def test_recompute_managed_tensor_bypasses_paged_stash_save_hook(): + tensor = torch.randn(8, 4) + tensor.grouped_tensor_scale_inv = False + mark_paged_stash_recompute_managed(tensor) + + # The ownership check happens before the manager needs CUDA streams or capture state. + manager = object.__new__(PagedStashManager) + assert manager.on_save_for_backward(tensor) is tensor + + +@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.""" + + 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, + recompute_granularity="selective", + recompute_modules=["moe_act"], + ) + + assert container.config.use_transformer_engine_op_fuser is False + assert container.config.moe_use_grouped_tensor is True + assert container.config.recompute_modules == ["moe_act"] + + 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(),