diff --git a/megatron/core/models/common/utils.py b/megatron/core/models/common/utils.py index d1b7c54fefb..60dd9610cda 100644 --- a/megatron/core/models/common/utils.py +++ b/megatron/core/models/common/utils.py @@ -77,6 +77,10 @@ def should_free_input(name, is_moe, config, num_local_experts): config.moe_token_dispatcher_type == "flex" and config.moe_flex_dispatcher_backend == "ncclep" ) + enable_moonep = ( + config.moe_token_dispatcher_type == "flex" + and config.moe_flex_dispatcher_backend == "moonep" + ) # Define which nodes should free input memory. # Since we split the computing graph into multiple nodes, we can manually control # when and how to free the input memory. @@ -87,22 +91,22 @@ def should_free_input(name, is_moe, config, num_local_experts): # original bf16 tensors are safe to be freed. free_mlp = config.fp8 is not None or config.fp4 is not None if not free_mlp: - # AlltoAll dispatcher with local_num_experts=1, HybridEP, and NCCL EP all use + # AlltoAll dispatcher with local_num_experts=1, HybridEP, MoonEP, and NCCL EP all use # identity operation for `dispatch_postprocess`, hence the mlp inputs will be # directly passed to GroupedGemm and should be saved for backward pass. free_mlp = num_local_experts > 1 or config.moe_token_dispatcher_type != "alltoall" - free_mlp = free_mlp and not (enable_hybridep or enable_ncclep) + free_mlp = free_mlp and not (enable_hybridep or enable_moonep or enable_ncclep) free_input_nodes = { "mlp": free_mlp, "moe_combine": True, - # For non-DeepEP/HybridEP/NCCL-EP dispatcher mode, the input is the un-dispatched + # For non-DeepEP/HybridEP/MoonEP/NCCL-EP dispatcher mode, the input is the un-dispatched # tokens and probs before dispatch A2A and it's not needed anymore after the - # forward pass. For DeepEP, HybridEP, and NCCL EP dispatcher mode, they are all + # forward pass. For DeepEP, HybridEP, MoonEP, and NCCL EP dispatcher mode, they are all # needed in backward pass and cannot be freed. # If moe_preprocess is in cuda graph scope, tokens and probs are fixed size # tensors, so they cannot be freed. - "moe_dispatch": not (enable_deepep or enable_hybridep or enable_ncclep) + "moe_dispatch": not (enable_deepep or enable_hybridep or enable_moonep or enable_ncclep) and (CudaGraphModule.moe_preprocess not in config.cuda_graph_modules), } diff --git a/megatron/core/models/gpt/fine_grained_callables.py b/megatron/core/models/gpt/fine_grained_callables.py index ba1fd3af499..df6caef7b22 100644 --- a/megatron/core/models/gpt/fine_grained_callables.py +++ b/megatron/core/models/gpt/fine_grained_callables.py @@ -63,6 +63,10 @@ def build_transformer_layer_callables(layer: TransformerLayer): layer.config.moe_token_dispatcher_type == "flex" and layer.config.moe_flex_dispatcher_backend == "ncclep" ) + enable_moonep = ( + layer.config.moe_token_dispatcher_type == "flex" + and layer.config.moe_flex_dispatcher_backend == "moonep" + ) def submodule_pre_dispatch_forward(node: ScheduleNode, hidden_states: torch.Tensor): """ @@ -160,7 +164,7 @@ def submodule_dispatch_forward( Dispatches tokens to the experts based on the router output. """ token_dispatcher = layer.mlp.token_dispatcher - if enable_deepep or enable_hybridep or enable_ncclep: + if enable_deepep or enable_hybridep or enable_moonep or enable_ncclep: # update token_probs to be the detached version, prevents # backward graph from connecting to pre_dispatch_computation submodule token_dispatcher._comm_manager.token_probs = probs @@ -186,16 +190,16 @@ def submodule_moe_forward(node: ScheduleNode, dispatched_tokens: torch.Tensor): """ dispatched_probs = node.layer_state.dispatched_probs token_dispatcher = layer.mlp.token_dispatcher - if enable_deepep or enable_hybridep or enable_ncclep: + if enable_deepep or enable_hybridep or enable_moonep or enable_ncclep: # update dispatched_probs to be detached version, prevents # backward graph from connecting to dispatch submodule token_dispatcher._comm_manager.dispatched_probs = dispatched_probs expert_output, _ = layer.mlp.routed_experts_compute(dispatched_tokens, dispatched_probs) - # For HybridEP and NCCL EP, tokens_per_expert is generated on comm stream, as the + # For HybridEP, MoonEP, and NCCL EP, tokens_per_expert is generated by dispatch, as the # input to `routed_experts_compute`, a ref is needed to prevent it from being freed. - if enable_hybridep or enable_ncclep: + if enable_hybridep or enable_moonep or enable_ncclep: tokens_per_expert = token_dispatcher._comm_manager.get_number_of_tokens_per_expert() node.layer_state.tokens_per_expert = tokens_per_expert diff --git a/megatron/core/parallel_state.py b/megatron/core/parallel_state.py index 2a3c7581122..0e1daa7bdbc 100644 --- a/megatron/core/parallel_state.py +++ b/megatron/core/parallel_state.py @@ -2505,10 +2505,14 @@ def get_all_ranks(): def destroy_model_parallel(): """Set the groups to none.""" - # Release the NCCL EP context (if the 'ncclep' flex dispatcher bootstrapped one) before the - # process group's communicator is torn down. TE registers an atexit ep_finalize that would - # otherwise run after dist.destroy_process_group() and hit a "corrupted comm object" at exit. - # Idempotent and a no-op when NCCL EP was never bootstrapped. + # Release flex-dispatcher contexts before their process-group communicators are torn down. + # The finalizers are idempotent no-ops when their corresponding backends were never used. + try: + from megatron.core.transformer.moe.fused_a2a import moonep_finalize + + moonep_finalize() + except Exception: # finalize must never block teardown + pass try: from megatron.core.transformer.moe.fused_a2a import nccl_ep_finalize diff --git a/megatron/core/transformer/moe/README.md b/megatron/core/transformer/moe/README.md index f3c35fe4f6e..d9d3f064eac 100644 --- a/megatron/core/transformer/moe/README.md +++ b/megatron/core/transformer/moe/README.md @@ -309,8 +309,53 @@ After routing, tokens are **dispatched** to the GPU hosting the assigned expert. | **alltoall** | NCCL-based All-to-All communication for token exchange | Standard EP > 1 setups | `--moe-token-dispatcher-type alltoall` | | **FlexDispatcher with [DeepEP](https://github.com/deepseek-ai/DeepEP) backend** | Removes redundant tokens during cross-node communication, fuses intra/inter-node communication into single kernel | Cross-node EP, fine-grained MoE (DeepSeek-V3) | `--moe-token-dispatcher-type flex --moe-flex-dispatcher-backend deepep` | | **FlexDispatcher with [HybridEP](https://github.com/deepseek-ai/DeepEP/tree/hybrid-ep) backend** | NVIDIA's optimized dispatcher using TMA and IBGDA, fewer SMs, native MNNVL support | GB200 NVL72, Multi-Node NVLink | `--moe-token-dispatcher-type flex --moe-flex-dispatcher-backend hybridep` | +| **FlexDispatcher with MoonEP backend** | Prefetches redundant experts into NVLink-visible slots while retaining Megatron grouped parameters for optimization and checkpoints | Eager BF16 training on one NVLink-connected node | `--moe-token-dispatcher-type flex --moe-flex-dispatcher-backend moonep` | | **allgather** | Gathers all tokens to each GPU, no inter-GPU token movement | TP-only setups, small EP, large Top-K | `--moe-token-dispatcher-type allgather` | +#### MoonEP + +MoonEP is an optional external dependency and is not installed by Megatron Core. Its initial +integration supports eager BF16 training on one NVLink-connected host with expert tensor +parallelism set to 1. Enable it with: + +```bash +--moe-token-dispatcher-type flex +--moe-flex-dispatcher-backend moonep +--moe-grouped-gemm +--moe-single-grouped-weight +--use-transformer-engine-op-fuser +--gradient-accumulation-fusion +--disable-bias-linear +--moe-router-dtype fp32 +``` + +The supported expert activations are fused SwiGLU, quick-GeGLU, and weighted squared ReLU. Latent +MoE is supported with `--moe-latent-size`; routing remains in the model hidden dimension while +MoonEP dispatch, expert computation, and combine operate in the latent dimension. Packed and +GLU-interleaved FC1 layouts remain one contiguous projection. The registered +`moe_single_grouped_weight` parameters remain the optimizer and checkpoint source of truth; +MoonEP's `E+B` runtime weights and FP32 gradient storage are mirrors, so the distributed optimizer +and existing sharded checkpoint layout remain supported. + +The dispatched activation and probability tensors retain MoonEP's fixed `[NvS, H]` and `[NvS]` +capacity through the fused expert MLP. Device-side `E+B` counts delimit the live rows; unused tail +rows are ignored by grouped GEMM and combine. This is padding-only—the router's capacity-based +drop-and-pad mode remains disabled, and MoonEP does not drop routed tokens. Steady-state dispatch, +FC2 output/combine, combine-backward, and FC1 dgrad/dispatch-backward use symmetric token buffers +directly, without full-activation boundary copies. Per-forward dispatch buffers are pooled until +FC1 backward so multiple outstanding microbatches cannot overwrite saved activations; the pool +grows only to the maximum in-flight depth and is then reused. Router weights remain in independent +small tensors. Steady-state dispatch, expert compute, weight publication, and gradient reduction do +not read routing counts back to the CPU or use device-wide synchronization. + +MoonEP requires fixed, equal local token counts across its ranks, top-k no greater than 32, +128-element-aligned expert projection dimensions (including `moe_latent_size`, when set), and +VMM-aligned per-rank expert chunks. CUDA graphs, FP8/FP4, delayed wgrad, expert +communication/shared-expert overlap, and capacity-based dropping or padding are rejected. +Megatron's +`destroy_model_parallel()` tears down live MoonEP VMM buffers collectively before process-group +teardown. + ### Upcycling Use `--moe-use-upcycling` to enable upcycling, which loads the dense model from the `--load` directory, converts it to an MoE model at runtime, and starts training. The converted model is saved to the `--save` path before training begins. Upcycling is built on distributed checkpointing, supporting parallel modes different from existing dense checkpoints, such as arbitrary expert parallelism during upcycling. @@ -549,6 +594,7 @@ For MoE models, certain configurations may prevent CUDA Graph capture of MoE lay | Argument | Description | Default | |----------|-------------|---------| | --moe-token-dispatcher-type | Dispatcher: allgather, alltoall, flex | allgather | +| --moe-flex-dispatcher-backend | Flex backend: deepep, hybridep, moonep, ncclep | deepep | | --moe-enable-deepep | Enable DeepEP (with flex) | False | | --moe-expert-capacity-factor | Capacity factor | None | | --moe-pad-expert-input-to-capacity | Pad to capacity | False | diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 7fc0725b9de..248f57a5ced 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -3,6 +3,7 @@ import inspect import logging +import os from collections.abc import Callable from contextlib import nullcontext from copy import deepcopy @@ -208,6 +209,18 @@ def __init__( self.ep_group = pg_collection.ep self.tp_group = pg_collection.expt_tp + # Transformer Engine currently guards its single grouped parameter + # implementation with this environment variable in addition to the + # ``single_grouped_weight`` constructor argument. MoonEP cannot fall + # back to per-expert parameters because its VMM bridge requires one + # contiguous source parameter. + if self.config.moe_flex_dispatcher_backend == "moonep": + os.environ.setdefault("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") + # TE 2.15's CuTe DSL grouped-MLP wgrad kernel has a known + # correctness issue with single grouped weights. Keep the fused + # forward/dgrad path, but use TE's grouped-GEMM wgrad fallback. + os.environ.setdefault("NVTE_DISABLE_CUTEDSL_WGRAD_FUSED_GROUPED_MLP", "1") + # Double the output width with gated linear unit, see https://arxiv.org/pdf/2002.05202.pdf ffn_hidden_size = not_none(self.config.moe_ffn_hidden_size) if self.config.gated_linear_unit: @@ -215,14 +228,18 @@ def __init__( self.linear_fc1 = submodules.linear_fc1( self.num_local_experts, - self.input_size if self.config.moe_latent_size is None else self.config.moe_latent_size, + ( + self.input_size + if self.config.moe_latent_size is None + else self.config.moe_latent_size + ), ffn_hidden_size, config=self.config, init_method=not_none(self.config.init_method), bias=self.config.add_bias_linear, skip_bias_add=False, is_expert=True, - tp_comm_buffer_name='fc1', + tp_comm_buffer_name="fc1", pg_collection=pg_collection, name=(name + ".linear_fc1") if name is not None else None, ) @@ -245,7 +262,7 @@ def __init__( bias=self.config.add_bias_linear, skip_bias_add=True, is_expert=True, - tp_comm_buffer_name='fc2', + tp_comm_buffer_name="fc2", pg_collection=pg_collection, name=(name + ".linear_fc2") if name is not None else None, ) @@ -266,7 +283,7 @@ def __init__( ) self.activation_recompute = ( - self.config.recompute_granularity == 'selective' + self.config.recompute_granularity == "selective" and "moe_act" in self.config.recompute_modules ) if self.activation_recompute and (self.config.fp8 or self.config.fp4): @@ -289,6 +306,7 @@ def __init__( ), "Fused GroupedMLP is not supported for this configuration." self._with_fused_impl: bool = self.config.use_transformer_engine_op_fuser self._fused_ops: Optional[Tuple[torch.nn.Module]] = None + self._moonep_weight_bridge = None if ( self.config.gated_linear_unit and self.config.moe_mlp_glu_interleave_size is not None @@ -309,6 +327,12 @@ def __init__( self.num_local_experts, align_size=align_size ) + def set_moonep_weight_bridge(self, bridge) -> None: + """Use MoonEP's ``[E+B]`` runtime grouped weights for fused expert compute.""" + if self._fused_ops is not None: + raise RuntimeError("MoonEP weights must be bound before the first expert forward.") + self._moonep_weight_bridge = bridge + @staticmethod def _apply_bias(intermediate_parallel, bias_parallel, tokens_per_expert, permuted_probs): if bias_parallel is None: @@ -393,8 +417,6 @@ def _is_fused_impl_supported(self) -> bool: # Check TE CuTe DSL fused kernel conditions (must match TE's # fuse_grouped_mlp_ops matching logic). - import os - if use_glu_fusion and int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: return False return True @@ -432,8 +454,15 @@ def register_grouped_linear_params( for idx in range(linear.num_gemms): op.register_parameter(f"bias{idx}", linear.get_parameter(f"bias{idx}")) + def register_moonep_weight(op: torch.nn.Module, runtime_weight: torch.nn.Parameter) -> None: + """Attach an unregistered ``[E+B]`` MoonEP weight to a TE op shell.""" + op.register_parameter("weight", runtime_weight) + for idx in range(op.num_groups): + op.register_parameter(f"weight{idx}", None) + # Container for fusible ops ops = te.pytorch.ops.Sequential() + moonep_bridge = self._moonep_weight_bridge # Check if there are 1 or "num_gemms" params in the GroupedLinear module. fc1_single_grouped_weight = self.linear_fc1.single_grouped_weight @@ -456,27 +485,42 @@ def register_grouped_linear_params( # for runs that enable it via overlap_dispatch_backward_with_experts_wgrad. fc1_delay_wgrad_compute = self.linear_fc1.delay_wgrad_compute fc2_delay_wgrad_compute = self.linear_fc2.delay_wgrad_compute + fc1_num_gemms = ( + moonep_bridge.num_runtime_experts + if moonep_bridge is not None + else self.linear_fc1.num_gemms + ) + fc2_num_gemms = ( + moonep_bridge.num_runtime_experts + if moonep_bridge is not None + else self.linear_fc2.num_gemms + ) # Create a parameterless op shell and then attach the existing GroupedLinear weights below. # Using meta avoids allocating duplicate weights for the fused wrapper. op = te.pytorch.ops.GroupedLinear( - self.linear_fc1.num_gemms, + fc1_num_gemms, self.linear_fc1.in_features, self.linear_fc1.out_features, bias=self.linear_fc1.use_bias, device="meta", dtype=fc1_weight_dtype, accumulate_into_main_grad=self.linear_fc1.fuse_wgrad_accumulation, - single_grouped_weight=fc1_single_grouped_weight, + single_grouped_weight=( + True if moonep_bridge is not None else fc1_single_grouped_weight + ), single_grouped_bias=fc1_single_grouped_bias, delay_wgrad_compute=fc1_delay_wgrad_compute, ) # In single grouped mode, clear stale per-expert meta params so TE does not reset # the op and replace the shared DDP parameter with a fresh one lacking main_grad. - register_grouped_linear_params( - op, self.linear_fc1, fc1_single_grouped_weight, fc1_single_grouped_bias - ) + if moonep_bridge is not None: + register_moonep_weight(op, moonep_bridge.runtime_fc1_weight) + else: + register_grouped_linear_params( + op, self.linear_fc1, fc1_single_grouped_weight, fc1_single_grouped_bias + ) ops.append(op) # Activation and post-multiply probs (SwiGLU, clamped quick-GeGLU, or SReLU) @@ -494,32 +538,29 @@ def register_grouped_linear_params( else: op = te.pytorch.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave) elif self.config.activation_func == quick_gelu and self.config.gated_linear_unit: - clamp = self.config.activation_func_clamp_value - if clamp is not None: - if ( - "activation_recompute_in_mlp" - in inspect.signature(te.pytorch.ops.ScaledClampedQGeGLU).parameters - ): - op = te.pytorch.ops.ScaledClampedQGeGLU( - glu_interleave_size=glu_interleave, - activation_recompute_in_mlp=activation_recompute_in_mlp, - limit=clamp, - ) - else: - op = te.pytorch.ops.ScaledClampedQGeGLU( - glu_interleave_size=glu_interleave, limit=clamp - ) - else: - if ( - "activation_recompute_in_mlp" - in inspect.signature(te.pytorch.ops.ScaledClampedQGeGLU).parameters - ): - op = te.pytorch.ops.ScaledClampedQGeGLU( - glu_interleave_size=glu_interleave, - activation_recompute_in_mlp=activation_recompute_in_mlp, - ) - else: - op = te.pytorch.ops.ScaledClampedQGeGLU(glu_interleave_size=glu_interleave) + qgeglu_signature = inspect.signature(te.pytorch.ops.ScaledClampedQGeGLU) + qgeglu_kwargs = { + "glu_interleave_size": glu_interleave, + # Megatron's None means no clamp, whereas TE defaults to 7. + "limit": ( + float("inf") + if self.config.activation_func_clamp_value is None + else self.config.activation_func_clamp_value + ), + } + if "alpha" in qgeglu_signature.parameters: + qgeglu_kwargs["alpha"] = 1.702 + if "glu_linear_offset" in qgeglu_signature.parameters: + # Newer TE defaults to 1, while Megatron defaults to 0. + qgeglu_kwargs["glu_linear_offset"] = self.config.glu_linear_offset + elif self.config.glu_linear_offset != 0.0: + raise RuntimeError( + "The installed Transformer Engine ScaledClampedQGeGLU does not support " + "Megatron's nonzero glu_linear_offset." + ) + if "activation_recompute_in_mlp" in qgeglu_signature.parameters: + qgeglu_kwargs["activation_recompute_in_mlp"] = activation_recompute_in_mlp + op = te.pytorch.ops.ScaledClampedQGeGLU(**qgeglu_kwargs) elif ( self.config.activation_func == squared_relu and self.config.use_fused_weighted_squared_relu @@ -544,14 +585,16 @@ def register_grouped_linear_params( # FC2 fc2_bias_kwargs = {"scale_bias": True} if self.linear_fc2.use_bias else {} op = te.pytorch.ops.GroupedLinear( - self.linear_fc2.num_gemms, + fc2_num_gemms, self.linear_fc2.in_features, self.linear_fc2.out_features, bias=self.linear_fc2.use_bias, device="meta", dtype=fc2_weight_dtype, accumulate_into_main_grad=self.linear_fc2.fuse_wgrad_accumulation, - single_grouped_weight=fc2_single_grouped_weight, + single_grouped_weight=( + True if moonep_bridge is not None else fc2_single_grouped_weight + ), single_grouped_bias=fc2_single_grouped_bias, delay_wgrad_compute=fc2_delay_wgrad_compute, # Preserve p * (FC2(x) + bias) after the scaled activation moves p before FC2. @@ -560,9 +603,12 @@ def register_grouped_linear_params( # In single grouped mode, clear stale per-expert meta params so TE does not reset # the op and replace the shared DDP parameter with a fresh one lacking main_grad. - register_grouped_linear_params( - op, self.linear_fc2, fc2_single_grouped_weight, fc2_single_grouped_bias - ) + if moonep_bridge is not None: + register_moonep_weight(op, moonep_bridge.runtime_fc2_weight) + else: + register_grouped_linear_params( + op, self.linear_fc2, fc2_single_grouped_weight, fc2_single_grouped_bias + ) ops.append(op) # Emulate submodule pre-forward hooks @@ -593,6 +639,9 @@ def forward_pre_hook(module, *_) -> None: "has a pre-forward hook that modifies the input tensor." ) self._ensure_main_grad_for_fused_impl() + if self._moonep_weight_bridge is not None: + self._moonep_weight_bridge.prepare_forward() + self._moonep_weight_bridge.prefetch(self._moonep_weight_bridge.last_plan) return forward_pre_hook @@ -676,7 +725,7 @@ def _fused_forward( fine_grained_activation_offloading, permuted_local_hidden_states, offload_name ) with fused_group_mlp_manager as permuted_local_hidden_states: - # NCCL-EP zero-copy is active exactly when ``output_buffer`` is not None, and then the + # Backend zero-copy is active exactly when ``output_buffer`` is not None, and then the # fused-MLP input aliases the persistent symm buffer (also the fc2 output combine # reads), whose storage is non-resizable — so skip the force-release in that case. forced_released_tensors = ( @@ -685,7 +734,7 @@ def _fused_forward( else [] ) with stash_context: - # NCCL-EP zero-copy: route the fc2 output (fwd combine reads it one-sided) and the + # Backend zero-copy: route the fc2 output (fwd combine reads it one-sided) and the # fc1 dgrad (bwd dispatch scatters it one-sided) into caller-provided symm buffers. # op_kwargs keys are basic-op indices into [fc1, activation, fc2]: 0=fc1, -1=fc2. op_kwargs = {} @@ -741,9 +790,9 @@ def forward( 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. + (backend 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. + into (backend zero-copy bwd dispatch); only the fused op-fuser path supports it. Return: output (torch.Tensor): The output of the local experts. @@ -910,7 +959,7 @@ def glu(x): return output, output_bias def sharded_state_dict( - self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None + self, prefix: str = "", sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: """ Maps local expert to global experts. @@ -918,13 +967,13 @@ def sharded_state_dict( """ # Guard for cases metadata is not provided metadata = ensure_metadata_has_dp_cp_group(metadata) - singleton_local_shards = (metadata or {}).get('singleton_local_shards', False) + singleton_local_shards = (metadata or {}).get("singleton_local_shards", False) sharded_state_dict = {} for name, module in self._modules.items(): sub_sd = sharded_state_dict_default( - module, f'{name}.', sharded_offsets, metadata, tp_group=self.tp_group + module, f"{name}.", sharded_offsets, metadata, tp_group=self.tp_group ) - if name == 'linear_fc1' and self.config.gated_linear_unit: + if name == "linear_fc1" and self.config.gated_linear_unit: num_global_experts = self.ep_group.size() * self.num_local_experts local_expert_indices_offset = self.ep_group.rank() * self.num_local_experts ep_axis = len(sharded_offsets) @@ -936,20 +985,20 @@ def sharded_state_dict( *sharded_offsets, (ep_axis, local_expert_indices_offset + i, num_global_experts), ) - for k in (f'{name}.weight{i}', f'{name}.bias{i}'): + for k in (f"{name}.weight{i}", f"{name}.bias{i}"): if k in sub_sd: sub_sd[k] = apply_swiglu_sharded_factory( sub_sd[k], new_sharded_offsets, singleton_local_shards, tp_group=self.tp_group, - dp_group=metadata['dp_cp_group'], + dp_group=metadata["dp_cp_group"], ) if singleton_local_shards: - replace_prefix_for_sharding(sub_sd, '', f'{prefix}experts.') + replace_prefix_for_sharding(sub_sd, "", f"{prefix}experts.") else: # Add prefix here to match sequential's keys - replace_prefix_for_sharding(sub_sd, f'{name}.', f'{prefix}experts.{name}.') + replace_prefix_for_sharding(sub_sd, f"{name}.", f"{prefix}experts.{name}.") sharded_state_dict.update({f"{prefix}{k}": v for k, v in sub_sd.items()}) return sharded_state_dict @@ -1026,7 +1075,7 @@ def __init__( self._mcore_activation_type = self._resolve_mcore_activation_type() self.inference_grouped_gemm_backend = config.inference_grouped_gemm_backend - self._nvls_dispatcher = config.inference_moe_token_dispatcher_type == 'nvls' + self._nvls_dispatcher = config.inference_moe_token_dispatcher_type == "nvls" def _resolve_flashinfer_activation_type(self): """Map megatron activation config to FlashInfer ActivationType.""" @@ -1067,14 +1116,14 @@ def _build_concatenated_mxfp8_weights(self): intended for non-colocated inference. """ - for linear_name, buf_name in [('linear_fc1', '_fc1_weight'), ('linear_fc2', '_fc2_weight')]: + for linear_name, buf_name in [("linear_fc1", "_fc1_weight"), ("linear_fc2", "_fc2_weight")]: linear = getattr(self, linear_name) q_list, s_list = [], [] for i in range(self.num_local_experts): - w = getattr(linear, f'weight{i}') + w = getattr(linear, f"weight{i}") if isinstance(w, MXFP8Tensor): mxfp8 = w - elif hasattr(w, 'data') and isinstance(w.data, MXFP8Tensor): + elif hasattr(w, "data") and isinstance(w.data, MXFP8Tensor): mxfp8 = w.data else: raise RuntimeError( @@ -1093,11 +1142,11 @@ def _build_concatenated_mxfp8_weights(self): # mirroring _build_concatenated_weights. This frees the original # allocations while keeping the Parameter objects intact. for i in range(self.num_local_experts): - w = getattr(linear, f'weight{i}') + w = getattr(linear, f"weight{i}") if isinstance(w, MXFP8Tensor): w.data = stacked_data[i] w.scale = stacked_scale[i] - elif hasattr(w, 'data') and isinstance(w.data, MXFP8Tensor): + elif hasattr(w, "data") and isinstance(w.data, MXFP8Tensor): w.data.data = stacked_data[i] w.data.scale = stacked_scale[i] @@ -1129,8 +1178,8 @@ def _build_concatenated_weights(self): # Copy existing TE weights into big tensors, then point param.data to the views for i in range(self.num_local_experts): - fc1_param = getattr(self.linear_fc1, f'weight{i}') - fc2_param = getattr(self.linear_fc2, f'weight{i}') + fc1_param = getattr(self.linear_fc1, f"weight{i}") + fc2_param = getattr(self.linear_fc2, f"weight{i}") # Copy initialized data into contiguous buffer _fc1_weight[i].copy_(fc1_param.data) @@ -1142,8 +1191,8 @@ def _build_concatenated_weights(self): fc2_param.data = _fc2_weight[i] # Register big tensors as non-persistent buffers (for .to() device movement, not saved) - self.register_buffer('_fc1_weight', _fc1_weight, persistent=False) - self.register_buffer('_fc2_weight', _fc2_weight, persistent=False) + self.register_buffer("_fc1_weight", _fc1_weight, persistent=False) + self.register_buffer("_fc2_weight", _fc2_weight, persistent=False) def _flashinfer_forward(self, hidden_states, routing_map, probs): """FlashInfer fused MoE kernel for CUDA-graphed inference iterations.""" @@ -1160,7 +1209,7 @@ def _flashinfer_forward(self, hidden_states, routing_map, probs): activation_type=self._flashinfer_activation_type, ep_size=self.ep_group.size(), ep_rank=self.ep_group.rank(), - output=NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None, + output=(NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None), )[0] return output, None @@ -1178,7 +1227,7 @@ def _mcore_fused_moe_forward(self, hidden_states, probs, routing_map): valid_tokens=InferenceAllGatherDispatcherBase._valid_tokens(), routing_map=routing_map, disable_fused_quant_kernels=self.config.inference_moe_disable_fused_quant_kernels, - out=NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None, + out=(NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None), ) return output, None @@ -1195,7 +1244,7 @@ def _vllm_forward(self, hidden_states, probs, routing_map): local_expert_start=local_expert_start, valid_tokens=InferenceAllGatherDispatcherBase._valid_tokens(), routing_map=routing_map, - out=NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None, + out=(NVLSAllGatherVDispatcher._get_rsv_tensor() if self._nvls_dispatcher else None), num_tokens_hint=InferenceAllGatherDispatcherBase._get_host_valid_tokens_estimate(), ) return output, None @@ -1234,7 +1283,7 @@ def forward( if not self._concatenated_weights_built: w = self.linear_fc1.weight0 if isinstance(w, MXFP8Tensor) or ( - hasattr(w, 'data') and isinstance(w.data, MXFP8Tensor) + hasattr(w, "data") and isinstance(w.data, MXFP8Tensor) ): self._build_concatenated_mxfp8_weights() else: @@ -1297,7 +1346,7 @@ def __init__( ffn_hidden_size=self.config.moe_ffn_hidden_size, is_expert=True, tp_group=pg_collection.expt_tp, - name=(name + f".local_experts.{expert_idx}") if name is not None else None, + name=((name + f".local_experts.{expert_idx}") if name is not None else None), ) self.local_experts.append(expert) @@ -1374,7 +1423,7 @@ def backward_dw(self): for expert in self.local_experts: expert.backward_dw() - def sharded_state_dict(self, prefix='', sharded_offsets=(), metadata=None): + def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): """Maps local expert to global experts.""" # Guard for cases metadata is not provided metadata = ensure_metadata_has_dp_cp_group(metadata) @@ -1383,16 +1432,16 @@ def sharded_state_dict(self, prefix='', sharded_offsets=(), metadata=None): num_global_experts = self.ep_group.size() * self.num_local_experts local_expert_indices_offset = self.ep_group.rank() * self.num_local_experts - singleton_local_shards = (metadata or {}).get('singleton_local_shards', False) + singleton_local_shards = (metadata or {}).get("singleton_local_shards", False) for expert_local_idx, expert in enumerate(self.local_experts): expert_global_idx = local_expert_indices_offset + expert_local_idx - expert_state_dict_prefix = f'{prefix}local_experts.{expert_local_idx}.' + expert_state_dict_prefix = f"{prefix}local_experts.{expert_local_idx}." if singleton_local_shards: - expert_sharded_prefix = f'{prefix}experts.{expert_global_idx}.' + expert_sharded_prefix = f"{prefix}experts.{expert_global_idx}." expert_sharded_offsets = sharded_offsets else: - expert_sharded_prefix = f'{prefix}experts.' + expert_sharded_prefix = f"{prefix}experts." expert_sharded_offsets = ( *sharded_offsets, (len(sharded_offsets), expert_global_idx, num_global_experts), @@ -1410,7 +1459,7 @@ def sharded_state_dict(self, prefix='', sharded_offsets=(), metadata=None): replica_id = sh_ten.replica_id assert ( len(replica_id) == 3 - ), f'Expected replica_id for {k} to be in (PP, TP, DP) format, got: {replica_id}' + ), f"Expected replica_id for {k} to be in (PP, TP, DP) format, got: {replica_id}" sh_ten.replica_id = (*replica_id[:2], self.dp_group.rank()) diff --git a/megatron/core/transformer/moe/fused_a2a.py b/megatron/core/transformer/moe/fused_a2a.py index 444760655a6..76629cd5070 100644 --- a/megatron/core/transformer/moe/fused_a2a.py +++ b/megatron/core/transformer/moe/fused_a2a.py @@ -3,6 +3,9 @@ # Copyright (c) 2025 DeepSeek # Licensed under the MIT License - https://github.com/deepseek-ai/DeepEP/blob/main/LICENSE +import os +import weakref +from dataclasses import dataclass from typing import Optional from megatron.core.utils import internal_api @@ -19,6 +22,689 @@ _buffer = None +try: + from moonep import Buffer as MoonEPBuffer + from moonep._C import nvl_dist_alloc as _moonep_nvl_dist_alloc + from moonep._C import nvl_dist_map as _moonep_nvl_dist_map + from moonep._C import nvl_release_mem_handle as _moonep_nvl_release_mem_handle + from moonep.buffer import _exchange_ipc_fds as _moonep_exchange_ipc_fds + from moonep.buffer import create_nvl_dist_tensor as _moonep_create_nvl_dist_tensor + from moonep.buffer import get_vmm_granularity as _moonep_get_vmm_granularity + from moonep.grad_reduce import launch_grad_reduce as _moonep_launch_grad_reduce + from moonep.inter_rank_sync import launch_inter_rank_sync as _moonep_inter_rank_sync + from moonep.prefetch import launch_prefetch as _moonep_launch_prefetch + + HAVE_MOONEP = True + _MOONEP_IMPORT_ERROR = None +except ImportError as exc: + MoonEPBuffer = None + HAVE_MOONEP = False + _MOONEP_IMPORT_ERROR = exc + + +_moonep_buffers = weakref.WeakSet() +_moonep_bridges = weakref.WeakSet() +_moonep_dispatch_buffer_pools = weakref.WeakSet() +_moonep_token_buffer_pools = {} + + +def is_moonep_available() -> bool: + """Return whether the optional MoonEP package was imported successfully.""" + return HAVE_MOONEP + + +def new_moonep_buffer(**kwargs): + """Create and register a MoonEP buffer for explicit collective teardown.""" + if not HAVE_MOONEP: + raise ImportError( + "MoonEP is not installed. Install the optional 'moonep' package before using " + "moe_flex_dispatcher_backend='moonep'." + ) from _MOONEP_IMPORT_ERROR + buffer = MoonEPBuffer(explicitly_destroy=True, **kwargs) + _moonep_buffers.add(buffer) + return buffer + + +def _allocate_moonep_token_buffer(ctx): + """Collectively allocate one symmetric hidden-token buffer pair.""" + rank = int(ctx["rank"]) + world_size = int(ctx["R"]) + num_slots = int(ctx["NvS"]) + padded_slots = int(ctx["NvS_padded"]) + full = _moonep_create_nvl_dist_tensor( + [padded_slots, int(ctx["H"])], torch.bfloat16, rank, world_size, group=ctx["group"] + ) + local = full[rank * padded_slots : rank * padded_slots + num_slots] + return full, local + + +class MoonEPDispatchBufferPool: + """Pool per-forward symmetric dispatch outputs until their FC1 backward.""" + + def __init__(self, buffer): + self._ctx = buffer._require_ctx() + self._free = [(self._ctx["hidden_buf"], self._ctx["hidden_buf_local"])] + self._allocated = list(self._free) + self._destroyed = False + _moonep_dispatch_buffer_pools.add(self) + + def acquire(self): + """Acquire a buffer, growing collectively to the maximum in-flight depth.""" + if self._destroyed: + raise RuntimeError("MoonEP dispatch buffer pool has been destroyed.") + if self._free: + return self._free.pop() + pair = _allocate_moonep_token_buffer(self._ctx) + self._allocated.append(pair) + return pair + + def release(self, pair) -> None: + """Recycle a dispatch buffer after dispatch backward consumes FC1 dgrad.""" + if not self._destroyed: + self._free.append(pair) + + def destroy(self) -> None: + """Drop all VMM tensor references after MoonEP work has synchronized.""" + if self._destroyed: + return + self._free.clear() + self._allocated.clear() + self._ctx = None + self._destroyed = True + + +class _MoonEPSharedTokenBufferPool: + """Own the two process-group-wide transient expert boundary buffers.""" + + def __init__(self, ctx): + self._buffers = tuple(_allocate_moonep_token_buffer(ctx) for _ in range(2)) + + @property + def forward(self): + """FC2-output / combine-backward buffer pair.""" + return self._buffers[0] + + @property + def backward(self): + """FC1-dgrad / dispatch-backward buffer pair.""" + return self._buffers[1] + + def destroy(self) -> None: + """Drop VMM tensor references after MoonEP work has synchronized.""" + self._buffers = () + + +def get_moonep_zero_copy_token_buffers(buffer): + """Return shared symmetric buffers for MoonEP's two transient expert boundaries. + + Dispatch output uses a separate per-forward pool because FC1 autograd saves + it. FC2 output/combine-backward and FC1 dgrad/dispatch-backward have + non-overlapping lifetimes across layers, so two process-group-wide buffers + are sufficient and avoid allocating those two boundaries per layer. + """ + ctx = buffer._require_ctx() + group = ctx["group"] + key = ( + id(group), + str(ctx["device"]), + int(ctx["R"]), + int(ctx["NvS"]), + int(ctx["NvS_padded"]), + int(ctx["H"]), + ) + pool = _moonep_token_buffer_pools.get(key) + if pool is not None: + return pool + + pool = _MoonEPSharedTokenBufferPool(ctx) + _moonep_token_buffer_pools[key] = pool + return pool + + +def moonep_finalize() -> None: + """Destroy all live MoonEP buffers and runtime VMM mappings.""" + for buffer in list(_moonep_buffers): + buffer.destroy() + _moonep_buffers.clear() + for pool in list(_moonep_dispatch_buffer_pools): + pool.destroy() + _moonep_dispatch_buffer_pools.clear() + for bridge in list(_moonep_bridges): + bridge.destroy() + _moonep_bridges.clear() + for pool in _moonep_token_buffer_pools.values(): + pool.destroy() + _moonep_token_buffer_pools.clear() + + +def _close_fds(fds) -> None: + """Close a collection of POSIX file descriptors exactly once.""" + for fd in set(fds): + os.close(fd) + + +def _allocate_moonep_mapping(*, chunk_shape, dtype, group, with_reduce_view: bool): + """Allocate an ``[E+B]`` VMM mapping and an optional all-rank slot view. + + Each rank owns one expert chunk and one equally-sized prefetch/gradient-slot + chunk. The returned composite maps all expert chunks followed by this + rank's slot chunk. For gradients, the second returned tensor maps every + rank's slot chunk as ``[R, B, ...]`` for MoonEP's owner-side reduction. + """ + rank = torch.distributed.get_rank(group=group) + world_size = torch.distributed.get_world_size(group=group) + chunk_bytes = int(torch.tensor([], dtype=dtype).element_size()) + for dim in chunk_shape: + chunk_bytes *= int(dim) + granularity = int(_moonep_get_vmm_granularity()) + if chunk_bytes % granularity != 0: + raise ValueError( + "MoonEP expert chunks must be VMM aligned: " + f"shape={tuple(chunk_shape)}, dtype={dtype}, bytes={chunk_bytes}, " + f"granularity={granularity}." + ) + + expert_keepalive, expert_fd, expert_handle = _moonep_nvl_dist_alloc( + shape=list(chunk_shape), dtype=dtype + ) + slot_keepalive, slot_fd, slot_handle = _moonep_nvl_dist_alloc( + shape=list(chunk_shape), dtype=dtype + ) + _moonep_nvl_release_mem_handle(expert_handle) + _moonep_nvl_release_mem_handle(slot_handle) + + expert_fds = _moonep_exchange_ipc_fds( + expert_fd, list(range(world_size)), rank, world_size, group + ) + os.close(expert_fd) + + slot_fds = None + if with_reduce_view: + slot_fds = _moonep_exchange_ipc_fds( + slot_fd, list(range(world_size)), rank, world_size, group + ) + os.close(slot_fd) + local_slot_fd = slot_fds[rank] + else: + local_slot_fd = slot_fd + + expert_fd_list = [expert_fds[idx] for idx in range(world_size)] + full = _moonep_nvl_dist_map( + chunk_shape=list(chunk_shape), + dtype=dtype, + fds=[*expert_fd_list, local_slot_fd], + local_rank=rank, + world_size=world_size + 1, + ) + + reduce_view = None + fds_to_close = [*expert_fd_list, local_slot_fd] + if slot_fds is not None: + slot_fd_list = [slot_fds[idx] for idx in range(world_size)] + reduce_view = _moonep_nvl_dist_map( + chunk_shape=list(chunk_shape), + dtype=dtype, + fds=slot_fd_list, + local_rank=rank, + world_size=world_size, + ) + fds_to_close.extend(slot_fd_list) + _close_fds(fds_to_close) + + # The mappings own their virtual addresses. Keep the local physical + # allocations alive for exactly as long as the bridge owns the mappings. + return full, reduce_view, (expert_keepalive, slot_keepalive) + + +def _allocate_moonep_grad_mapping(*, chunk_shape, group): + """Allocate a rank-private ``[E+B]`` wgrad view and shared slot view. + + The local owner chunk occupies this rank's global expert range. All + nonlocal expert ranges alias a private disposable sink chunk, so TE can + write zero-token group outputs without touching peer-owned gradients. The + final chunk is the local redundant-slot storage and is also mapped from + every rank as ``[R, B, ...]`` for MoonEP's owner-side reducer. + """ + rank = torch.distributed.get_rank(group=group) + world_size = torch.distributed.get_world_size(group=group) + dtype = torch.float32 + + keepalives = [] + local_fds = [] + for _ in range(3): + keepalive, fd, handle = _moonep_nvl_dist_alloc(shape=list(chunk_shape), dtype=dtype) + _moonep_nvl_release_mem_handle(handle) + keepalives.append(keepalive) + local_fds.append(fd) + owner_fd, sink_fd, slot_fd = local_fds + + slot_fds = _moonep_exchange_ipc_fds(slot_fd, list(range(world_size)), rank, world_size, group) + os.close(slot_fd) + local_slot_fd = slot_fds[rank] + + grad_fds = [sink_fd] * world_size + grad_fds[rank] = owner_fd + grad_fds.append(local_slot_fd) + full_grad = _moonep_nvl_dist_map( + chunk_shape=list(chunk_shape), + dtype=dtype, + fds=grad_fds, + local_rank=rank, + world_size=world_size + 1, + ) + + slot_fd_list = [slot_fds[idx] for idx in range(world_size)] + reduce_view = _moonep_nvl_dist_map( + chunk_shape=list(chunk_shape), + dtype=dtype, + fds=slot_fd_list, + local_rank=rank, + world_size=world_size, + ) + _close_fds([owner_fd, sink_fd, *slot_fd_list]) + return full_grad, reduce_view, tuple(keepalives) + + +@dataclass +class _MoonEPProjection: + """MoonEP runtime storage for one grouped expert projection.""" + + linear: torch.nn.Module + parameter: torch.nn.Parameter + full_weight: torch.Tensor + full_grad: torch.Tensor + reduce_buffers: torch.Tensor + runtime_parameter: torch.nn.Parameter + dummy_grad: torch.Tensor + keepalives: tuple + + +class MoonEPWeightBridge: + """Connect Megatron grouped expert parameters to MoonEP VMM runtime weights. + + Registered Megatron parameters remain the optimizer/checkpoint source of + truth. Their contiguous grouped storage is copied into rank-owned MoonEP + source chunks before each dispatch. Transformer Engine executes against an + unregistered ``[E+B]`` GroupedTensor whose FP32 ``main_grad`` points at the + corresponding MoonEP gradient mapping. + """ + + def __init__( + self, + *, + experts, + group: torch.distributed.ProcessGroup, + num_experts: int, + num_local_experts: int, + num_sms: Optional[int], + ) -> None: + if not HAVE_MOONEP: + raise ImportError( + "MoonEP is not installed. Install the optional 'moonep' package before using " + "moe_flex_dispatcher_backend='moonep'." + ) from _MOONEP_IMPORT_ERROR + + from transformer_engine.pytorch.tensor.grouped_tensor import GroupedTensor + + self.group = group + self.rank = torch.distributed.get_rank(group=group) + self.world_size = torch.distributed.get_world_size(group=group) + self.num_experts = int(num_experts) + self.num_local_experts = int(num_local_experts) + self.num_slots = self.num_local_experts + self.num_runtime_experts = self.num_experts + self.num_slots + self.num_sms = 32 if num_sms is None else int(num_sms) + self.buffer = None + self.last_plan = None + self._experts_ref = weakref.ref(experts) + self._destroyed = False + + if self.num_experts != self.world_size * self.num_local_experts: + raise ValueError( + "MoonEP requires an even expert distribution: " + f"num_experts={self.num_experts}, world_size={self.world_size}, " + f"num_local_experts={self.num_local_experts}." + ) + + self.projections = [] + for linear in (experts.linear_fc1, experts.linear_fc2): + parameter = dict(linear.named_parameters(recurse=False)).get("weight") + if parameter is None: + raise ValueError( + "MoonEP requires Transformer Engine to create one contiguous grouped " + "weight parameter. Ensure moe_single_grouped_weight=True and " + "NVTE_GROUPED_LINEAR_SINGLE_PARAM is not explicitly disabled." + ) + rowwise_data = getattr(parameter, "rowwise_data", None) + if rowwise_data is None or rowwise_data.dtype != torch.bfloat16: + raise ValueError( + "MoonEP requires BF16 moe_single_grouped_weight parameters with contiguous " + "rowwise_data." + ) + member_shape = (int(linear.out_features), int(linear.in_features)) + expected_numel = self.num_local_experts * member_shape[0] * member_shape[1] + if rowwise_data.numel() != expected_numel or not rowwise_data.is_contiguous(): + raise ValueError( + "MoonEP grouped parameter storage has an unexpected layout: " + f"expected {self.num_local_experts}x{member_shape}, " + f"got numel={rowwise_data.numel()}, contiguous={rowwise_data.is_contiguous()}." + ) + if member_shape[0] % 128 != 0 or member_shape[1] % 128 != 0: + raise ValueError( + "MoonEP weight prefetch requires both projection dimensions to be multiples " + f"of 128, got {member_shape}." + ) + + chunk_shape = (self.num_local_experts, *member_shape) + full_weight, _, weight_keepalives = _allocate_moonep_mapping( + chunk_shape=chunk_shape, + dtype=torch.bfloat16, + group=self.group, + with_reduce_view=False, + ) + full_grad, reduce_full, grad_keepalives = _allocate_moonep_grad_mapping( + chunk_shape=chunk_shape, group=self.group + ) + full_weight = full_weight.view(self.num_runtime_experts, *member_shape) + full_grad = full_grad.view(self.num_runtime_experts, *member_shape) + reduce_buffers = reduce_full.view(self.world_size, self.num_slots, *member_shape) + full_grad[ + self.rank * self.num_local_experts : (self.rank + 1) * self.num_local_experts + ].zero_() + full_grad[self.num_experts :].zero_() + + grouped_weight = GroupedTensor.make_grouped_tensor_from_rowwise_data( + num_tensors=self.num_runtime_experts, + tensor_shape=member_shape, + rowwise_data=full_weight, + dtype=torch.bfloat16, + ) + grouped_weight.requires_grad_(True) + runtime_parameter = torch.nn.Parameter(grouped_weight) + runtime_parameter.main_grad = full_grad + runtime_parameter.grad_added_to_main_grad = True + # Nonlocal rows alias a rank-private sink, so overwrite mode + # cannot corrupt another rank's owner gradients. + runtime_parameter.overwrite_main_grad = True + + # This cached zero tensor exists only to run the registered + # parameter's AccumulateGrad/DDP hook. The real gradient has + # already been accumulated into parameter.main_grad. + dummy_grad = torch.zeros_like(rowwise_data).view(parameter.shape) + self.projections.append( + _MoonEPProjection( + linear=linear, + parameter=parameter, + full_weight=full_weight, + full_grad=full_grad, + reduce_buffers=reduce_buffers, + runtime_parameter=runtime_parameter, + dummy_grad=dummy_grad, + keepalives=(*weight_keepalives, *grad_keepalives), + ) + ) + _moonep_bridges.add(self) + + @property + def runtime_fc1_weight(self) -> torch.nn.Parameter: + """Return the ``[E+B]`` FC1 runtime grouped parameter.""" + return self.projections[0].runtime_parameter + + @property + def runtime_fc2_weight(self) -> torch.nn.Parameter: + """Return the ``[E+B]`` FC2 runtime grouped parameter.""" + return self.projections[1].runtime_parameter + + @property + def source_parameters(self): + """Return Megatron's registered FC1/FC2 grouped parameters.""" + return tuple(projection.parameter for projection in self.projections) + + @property + def dummy_grads(self): + """Return cached dummy grads used to trigger registered-parameter hooks.""" + return tuple(projection.dummy_grad for projection in self.projections) + + def attach_buffer(self, buffer) -> None: + """Attach the layer's MoonEP communication buffer.""" + self.buffer = buffer + + def destroy(self) -> None: + """Release runtime grouped weights and their VMM mappings.""" + if self._destroyed: + return + experts = self._experts_ref() + if experts is not None: + experts._fused_ops = None + experts._moonep_weight_bridge = None + for projection in self.projections: + projection.runtime_parameter.main_grad = None + self.projections.clear() + self.buffer = None + self.last_plan = None + self._destroyed = True + + def prepare_forward(self) -> None: + """Refresh local source weights and clear this rank's gradient scratch.""" + local_start = self.rank * self.num_local_experts + local_end = local_start + self.num_local_experts + for projection in self.projections: + source = projection.parameter.rowwise_data.view_as( + projection.full_weight[local_start:local_end] + ) + projection.full_weight[local_start:local_end].copy_(source) + projection.full_grad[local_start:local_end].zero_() + projection.full_grad[self.num_experts :].zero_() + # Distributed-optimizer parameter all-gathers are run by the original + # linear pre-forward hooks immediately before this method. Publish every + # rank's local mirror refresh before any peer starts remote prefetch, + # entirely on the current CUDA stream. + _moonep_inter_rank_sync(self.buffer._require_ctx()) + + def prefetch(self, plan) -> None: + """Prefetch the plan's redundant FC1/FC2 experts into local slots.""" + experts_to_copy = plan.experts_to_copy[self.rank].contiguous() + for projection in self.projections: + _moonep_launch_prefetch( + projection.full_weight[: self.num_experts], + projection.full_weight[self.num_experts :], + experts_to_copy, + num_sms=self.num_sms, + ) + + def reduce_grads(self, plan) -> None: + """Reduce redundant wgrads and hand local results to Megatron DDP.""" + if self.buffer is None: + raise RuntimeError("MoonEPWeightBridge has no attached communication buffer.") + ctx = self.buffer._require_ctx() + local_start = self.rank * self.num_local_experts + local_end = local_start + self.num_local_experts + + # The TE op obeys PyTorch stream semantics: when its backward returns, + # its wgrad writes are ordered before subsequent work on the current + # stream. Align the EP ranks on-device before any reducer remote-reads + # peer slots. No host wait or device-wide fence is needed. + _moonep_inter_rank_sync(ctx) + + for projection in self.projections: + main_grad = getattr(projection.parameter, "main_grad", None) + if main_grad is None: + raise RuntimeError( + "MoonEP requires gradient-accumulation fusion and an initialized " + "parameter.main_grad buffer." + ) + _moonep_launch_grad_reduce( + projection.full_grad[: self.num_experts], + projection.reduce_buffers, + plan.experts_to_copy, + rank=self.rank, + num_sms=self.num_sms, + meta_buf=ctx["meta_buf"], + meta_stride=int(ctx["meta_chunk_padded"]), + barrier_off=int(ctx["BARRIER_OFF"]), + grid_sync_bar=ctx["grid_sync_bar"], + ) + + # FC1 and FC2 reducers are launched consecutively on one stream, + # matching MoonEP Buffer.reduce_grad(). Each reducer contains its + # own GPU-side cross-rank barrier and resets the shared barrier state. + main_grad.add_(projection.full_grad[local_start:local_end].view_as(main_grad)) + projection.full_grad[local_start:local_end].zero_() + projection.full_grad[self.num_experts :].zero_() + projection.parameter.grad_added_to_main_grad = True + + +class MoonEPDispatch(torch.autograd.Function): + """Autograd-aware MoonEP dispatch, probability gather, and wgrad reduction.""" + + @staticmethod + def forward( + ctx, + hidden_states, + topk_probs, + topk_indices, + tokens_per_expert, + fc1_parameter, + fc2_parameter, + buffer, + bridge, + dispatch_buffer_pool, + dgrad_hidden_buffer, + ): + """Dispatch activations and route weights while saving the MoonEP plan.""" + dispatch_hidden_buffer = ( + dispatch_buffer_pool.acquire() if dispatch_buffer_pool is not None else None + ) + try: + dispatched, dispatched_probs, cu_seqlens, plan = buffer.dispatch( + hidden_states.contiguous(), + topk_probs.float().contiguous(), + topk_indices.to(dtype=torch.int32).contiguous(), + tokens_per_expert.to(dtype=torch.int32).contiguous(), + zero_copy=dispatch_hidden_buffer is not None, + zero_copy_weights=False, + hidden_buffer=dispatch_hidden_buffer, + ) + except Exception: + if dispatch_buffer_pool is not None: + dispatch_buffer_pool.release(dispatch_hidden_buffer) + raise + bridge.last_plan = plan + + starts = torch.cat( + [torch.zeros(1, dtype=cu_seqlens.dtype, device=cu_seqlens.device), cu_seqlens[:-1]] + ) + runtime_tokens_per_expert = (cu_seqlens - starts).to(torch.int64) + ctx.buffer = buffer + ctx.bridge = bridge + ctx.plan = plan + ctx.dispatch_buffer_pool = dispatch_buffer_pool + ctx.dispatch_hidden_buffer = dispatch_hidden_buffer + ctx.dgrad_hidden_buffer = dgrad_hidden_buffer + ctx.mark_non_differentiable(runtime_tokens_per_expert) + return dispatched, dispatched_probs, runtime_tokens_per_expert + + @staticmethod + def backward(ctx, grad_hidden, grad_probs, _grad_tokens_per_expert): + """Combine activation/probability gradients and reduce duplicated wgrads.""" + try: + ctx.bridge.reduce_grads(ctx.plan) + grad_hidden = grad_hidden.contiguous() + use_zero_copy = ( + ctx.dgrad_hidden_buffer is not None + and grad_hidden.data_ptr() == ctx.dgrad_hidden_buffer[1].data_ptr() + ) + grad_hidden_states, grad_topk_probs, _ = ctx.buffer.combine( + plan=ctx.plan, + hidden_nvsh=grad_hidden, + route_weights_nvs=grad_probs.float().contiguous(), + zero_copy=use_zero_copy, + hidden_buffer=ctx.dgrad_hidden_buffer, + ) + finally: + if ctx.dispatch_buffer_pool is not None: + ctx.dispatch_buffer_pool.release(ctx.dispatch_hidden_buffer) + dummy_fc1_grad, dummy_fc2_grad = ctx.bridge.dummy_grads + return ( + grad_hidden_states, + grad_topk_probs, + None, + None, + dummy_fc1_grad, + dummy_fc2_grad, + None, + None, + None, + None, + ) + + +class MoonEPCombine(torch.autograd.Function): + """Autograd-aware MoonEP combine and saved-plan backward redispatch.""" + + @staticmethod + def forward(ctx, expert_output, buffer, plan, bridge, fwd_hidden_buffer): + """Combine expert outputs using the matching forward dispatch plan.""" + expert_output = expert_output.contiguous() + use_zero_copy = ( + fwd_hidden_buffer is not None + and expert_output.data_ptr() == fwd_hidden_buffer[1].data_ptr() + ) + combined, _, _ = buffer.combine( + plan=plan, + hidden_nvsh=expert_output, + zero_copy=use_zero_copy, + hidden_buffer=fwd_hidden_buffer, + ) + ctx.buffer = buffer + ctx.plan = plan + ctx.bridge = bridge + ctx.fwd_hidden_buffer = fwd_hidden_buffer + return combined + + @staticmethod + def backward(ctx, grad_output): + """Restore plan weights and redispatch the combined output gradient.""" + # Prefetch slots are shared and may have been overwritten by a later + # layer/microbatch. Restore this plan before the expert dgrad runs. + ctx.bridge.prefetch(ctx.plan) + grad_expert_output, _, _, _ = ctx.buffer.dispatch( + grad_output.contiguous(), + plan=ctx.plan, + zero_copy=ctx.fwd_hidden_buffer is not None, + hidden_buffer=ctx.fwd_hidden_buffer, + ) + return grad_expert_output, None, None, None, None + + +def moonep_dispatch( + hidden_states, + topk_probs, + topk_indices, + tokens_per_expert, + buffer, + bridge, + dispatch_buffer_pool=None, + dgrad_hidden_buffer=None, +): + """Dispatch tokens with MoonEP while preserving activation and router gradients.""" + return MoonEPDispatch.apply( + hidden_states, + topk_probs, + topk_indices, + tokens_per_expert, + *bridge.source_parameters, + buffer, + bridge, + dispatch_buffer_pool, + dgrad_hidden_buffer, + ) + + +def moonep_combine(expert_output, buffer, plan, bridge, fwd_hidden_buffer=None): + """Combine MoonEP expert output and install its saved-plan backward.""" + return MoonEPCombine.apply(expert_output, buffer, plan, bridge, fwd_hidden_buffer) + def get_hidden_bytes(x: torch.Tensor) -> int: """Calculate the number of hidden bytes for a tensor. @@ -293,7 +979,7 @@ def init_hybrid_ep_buffer( fp8_dispatch: bool = False, num_sms_preprocessing_api: Optional[int] = None, ) -> None: - ''' + """ Initialize the HybridEP buffer, including buffer allocation and metadata initialization. @@ -322,20 +1008,20 @@ def init_hybrid_ep_buffer( Whether to use FP8 communication during the dispatch phase. num_sms_preprocessing_api (Optional[int]): Number of SMs used by the preprocessing (metadata scan) kernel. - ''' + """ assert not fp8_dispatch, "HybridEP dispatcher does not support fp8 dispatch now" global _hybrid_ep_buffer kwargs = {} if num_sms_dispatch_api is not None: - kwargs['num_sms_dispatch_api'] = num_sms_dispatch_api + kwargs["num_sms_dispatch_api"] = num_sms_dispatch_api if num_sms_combine_api is not None: - kwargs['num_sms_combine_api'] = num_sms_combine_api + kwargs["num_sms_combine_api"] = num_sms_combine_api if num_blocks_permute is not None: - kwargs['num_blocks_permute'] = num_blocks_permute + kwargs["num_blocks_permute"] = num_blocks_permute if num_blocks_unpermute is not None: - kwargs['num_blocks_unpermute'] = num_blocks_unpermute + kwargs["num_blocks_unpermute"] = num_blocks_unpermute if num_sms_preprocessing_api is not None: - kwargs['num_sms_preprocessing_api'] = num_sms_preprocessing_api + kwargs["num_sms_preprocessing_api"] = num_sms_preprocessing_api _hybrid_ep_buffer = HybridEPBuffer( group=group, hidden_dim=hidden_dim, @@ -347,17 +1033,17 @@ def init_hybrid_ep_buffer( def reset_hybrid_ep_buffer(): - ''' + """ Reset the HybridEP buffer - ''' + """ global _hybrid_ep_buffer _hybrid_ep_buffer = None class HybridEPDispatch(torch.autograd.Function): - ''' + """ Fused dispatch operation for permute + dispatch a2a + permute using the HybridEP backend - ''' + """ @staticmethod def forward( @@ -376,15 +1062,15 @@ def forward( pad_multiple=None, num_sms_preprocessing_api=108, ): - ''' + """ Forward pass of fused dispatch of the HybridEP backend - ''' + """ if fused or num_blocks_permute is not None or num_blocks_unpermute is not None: import inspect import warnings sig = inspect.signature(HybridEPBuffer.dispatch_with_permute) - if 'fuse_permute_dispatch' not in sig.parameters: + if "fuse_permute_dispatch" not in sig.parameters: warnings.warn( "Current DeepEP version does not support fused permute dispatch or " "num_blocks_permute/num_blocks_unpermute. Falling back to unfused " @@ -446,9 +1132,9 @@ def forward( @staticmethod def backward(ctx, grad_x, grad_probs, grad_scaling_factor, grad_tokens_per_expert, grad_handle): - ''' + """ Backward pass of fused dispatch of the HybridEP backend - ''' + """ handle = ctx.handle combined_hidden, combined_probs = _hybrid_ep_buffer.combine_with_unpermute( hidden=grad_x, @@ -476,15 +1162,15 @@ def backward(ctx, grad_x, grad_probs, grad_scaling_factor, grad_tokens_per_exper @internal_api class HybridEPCombine(torch.autograd.Function): - ''' + """ Fused combine operation for permute + combine a2a + permute using the HybridEP backend - ''' + """ @staticmethod def forward(ctx, x, handle, num_permuted_tokens=None, pad_multiple=None, fused=False): - ''' + """ Forward pass of fused combine of the HybridEP backend - ''' + """ combined_hidden, _ = _hybrid_ep_buffer.combine_with_unpermute( hidden=x, handle=handle, @@ -499,9 +1185,9 @@ def forward(ctx, x, handle, num_permuted_tokens=None, pad_multiple=None, fused=F @staticmethod def backward(ctx, grad_x): - ''' + """ Backward pass of fused combine of the HybridEP backend - ''' + """ handle = ctx.handle dispatched_hidden, _, _, _, _ = _hybrid_ep_buffer.dispatch_with_permute( hidden=grad_x, @@ -532,7 +1218,7 @@ def hybrid_ep_dispatch( pad_multiple=None, num_sms_preprocessing_api=108, ): - ''' + """ Perform fused dispatch for "permute + dispatch a2a + permute" using the HybridEP backend. @@ -564,7 +1250,7 @@ def hybrid_ep_dispatch( is performed. num_sms_preprocessing_api (int): Number of SMs used by the preprocessing (metadata scan) kernel. - ''' + """ return HybridEPDispatch.apply( x, routing_map, @@ -583,7 +1269,7 @@ def hybrid_ep_dispatch( @internal_api def hybrid_ep_combine(x, handle, num_permuted_tokens, pad_multiple, fused=False): - ''' + """ Perform fused combine operation for unpermute + combine a2a + unpermute using the HybridEP backend @@ -598,7 +1284,7 @@ def hybrid_ep_combine(x, handle, num_permuted_tokens, pad_multiple, fused=False) pad_multiple (int): The alignment multiple required for FP8 GEMM. If not provided, no padding is performed. - ''' + """ return HybridEPCombine.apply(x, handle, num_permuted_tokens, pad_multiple, fused) else: diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 48e78775a84..a3ef5d17b7d 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -244,12 +244,12 @@ def __init__( ) # If using mcore cudagraphs, recompute is handled by transformer_layer.MoETransformerLayer self.moe_layer_recompute = ( - config.recompute_granularity == 'selective' + config.recompute_granularity == "selective" and "moe" in config.recompute_modules - and config.cuda_graph_impl != 'local' + and config.cuda_graph_impl != "local" ) self.shared_experts_recompute = ( - config.recompute_granularity == 'selective' + config.recompute_granularity == "selective" and "shared_experts" in config.recompute_modules ) @@ -329,6 +329,7 @@ def __init__( pg_collection=pg_collection, name=(name + ".experts") if name is not None else None, ) + self.token_dispatcher.set_experts(self.experts) # Initialize shared experts if self.use_shared_expert: @@ -346,7 +347,7 @@ def __init__( # Inference-optimized mode setup if config.transformer_impl == "inference_optimized": - if config.inference_grouped_gemm_backend == 'auto': + if config.inference_grouped_gemm_backend == "auto": assert HAVE_FLASHINFER, ( "inference_grouped_gemm_backend='auto'" "requires flashinfer-python. " @@ -361,14 +362,14 @@ def __init__( from megatron.core.inference.utils import check_flashinfer_jit_cache_installed check_flashinfer_jit_cache_installed() - elif config.inference_grouped_gemm_backend == 'torch': - assert hasattr(torch.nn.functional, 'grouped_mm') or hasattr( - torch, '_grouped_mm' + elif config.inference_grouped_gemm_backend == "torch": + assert hasattr(torch.nn.functional, "grouped_mm") or hasattr( + torch, "_grouped_mm" ), ( "inference_grouped_gemm_backend='torch' requires " "torch.nn.functional.grouped_mm (> torch 2.10) or torch._grouped_mm (<= 2.10)." ) - elif config.inference_grouped_gemm_backend == 'vllm': + elif config.inference_grouped_gemm_backend == "vllm": assert HAVE_TRITON, ( "inference_grouped_gemm_backend='vllm' requires Triton. " "Install triton (pip install triton)." @@ -393,7 +394,7 @@ def _setup_inference_mode(self, pg_collection): """ dispatcher_type = self.config.inference_moe_token_dispatcher_type dispatcher_cls = ( - NVLSAllGatherVDispatcher if dispatcher_type == 'nvls' else NCCLAllGatherDispatcher + NVLSAllGatherVDispatcher if dispatcher_type == "nvls" else NCCLAllGatherDispatcher ) self._training_token_dispatcher = self.token_dispatcher @@ -408,7 +409,7 @@ def _setup_inference_mode(self, pg_collection): # The dispatcher launches the shared-expert forward on SharedExpertMLP.stream # concurrently with AGV+experts+RSV and adds it back in combine_postprocess. if ( - dispatcher_type == 'nvls' + dispatcher_type == "nvls" and self.use_shared_expert and self.config.moe_shared_expert_overlap ): @@ -546,7 +547,7 @@ def routed_experts_compute(self, hidden_states: torch.Tensor, probs: torch.Tenso dispatched_input, tokens_per_expert, permuted_probs, routing_map=routing_map ) else: - # NCCL-EP zero-copy: experts write fc2 output and fc1 dgrad straight into the combine / + # Backend zero-copy: experts write fc2 output and fc1 dgrad straight into the combine / # dispatch symm buffers. Passed only when set (non-TEGroupedMLP experts don't accept # these kwargs). output_buffer, grad_input_buffer = self.token_dispatcher.get_expert_zero_copy_buffers() diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index 9b7cf177c79..de53adaa843 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -1418,12 +1418,13 @@ 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. + HybridEP / MoonEP / NCCL-EP dispatcher. """ if config.moe_router_padding_for_quantization: return True if config.moe_token_dispatcher_type == "flex" and config.moe_flex_dispatcher_backend in ( "hybridep", + "moonep", "ncclep", ): return True diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 5743a047960..9df01984cc4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -2,6 +2,7 @@ import logging import os +import socket import warnings from abc import ABC, abstractmethod from typing import List, Optional, Tuple @@ -21,14 +22,21 @@ from megatron.core.transformer.enums import CudaGraphModule from megatron.core.transformer.moe.fused_a2a import ( HYBRIDEP_TOKEN_ALIGNMENT, + MoonEPDispatchBufferPool, + MoonEPWeightBridge, alloc_ep_symm_buffer, ensure_nccl_ep_bootstrapped, fused_combine, fused_dispatch, + get_moonep_zero_copy_token_buffers, hybrid_ep_combine, hybrid_ep_dispatch, + is_moonep_available, + moonep_combine, + moonep_dispatch, nccl_ep_combine, nccl_ep_dispatch, + new_moonep_buffer, new_nccl_ep_buffer, set_deepep_num_sms, ) @@ -216,6 +224,10 @@ def set_shared_experts(self, shared_experts): self.shared_experts = shared_experts self.use_nccl_stream = True + def set_experts(self, experts): + """Bind routed experts when a dispatcher backend needs their runtime state.""" + del experts + def get_expert_zero_copy_buffers(self): """Buffers the experts should write their output / grad input into, if any. @@ -263,7 +275,7 @@ def __init__( # Attributes that need to be captured in cudagraph. These attributes are returned # as cudagraph outputs when the cuda_graph_modules contains moe_preprocess. - self.cudagraph_attrs = ['routing_map'] + self.cudagraph_attrs = ["routing_map"] def dispatch_preprocess( self, hidden_states: torch.Tensor, routing_map: torch.Tensor, probs: torch.Tensor @@ -316,7 +328,7 @@ def dispatch_postprocess(self, hidden_states, probs): tokens_per_expert = self.local_map.sum(dim=0).long().cpu() - permuted_local_hidden_states, _, self.reversed_local_input_permutation_mapping, _, _ = ( + (permuted_local_hidden_states, _, self.reversed_local_input_permutation_mapping, _, _) = ( permute( hidden_states, self.local_map, @@ -471,16 +483,16 @@ def __init__( # Attributes that need to be captured in cudagraph. These attributes are returned # as cudagraph outputs when the cuda_graph_modules contains moe_preprocess. self.cudagraph_attrs = [ - 'tokens_per_expert', - 'input_splits', - 'output_splits', - 'output_splits_tp', - 'num_out_tokens', - 'num_global_tokens_per_local_expert', - 'reversed_local_input_permutation_mapping', - 'routing_map', - 'hidden_shape', - 'probs', + "tokens_per_expert", + "input_splits", + "output_splits", + "output_splits_tp", + "num_out_tokens", + "num_global_tokens_per_local_expert", + "reversed_local_input_permutation_mapping", + "routing_map", + "hidden_shape", + "probs", ] self.shared_experts = None @@ -489,8 +501,8 @@ def set_shared_experts(self, shared_experts): """Set shared expert to the dispatcher.""" super().set_shared_experts(shared_experts) if shared_experts.use_shared_expert_gate: - self.cudagraph_attrs.append('shared_experts.gate_score') - self.cudagraph_attrs.append('shared_experts.cached_fc1_input') + self.cudagraph_attrs.append("shared_experts.gate_score") + self.cudagraph_attrs.append("shared_experts.cached_fc1_input") def preprocess(self, routing_map: torch.Tensor) -> torch.Tensor: """ @@ -1051,7 +1063,7 @@ def __init__( ) self.moe_expert_rank_capacity_factor = self.config.moe_expert_rank_capacity_factor - self.over_budget = torch.zeros(1, dtype=torch.bool, device='cuda') + self.over_budget = torch.zeros(1, dtype=torch.bool, device="cuda") # HybridEP dispatch expects equal per-rank input sizes. When requested, # variable token counts are padded to the group-wide max and trimmed in combine. self._original_num_tokens: Optional[int] = None @@ -1207,9 +1219,9 @@ def get_restored_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> return hidden_states def get_number_of_tokens_per_expert(self) -> torch.Tensor: - ''' + """ Get the number of tokens per expert. - ''' + """ return self.tokens_per_expert @@ -1307,7 +1319,7 @@ def dispatch( "DeepEP only supports float32 probs, please set --moe-router-dtype=fp32" ) self.token_probs = self.token_probs.float() # downcast or upcast - hidden_states, dispatched_indices, dispatched_probs, num_tokens_per_expert, handle = ( + (hidden_states, dispatched_indices, dispatched_probs, num_tokens_per_expert, handle) = ( fused_dispatch( hidden_states, self.token_indices, @@ -1459,6 +1471,190 @@ def get_restored_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> return hidden_states +class _MoonEPManager(_DispatchManager): + """MoonEP dispatch manager for fixed-shape, single-node BF16 training.""" + + def __init__( + self, + group: torch.distributed.ProcessGroup, + num_local_experts: int, + router_topk: int, + num_experts: int, + config: TransformerConfig, + ): + self.group = group + self.num_local_experts = int(num_local_experts) + self.router_topk = int(router_topk) + self.num_experts = int(num_experts) + self.config = config + # Latent MoE projects tokens before dispatch and combines them before + # projecting back to hidden_size in MoELayer.postprocess(). + self.hidden_dim = getattr(config, "moe_latent_size", None) or config.hidden_size + if not is_moonep_available(): + raise ImportError( + "MoonEP is not installed. Install the optional 'moonep' package before using " + "moe_flex_dispatcher_backend='moonep'." + ) + + self.token_probs: Optional[torch.Tensor] = None + self.token_indices: Optional[torch.Tensor] = None + self.tokens_per_expert: Optional[torch.Tensor] = None + self.dispatched_probs: Optional[torch.Tensor] = None + self.handle = None + self._buffer = None + self._bridge: Optional[MoonEPWeightBridge] = None + self._num_local_tokens: Optional[int] = None + self._buffer_num_tokens: Optional[int] = None + self._dispatch_capacity: Optional[int] = None + self._dispatch_hidden_buffer_pool = None + self._zero_copy_token_buffers = None + + def bind_experts(self, experts) -> None: + """Bind the registered grouped expert parameters to MoonEP runtime storage.""" + self._bridge = MoonEPWeightBridge( + experts=experts, + group=self.group, + num_experts=self.num_experts, + num_local_experts=self.num_local_experts, + num_sms=self.config.moe_flex_dispatcher_num_sms, + ) + experts.set_moonep_weight_bridge(self._bridge) + + def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): + num_tokens = int(routing_map.shape[0]) + routing_map = routing_map.reshape(num_tokens, self.num_experts) + probs = probs.reshape(num_tokens, self.num_experts) + self.token_probs, self.token_indices = torch.topk(probs, self.router_topk, dim=-1) + self.token_indices = self.token_indices.to(torch.int32) + # torch.bincount performs D2H scalar reads of the input min/max even + # when minlength is fixed. MoonEP knows E statically, so build the + # fixed-size histogram entirely on the GPU instead. + flat_indices = self.token_indices.reshape(-1).long() + self.tokens_per_expert = torch.zeros( + self.num_experts, dtype=torch.int32, device=flat_indices.device + ) + self.tokens_per_expert.scatter_add_( + 0, flat_indices, torch.ones_like(flat_indices, dtype=torch.int32) + ) + self._num_local_tokens = num_tokens + + def _ensure_buffer(self, hidden_states: torch.Tensor) -> None: + if self._bridge is None: + raise RuntimeError("MoonEP experts must be bound before the first dispatch.") + num_tokens = int(hidden_states.shape[0]) + if self._buffer is not None: + if num_tokens != self._buffer_num_tokens: + raise ValueError( + "MoonEP v1 requires a fixed local token count while a layer buffer is live: " + f"initialized with {self._buffer_num_tokens}, got {num_tokens}." + ) + return + rank_metadata = [None] * torch.distributed.get_world_size(group=self.group) + torch.distributed.all_gather_object( + rank_metadata, + (socket.gethostname(), num_tokens, int(hidden_states.shape[1])), + group=self.group, + ) + hostnames = {metadata[0] for metadata in rank_metadata} + if len(hostnames) != 1: + raise ValueError( + "MoonEP v1 requires all communication ranks on one NVLink-connected host, " + f"got hosts={sorted(hostnames)}." + ) + token_counts = {metadata[1] for metadata in rank_metadata} + if len(token_counts) != 1: + raise ValueError( + "MoonEP requires equal local token counts across its communication group, " + f"got counts={sorted(token_counts)}." + ) + hidden_dims = {metadata[2] for metadata in rank_metadata} + if len(hidden_dims) != 1: + raise ValueError( + "MoonEP requires equal dispatched hidden dimensions across its communication " + f"group, got dimensions={sorted(hidden_dims)}." + ) + self._buffer_num_tokens = num_tokens + self._buffer = new_moonep_buffer( + S=num_tokens, + H=int(hidden_states.shape[1]), + K=self.router_topk, + E=self.num_experts, + num_ep_ranks=torch.distributed.get_world_size(group=self.group), + num_sms=self.config.moe_flex_dispatcher_num_sms, + B=self.num_local_experts, + group=self.group, + ) + self._bridge.attach_buffer(self._buffer) + self._dispatch_hidden_buffer_pool = MoonEPDispatchBufferPool(self._buffer) + self._zero_copy_token_buffers = get_moonep_zero_copy_token_buffers(self._buffer) + + def dispatch( + self, + hidden_states: torch.Tensor, + async_finish: bool = False, + allocate_on_comm_stream: bool = False, + ) -> torch.Tensor: + del async_finish, allocate_on_comm_stream + self._ensure_buffer(hidden_states) + dispatched_hidden, self.dispatched_probs, self.tokens_per_expert = moonep_dispatch( + hidden_states, + self.token_probs, + self.token_indices, + self.tokens_per_expert, + self._buffer, + self._bridge, + self._dispatch_hidden_buffer_pool, + self._zero_copy_token_buffers.backward, + ) + self.handle = self._bridge.last_plan + self._dispatch_capacity = int(dispatched_hidden.shape[0]) + return dispatched_hidden + + def combine( + self, + hidden_states: torch.Tensor, + async_finish: bool = False, + allocate_on_comm_stream: bool = False, + ) -> torch.Tensor: + del async_finish, allocate_on_comm_stream + hidden_states = moonep_combine( + hidden_states, + self._buffer, + self.handle, + self._bridge, + self._zero_copy_token_buffers.forward, + ) + self.handle = None + self.dispatched_probs = None + self._dispatch_capacity = None + return hidden_states + + def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: + # MoonEP dispatch returns a fixed-capacity [NvS, H] tensor. As in the + # NCCL-EP static-shape path, TE consumes the whole storage allocation + # and uses the device-side E+B counts to delimit valid expert rows. + # Rows beyond sum(counts) are ignored by grouped GEMM and combine. + return hidden_states, self.dispatched_probs + + def get_restored_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: + pad_rows = self._dispatch_capacity - int(hidden_states.shape[0]) + if pad_rows > 0: + hidden_states = torch.cat( + [hidden_states, hidden_states.new_zeros(pad_rows, hidden_states.shape[-1])], dim=0 + ) + return hidden_states + + def get_number_of_tokens_per_expert(self) -> torch.Tensor: + """Return the padded ``E+B`` expert-group token counts.""" + return self.tokens_per_expert + + def get_expert_zero_copy_buffers(self): + """Return local FC2-output and FC1-dgrad symmetric buffer views.""" + if self._zero_copy_token_buffers is None: + return None, None + return (self._zero_copy_token_buffers.forward[1], self._zero_copy_token_buffers.backward[1]) + + class _NCCLEPManager(_DispatchManager): """A manager class to handle dispatch/combine for MoE models using the NCCL Expert Parallelism backend, via TransformerEngine's transformer_engine.pytorch.ep API @@ -1701,9 +1897,9 @@ def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> return permuted_hidden, permuted_probs def get_number_of_tokens_per_expert(self) -> torch.Tensor: - ''' + """ Get the number of tokens per expert. - ''' + """ return self.tokens_per_expert def get_restored_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -1775,7 +1971,7 @@ def __init__( num_experts=self.tp_size * self.config.num_moe_experts, config=self.config, ) - self.cudagraph_attrs = ['_comm_manager.token_probs', '_comm_manager.token_indices'] + self.cudagraph_attrs = ["_comm_manager.token_probs", "_comm_manager.token_indices"] elif self.config.moe_flex_dispatcher_backend == "hybridep": self._comm_manager = _HybridEPManager( group=self.tp_ep_group, @@ -1783,7 +1979,17 @@ def __init__( num_experts=self.tp_size * self.config.num_moe_experts, config=self.config, ) - self.cudagraph_attrs = ['_comm_manager.token_probs', '_comm_manager.routing_map'] + self.cudagraph_attrs = ["_comm_manager.token_probs", "_comm_manager.routing_map"] + elif self.config.moe_flex_dispatcher_backend == "moonep": + assert self.tp_size == 1, "MoonEP dispatcher requires expert tensor parallel size 1" + self._comm_manager = _MoonEPManager( + group=self.tp_ep_group, + num_local_experts=self.num_local_experts, + router_topk=self.config.moe_router_topk, + num_experts=self.config.num_moe_experts, + config=self.config, + ) + self.cudagraph_attrs = [] elif self.config.moe_flex_dispatcher_backend == "ncclep": assert self.tp_size * self.ep_size > 1, "NCCL EP dispatcher requires TPxEP > 1" self._comm_manager = _NCCLEPManager( @@ -1793,23 +1999,38 @@ def __init__( num_experts=self.tp_size * self.config.num_moe_experts, config=self.config, ) - self.cudagraph_attrs = ['_comm_manager.token_probs', '_comm_manager.token_indices'] + self.cudagraph_attrs = ["_comm_manager.token_probs", "_comm_manager.token_indices"] else: raise ValueError( f"Invalid backend: {self.config.moe_flex_dispatcher_backend}" - "Please set --moe-flex-dispatcher-backend to deepep, hybridep, or ncclep" + "Please set --moe-flex-dispatcher-backend to deepep, hybridep, moonep, or ncclep" ) + def set_experts(self, experts) -> None: + """Bind backend-specific expert runtime state after expert construction.""" + bind_experts = getattr(self._comm_manager, "bind_experts", None) + if bind_experts is not None: + bind_experts(experts) + def get_expert_zero_copy_buffers(self): - """NCCL-EP zero-copy: ``(output_buffer, grad_input_buffer)`` — the shared symm buffers the - experts write the fc2 output / fc1 dgrad into, so combine (fwd) and dispatch (bwd) read and - scatter them one-sided. ``(None, None)`` for every other backend/mode. + """Return backend buffers that experts write FC2 output / FC1 dgrad into directly. + + MoonEP and NCCL-EP then consume these symmetric buffers in combine + (forward) and dispatch (backward), eliminating the large activation + boundary copies. ``(None, None)`` for other backends/modes. Returned detached: the op-fuser calls requires_grad_() on its output and returns it, so handing it the persistent buffer would permanently mark the shared classvar as requiring grad and break the next layer's reuse. The detached view shares storage (zero-copy intact). """ + if self.config.moe_flex_dispatcher_backend == "moonep": + output_buffer, grad_input_buffer = self._comm_manager.get_expert_zero_copy_buffers() + return ( + output_buffer.detach() if output_buffer is not None else None, + grad_input_buffer.detach() if grad_input_buffer is not None else None, + ) + def _detached(name): buf = getattr(self._comm_manager, name, None) return buf.detach() if buf is not None else None @@ -1987,12 +2208,12 @@ def combine_postprocess(self, hidden_states: torch.Tensor): def check_over_budget(self): """Check if the dispatcher has exceeded its budget.""" - if hasattr(self._comm_manager, 'over_budget'): + if hasattr(self._comm_manager, "over_budget"): return self._comm_manager.over_budget else: return None def reset_over_budget(self): """Reset the accumulated over-budget flag on the communication manager.""" - if hasattr(self._comm_manager, 'over_budget'): + if hasattr(self._comm_manager, "over_budget"): self._comm_manager.over_budget.fill_(0) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 4f9de546161..6d430a682a3 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -9,6 +9,7 @@ import torch import torch.nn.functional as F +from megatron.core.activations import squared_relu from megatron.core.enums import Fp4Recipe, Fp8Recipe from megatron.core.inference.moe import InferenceGroupedGemmBackend from megatron.core.quantization.quant_config import RecipeConfig @@ -866,10 +867,11 @@ class TransformerConfig(ModelParallelConfig): moe_enable_deepep: bool = False """[Experimental] Enable DeepEP for efficient token dispatching and combine in MoE models.""" - moe_flex_dispatcher_backend: Literal['deepep', 'hybridep', 'ncclep'] = "deepep" + moe_flex_dispatcher_backend: Literal['deepep', 'hybridep', 'moonep', 'ncclep'] = "deepep" """[Experimental] The backend to use for flex token dispatcher. The default is "deepep". - Options are "deepep", "hybridep", and "ncclep". Currently only "hybridep" backend supports - the MNNVL case. "ncclep" uses NVIDIA NCCL Expert Parallelism via TransformerEngine's + Options are "deepep", "hybridep", "moonep", and "ncclep". Currently only "hybridep" backend + supports the MNNVL case. "moonep" is a single-node NVLink backend using an externally installed + MoonEP package. "ncclep" uses NVIDIA NCCL Expert Parallelism via TransformerEngine's transformer_engine.pytorch.ep API.""" moe_permute_fusion_into_hybridep: bool = False @@ -925,8 +927,8 @@ class TransformerConfig(ModelParallelConfig): moe_flex_dispatcher_num_sms: Optional[int] = None """Number of SMs for the flex token dispatcher's dispatch/combine communication, for all - backends (deepep, hybridep, ncclep). None lets each backend use its own default. Unifies the - deprecated per-backend moe_{deepep,hybridep}_num_sms knobs (routed in __post_init__).""" + backends (deepep, hybridep, moonep, ncclep). None lets each backend use its own default. Unifies + the deprecated per-backend moe_{deepep,hybridep}_num_sms knobs (routed in __post_init__).""" moe_deepep_num_sms: Optional[int] = None """DEPRECATED: use moe_flex_dispatcher_num_sms. Number of SMs to use for DeepEP (historical @@ -1613,6 +1615,80 @@ def __post_init__(self): "moe_token_dispatcher_type='flex'." ) + if self.moe_flex_dispatcher_backend == "moonep": + moonep_errors = [] + if self.moe_token_dispatcher_type != "flex": + moonep_errors.append("moe_token_dispatcher_type='flex'") + if not self.bf16 or self.params_dtype != torch.bfloat16: + moonep_errors.append("BF16 execution and BF16 parameters") + if self.fp8 or self.fp4: + moonep_errors.append("FP8 and FP4 disabled") + if self.add_bias_linear: + moonep_errors.append("add_bias_linear=False") + if not self.moe_grouped_gemm: + moonep_errors.append("moe_grouped_gemm=True") + if not self.moe_single_grouped_weight: + moonep_errors.append("moe_single_grouped_weight=True") + if self.moe_single_grouped_bias: + moonep_errors.append("moe_single_grouped_bias=False") + if not self.use_transformer_engine_op_fuser: + moonep_errors.append("use_transformer_engine_op_fuser=True") + if not self.gradient_accumulation_fusion: + moonep_errors.append("gradient_accumulation_fusion=True") + if self.moe_router_dtype != "fp32": + moonep_errors.append("moe_router_dtype='fp32'") + if self.expert_tensor_parallel_size != 1: + moonep_errors.append("expert_tensor_parallel_size=1") + if self.moe_router_topk > 32: + moonep_errors.append("moe_router_topk<=32") + supported_glu = self.gated_linear_unit and self.activation_func in (F.silu, quick_gelu) + supported_squared_relu = ( + not self.gated_linear_unit + and self.activation_func == squared_relu + and self.use_fused_weighted_squared_relu + ) + if not (supported_glu or supported_squared_relu): + moonep_errors.append( + "fused SwiGLU, quick-GeGLU, or weighted squared-ReLU activation" + ) + if self.cuda_graph_impl != "none": + moonep_errors.append("cuda_graph_impl='none'") + if self.delay_wgrad_compute: + moonep_errors.append("delay_wgrad_compute=False") + if self.overlap_dispatch_backward_with_experts_wgrad: + moonep_errors.append("overlap_dispatch_backward_with_experts_wgrad=False") + if self.overlap_moe_expert_parallel_comm: + moonep_errors.append("overlap_moe_expert_parallel_comm=False") + if self.moe_shared_expert_overlap: + moonep_errors.append("moe_shared_expert_overlap=False") + expert_input_size = self.moe_latent_size or self.hidden_size + if expert_input_size % 128 != 0: + moonep_errors.append( + "moe_latent_size (or hidden_size when latent MoE is disabled) divisible by 128" + ) + if self.moe_ffn_hidden_size % 128 != 0: + moonep_errors.append("moe_ffn_hidden_size divisible by 128") + if self.moe_expert_capacity_factor is not None: + moonep_errors.append("moe_expert_capacity_factor=None") + if self.moe_expert_rank_capacity_factor is not None: + moonep_errors.append("moe_expert_rank_capacity_factor=None") + if self.moe_pad_expert_input_to_capacity: + moonep_errors.append("moe_pad_expert_input_to_capacity=False") + if self.moe_router_padding_for_quantization: + moonep_errors.append("moe_router_padding_for_quantization=False") + if self.moe_token_dropping: + moonep_errors.append("moe_token_dropping=False") + if self.moe_apply_probs_on_input: + moonep_errors.append("moe_apply_probs_on_input=False") + if self.enable_cuda_graph or self.external_cuda_graph: + moonep_errors.append("deprecated CUDA graph flags disabled") + if moonep_errors: + raise ValueError( + "MoonEP flex dispatcher configuration is unsupported; require " + + ", ".join(moonep_errors) + + "." + ) + # 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 # moe_flex_dispatcher_num_sms takes precedence) and warn. diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index b9e7fa346d2..022c0f0896c 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -452,11 +452,21 @@ def __init__(self, glu_interleave_size, *, activation_recompute_in_mlp=False): self.activation_recompute_in_mlp = activation_recompute_in_mlp class FakeScaledClampedQGeGLU(torch.nn.Module): - def __init__(self, glu_interleave_size, *, activation_recompute_in_mlp=False, limit=None): + def __init__( + self, + glu_interleave_size, + *, + activation_recompute_in_mlp=False, + limit=7.0, + alpha=1.702, + glu_linear_offset=1.0, + ): super().__init__() self.glu_interleave_size = glu_interleave_size self.activation_recompute_in_mlp = activation_recompute_in_mlp self.limit = limit + self.alpha = alpha + self.glu_linear_offset = glu_linear_offset class FakeScaledSReLU(torch.nn.Module): def __init__(self, *, activation_recompute_in_mlp=False): @@ -484,8 +494,9 @@ def register_forward_pre_hook(self, hook): ) -def test_make_fused_ops_uses_clamped_qgeglu_for_quick_gelu(monkeypatch): - """quick_gelu + clamp value → ScaledClampedQGeGLU(limit=clamp).""" +@pytest.mark.parametrize(("clamp_value", "expected_limit"), [(7.0, 7.0), (None, float("inf"))]) +def test_make_fused_ops_uses_qgeglu_for_quick_gelu(monkeypatch, clamp_value, expected_limit): + """Quick-GeGLU carries Megatron's clamp and linear-offset semantics into TE.""" fake_te, FakeGroupedLinear = _make_fake_te_namespace() monkeypatch.setattr(experts_module, "te", fake_te) @@ -494,9 +505,10 @@ def test_make_fused_ops_uses_clamped_qgeglu_for_quick_gelu(monkeypatch): module.config = SimpleNamespace( moe_mlp_glu_interleave_size=4, delay_wgrad_compute=False, - activation_func_clamp_value=7.0, + activation_func_clamp_value=clamp_value, activation_func=quick_gelu, gated_linear_unit=True, + glu_linear_offset=0.25, ) module.activation_func = quick_gelu module.activation_recompute = True @@ -519,7 +531,9 @@ def test_make_fused_ops_uses_clamped_qgeglu_for_quick_gelu(monkeypatch): assert type(activation).__name__ == "FakeScaledClampedQGeGLU" assert activation.glu_interleave_size == 4 assert activation.activation_recompute_in_mlp is True - assert activation.limit == 7.0 + assert activation.limit == expected_limit + assert activation.alpha == 1.702 + assert activation.glu_linear_offset == 0.25 def test_make_fused_ops_uses_scaled_srelu_for_weighted_squared_relu(monkeypatch): diff --git a/tests/unit_tests/transformer/moe/test_moonep.py b/tests/unit_tests/transformer/moe/test_moonep.py new file mode 100644 index 00000000000..8dc0eff53fd --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_moonep.py @@ -0,0 +1,463 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import os +import weakref +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +from megatron.core.activations import squared_relu +from megatron.core.fusions.fused_bias_geglu import quick_gelu +from megatron.core.transformer.moe import fused_a2a +from megatron.core.transformer.moe.fused_a2a import moonep_combine, moonep_dispatch + + +class _FakeMoonEPBuffer: + """CPU implementation of MoonEP's saved-plan dispatch/combine contract.""" + + def dispatch( + self, + hidden, + route_weights=None, + topk_experts=None, + tokens_per_expert=None, + plan=None, + *, + zero_copy=False, + zero_copy_weights=None, + hidden_buffer=None, + ): + del zero_copy, zero_copy_weights, hidden_buffer + if plan is None: + num_tokens, topk = topk_experts.shape + flat_experts = topk_experts.reshape(-1).long() + order = torch.argsort(flat_experts, stable=True) + source_tokens = torch.arange(num_tokens).repeat_interleave(topk)[order] + source_routes = torch.arange(topk).repeat(num_tokens)[order] + plan = SimpleNamespace( + source_tokens=source_tokens, + source_routes=source_routes, + num_tokens=num_tokens, + topk=topk, + experts_to_copy=torch.full((1, tokens_per_expert.numel()), -1, dtype=torch.int32), + ) + counts = torch.cat( + [tokens_per_expert, tokens_per_expert.new_zeros(tokens_per_expert.numel())] + ) + cu_seqlens = counts.cumsum(0) + else: + cu_seqlens = None + + dispatched_hidden = hidden[plan.source_tokens] + dispatched_weights = ( + None if route_weights is None else route_weights[plan.source_tokens, plan.source_routes] + ) + return dispatched_hidden, dispatched_weights, cu_seqlens, plan + + def combine( + self, *, plan, hidden_nvsh, route_weights_nvs=None, zero_copy=False, hidden_buffer=None + ): + del zero_copy, hidden_buffer + hidden = hidden_nvsh.new_zeros((plan.num_tokens, hidden_nvsh.shape[-1])) + hidden.index_add_(0, plan.source_tokens, hidden_nvsh) + weights = None + if route_weights_nvs is not None: + weights = route_weights_nvs.new_zeros((plan.num_tokens, plan.topk)) + weights[plan.source_tokens, plan.source_routes] = route_weights_nvs + return hidden, weights, None + + +class _FakeWeightBridge: + def __init__(self, device=None): + self.parameters = ( + torch.nn.Parameter(torch.ones((), device=device)), + torch.nn.Parameter(torch.ones((), device=device)), + ) + self.last_plan = None + self.reduced_plans = [] + self.prefetched_plans = [] + self.buffer = None + + @property + def source_parameters(self): + return self.parameters + + @property + def dummy_grads(self): + return tuple(torch.zeros_like(parameter) for parameter in self.parameters) + + def reduce_grads(self, plan): + self.reduced_plans.append(plan) + + def prefetch(self, plan): + self.prefetched_plans.append(plan) + + def attach_buffer(self, buffer): + self.buffer = buffer + + +class _FakeDispatchBufferPool: + def __init__(self): + self.acquired = [] + self.released = [] + + def acquire(self): + pair = (object(), object()) + self.acquired.append(pair) + return pair + + def release(self, pair): + self.released.append(pair) + + +def _run_fake_moonep(hidden, probs, indices, bridge, buffer, dispatch_buffer_pool=None): + tokens_per_expert = torch.bincount(indices.reshape(-1), minlength=int(indices.max()) + 1).to( + torch.int32 + ) + dispatched, dispatched_probs, runtime_counts = moonep_dispatch( + hidden, probs, indices, tokens_per_expert, buffer, bridge, dispatch_buffer_pool + ) + expert_output = dispatched * dispatched_probs.unsqueeze(-1) + output = moonep_combine(expert_output, buffer, bridge.last_plan, bridge) + return output, runtime_counts + + +def test_moonep_autograd_wrappers_preserve_hidden_and_probability_gradients(): + buffer = _FakeMoonEPBuffer() + bridge = _FakeWeightBridge() + indices = torch.tensor([[0, 2], [1, 2], [0, 1]], dtype=torch.int32) + hidden = torch.randn(3, 4, requires_grad=True) + probs = torch.randn(3, 2, requires_grad=True) + ref_hidden = hidden.detach().clone().requires_grad_(True) + ref_probs = probs.detach().clone().requires_grad_(True) + + output, runtime_counts = _run_fake_moonep(hidden, probs, indices, bridge, buffer) + ref_output = (ref_hidden.unsqueeze(1) * ref_probs.unsqueeze(2)).sum(dim=1) + grad = torch.randn_like(output) + output.backward(grad) + ref_output.backward(grad) + + torch.testing.assert_close(output, ref_output) + torch.testing.assert_close(hidden.grad, ref_hidden.grad) + torch.testing.assert_close(probs.grad, ref_probs.grad) + assert runtime_counts.numel() == 6 # E+B, with B=E for the one-rank fake. + assert runtime_counts.sum() == indices.numel() + assert list(map(id, bridge.prefetched_plans)) == list(map(id, bridge.reduced_plans)) + assert len(bridge.reduced_plans) == 1 + assert all(parameter.grad.item() == 0 for parameter in bridge.parameters) + + +def test_moonep_saved_plans_are_restored_for_multiple_outstanding_forwards(): + buffer = _FakeMoonEPBuffer() + bridge = _FakeWeightBridge() + dispatch_buffer_pool = _FakeDispatchBufferPool() + indices = torch.tensor([[0, 1], [1, 2]], dtype=torch.int32) + hidden_1 = torch.randn(2, 4, requires_grad=True) + hidden_2 = torch.randn(2, 4, requires_grad=True) + probs_1 = torch.randn(2, 2, requires_grad=True) + probs_2 = torch.randn(2, 2, requires_grad=True) + + output_1, _ = _run_fake_moonep(hidden_1, probs_1, indices, bridge, buffer, dispatch_buffer_pool) + plan_1 = bridge.last_plan + output_2, _ = _run_fake_moonep( + hidden_2, probs_2, indices.flip(0), bridge, buffer, dispatch_buffer_pool + ) + plan_2 = bridge.last_plan + assert len(dispatch_buffer_pool.acquired) == 2 + assert dispatch_buffer_pool.released == [] + (output_1.sum() + output_2.sum()).backward() + + assert set(map(id, bridge.prefetched_plans)) == {id(plan_1), id(plan_2)} + assert set(map(id, bridge.reduced_plans)) == {id(plan_1), id(plan_2)} + assert set(map(id, dispatch_buffer_pool.released)) == set( + map(id, dispatch_buffer_pool.acquired) + ) + + +def test_moonep_manager_preserves_static_dispatch_capacity(): + from megatron.core.transformer.moe.token_dispatcher import _MoonEPManager + + manager = _MoonEPManager.__new__(_MoonEPManager) + manager.dispatched_probs = torch.randn(12) + hidden = torch.randn(12, 4) + + expert_hidden, expert_probs = manager.get_permuted_hidden_states_by_experts(hidden) + + assert expert_hidden.data_ptr() == hidden.data_ptr() + assert expert_hidden.shape == (12, 4) + assert expert_probs.data_ptr() == manager.dispatched_probs.data_ptr() + + +def test_moonep_manager_exposes_shared_expert_zero_copy_buffers(): + from megatron.core.transformer.moe.token_dispatcher import _MoonEPManager + + manager = _MoonEPManager.__new__(_MoonEPManager) + output_buffer = torch.empty(12, 4) + dgrad_buffer = torch.empty(12, 4) + manager._zero_copy_token_buffers = SimpleNamespace( + forward=(object(), output_buffer), backward=(object(), dgrad_buffer) + ) + + actual_output, actual_dgrad = manager.get_expert_zero_copy_buffers() + + assert actual_output.data_ptr() == output_buffer.data_ptr() + assert actual_dgrad.data_ptr() == dgrad_buffer.data_ptr() + + +def test_moonep_metadata_uses_fixed_gpu_histogram(monkeypatch): + from megatron.core.transformer.moe.token_dispatcher import _MoonEPManager + + manager = _MoonEPManager.__new__(_MoonEPManager) + manager.num_experts = 4 + manager.router_topk = 2 + probs = torch.tensor([[4.0, 3.0, 2.0, 1.0], [1.0, 4.0, 3.0, 2.0], [4.0, 1.0, 3.0, 2.0]]) + routing_map = torch.zeros_like(probs, dtype=torch.bool) + + def unexpected_bincount(*_args, **_kwargs): + raise AssertionError("MoonEP metadata must not call torch.bincount") + + monkeypatch.setattr(torch, "bincount", unexpected_bincount) + manager.setup_metadata(routing_map, probs) + + torch.testing.assert_close( + manager.tokens_per_expert, torch.tensor([2, 2, 2, 0], dtype=torch.int32) + ) + + +def test_moonep_finalize_is_idempotent(monkeypatch): + class _Resource: + def __init__(self): + self.destroy_calls = 0 + + def destroy(self): + self.destroy_calls += 1 + + buffer = _Resource() + bridge = _Resource() + dispatch_pool = _Resource() + token_pool = _Resource() + token_buffers = {"test": token_pool} + monkeypatch.setattr(fused_a2a, "_moonep_buffers", weakref.WeakSet([buffer])) + monkeypatch.setattr(fused_a2a, "_moonep_bridges", weakref.WeakSet([bridge])) + monkeypatch.setattr( + fused_a2a, "_moonep_dispatch_buffer_pools", weakref.WeakSet([dispatch_pool]) + ) + monkeypatch.setattr(fused_a2a, "_moonep_token_buffer_pools", token_buffers) + + fused_a2a.moonep_finalize() + fused_a2a.moonep_finalize() + + assert buffer.destroy_calls == 1 + assert bridge.destroy_calls == 1 + assert dispatch_pool.destroy_calls == 1 + assert token_pool.destroy_calls == 1 + assert token_buffers == {} + + +@pytest.mark.skipif(not fused_a2a.HAVE_MOONEP, reason="MoonEP is not installed") +def test_moonep_availability_helper(): + assert fused_a2a.is_moonep_available() + + +@pytest.mark.internal +@pytest.mark.skipif( + not torch.cuda.is_available() or not fused_a2a.HAVE_MOONEP, + reason="CUDA and MoonEP are required", +) +def test_moonep_four_rank_dispatch_probability_grad_and_redundant_counts(): + """Run with ``torch.distributed.run --nproc_per_node=4`` on an NVLink node.""" + if int(os.environ.get("WORLD_SIZE", "1")) != 4: + pytest.skip("MoonEP distributed coverage requires a 4-rank torchrun launch") + + from megatron.core import parallel_state + from megatron.core.transformer.moe.token_dispatcher import _MoonEPManager + from tests.unit_tests.test_utilities import Utils + + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, expert_model_parallel_size=4, expert_tensor_parallel_size=1 + ) + group = parallel_state.get_expert_tensor_and_model_parallel_group() + config = SimpleNamespace(hidden_size=128, moe_flex_dispatcher_num_sms=None) + manager = _MoonEPManager( + group=group, num_local_experts=1, router_topk=2, num_experts=4, config=config + ) + manager._bridge = _FakeWeightBridge(device="cuda") + + try: + num_tokens = 16 + hidden = torch.randn(num_tokens, 128, device="cuda", dtype=torch.bfloat16) + hidden.requires_grad_(True) + # Every rank strongly favors experts 0 and 1. Non-owner ranks must use + # MoonEP's redundant B slot for at least one of those experts. + logits = torch.full((num_tokens, 4), -8.0, device="cuda") + logits[:, 0] = 8.0 + logits[:, 1] = 7.0 + logits.requires_grad_(True) + dense_probs = torch.softmax(logits, dim=-1) + _, indices = torch.topk(dense_probs, 2, dim=-1) + routing_map = torch.zeros_like(dense_probs, dtype=torch.bool) + routing_map.scatter_(1, indices, True) + + manager.setup_metadata(routing_map, dense_probs) + dispatched = manager.dispatch(hidden) + runtime_counts = manager.get_number_of_tokens_per_expert() + valid_hidden, valid_probs = manager.get_permuted_hidden_states_by_experts(dispatched) + expert_output = (valid_hidden * valid_probs.unsqueeze(-1)).to(hidden.dtype) + expert_output = manager.get_restored_hidden_states_by_experts(expert_output) + output = manager.combine(expert_output) + + expected = hidden * manager.token_probs.sum(dim=-1, keepdim=True).to(hidden.dtype) + torch.testing.assert_close(output, expected) + output.float().sum().backward() + assert hidden.grad is not None + assert logits.grad is not None and torch.count_nonzero(logits.grad) > 0 + assert runtime_counts.numel() == 5 # E+B with E=4 and B=1. + slot_tokens = runtime_counts[4:].sum().to(torch.int64) + torch.distributed.all_reduce(slot_tokens, group=group) + assert slot_tokens.item() > 0 + finally: + fused_a2a.moonep_finalize() + Utils.destroy_model_parallel() + + +def _set_main_grad(parameter): + rowwise_data = getattr(parameter, "rowwise_data", parameter) + parameter.main_grad = torch.zeros_like(rowwise_data).view(parameter.shape) + parameter.grad_added_to_main_grad = False + parameter.overwrite_main_grad = True + + +def _set_main_grads(layer): + for linear in (layer.experts.linear_fc1, layer.experts.linear_fc2): + _set_main_grad(linear.get_parameter("weight")) + if layer.config.moe_latent_size is not None: + _set_main_grad(layer.fc1_latent_proj.weight) + _set_main_grad(layer.fc2_latent_proj.weight) + + +@pytest.mark.internal +@pytest.mark.skipif( + not torch.cuda.is_available() or not fused_a2a.HAVE_MOONEP, + reason="CUDA and MoonEP are required", +) +@pytest.mark.parametrize( + ( + "activation_func", + "gated_linear_unit", + "weighted_squared_relu", + "glu_interleave", + "moe_latent_size", + ), + [ + (F.silu, True, False, None, None), + (F.silu, True, False, 128, None), + (quick_gelu, True, False, None, None), + (squared_relu, False, True, None, None), + (F.silu, True, False, None, 512), + ], +) +def test_moonep_full_layer_parity_with_alltoall( + monkeypatch, + activation_func, + gated_linear_unit, + weighted_squared_relu, + glu_interleave, + moe_latent_size, +): + """Compare full expert/router fwd+bwd and grouped main_grads on 4 NVLink GPUs.""" + if int(os.environ.get("WORLD_SIZE", "1")) != 4: + pytest.skip("MoonEP distributed coverage requires a 4-rank torchrun launch") + + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + from megatron.core.transformer.moe.moe_layer import MoELayer + from megatron.core.transformer.spec_utils import get_submodules + from megatron.core.transformer.transformer_config import TransformerConfig + from tests.unit_tests.test_utilities import Utils + + monkeypatch.setenv("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "1") + monkeypatch.setenv("NVTE_DISABLE_CUTEDSL_WGRAD_FUSED_GROUPED_MLP", "1") + monkeypatch.setenv("NVTE_GROUPED_LINEAR_SINGLE_PARAM", "1") + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, expert_model_parallel_size=4, expert_tensor_parallel_size=1 + ) + + common = { + "num_layers": 1, + "hidden_size": 1024, + "ffn_hidden_size": 1024, + "moe_ffn_hidden_size": 1024, + "num_attention_heads": 8, + "num_moe_experts": 4, + "expert_model_parallel_size": 4, + "expert_tensor_parallel_size": 1, + "moe_router_topk": 2, + "moe_router_load_balancing_type": "none", + "moe_router_dtype": "fp32", + "moe_grouped_gemm": True, + "moe_single_grouped_weight": True, + "use_transformer_engine_op_fuser": True, + "gradient_accumulation_fusion": True, + "add_bias_linear": False, + "bf16": True, + "params_dtype": torch.bfloat16, + "use_cpu_initialization": False, + "activation_func": activation_func, + "gated_linear_unit": gated_linear_unit, + "use_fused_weighted_squared_relu": weighted_squared_relu, + "moe_mlp_glu_interleave_size": glu_interleave, + "moe_latent_size": moe_latent_size, + } + alltoall_config = TransformerConfig(**common, moe_token_dispatcher_type="alltoall") + moonep_config = TransformerConfig( + **common, moe_token_dispatcher_type="flex", moe_flex_dispatcher_backend="moonep" + ) + mlp_spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=4, moe_grouped_gemm=True + ).submodules.mlp + submodules = get_submodules(mlp_spec) + + try: + ref_layer = MoELayer(alltoall_config, submodules).cuda() + moonep_layer = MoELayer(moonep_config, submodules).cuda() + moonep_layer.load_state_dict(ref_layer.state_dict()) + assert moonep_layer.state_dict().keys() == ref_layer.state_dict().keys() + _set_main_grads(ref_layer) + _set_main_grads(moonep_layer) + + torch.manual_seed(1234) + test_input = torch.randn(2, 4, 1024, device="cuda", dtype=torch.bfloat16) + + def run(layer): + hidden = test_input.detach().clone().requires_grad_(True) + output, _ = layer(hidden) + output.float().sum().backward() + values = [ + output.detach(), + hidden.grad.detach(), + layer.router.weight.grad.detach().clone(), + layer.experts.linear_fc1.weight.main_grad.detach().clone(), + layer.experts.linear_fc2.weight.main_grad.detach().clone(), + ] + if layer.config.moe_latent_size is not None: + values.extend( + [ + layer.fc1_latent_proj.weight.main_grad.detach().clone(), + layer.fc2_latent_proj.weight.main_grad.detach().clone(), + ] + ) + return values + + ref_values = run(ref_layer) + moonep_values = run(moonep_layer) + value_names = ["output", "input grad", "router grad", "FC1 main_grad", "FC2 main_grad"] + if moe_latent_size is not None: + value_names.extend(["latent FC1 main_grad", "latent FC2 main_grad"]) + for value_name, actual, expected in zip(value_names, moonep_values, ref_values): + torch.testing.assert_close( + actual, expected, rtol=2e-2, atol=2e-2, msg=lambda msg: f"{value_name}: {msg}" + ) + finally: + fused_a2a.moonep_finalize() + Utils.destroy_model_parallel() diff --git a/tests/unit_tests/transformer/test_transformer_config.py b/tests/unit_tests/transformer/test_transformer_config.py index febb3842789..682526831f4 100644 --- a/tests/unit_tests/transformer/test_transformer_config.py +++ b/tests/unit_tests/transformer/test_transformer_config.py @@ -1,7 +1,12 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import pytest +import torch +import torch.nn.functional as F +from megatron.core.activations import squared_relu +from megatron.core.fusions.fused_bias_geglu import quick_gelu +from megatron.core.transformer.moe.token_dispatcher import _MoonEPManager from megatron.core.transformer.transformer_config import TransformerConfig @@ -30,3 +35,129 @@ def test_ep_a2a_overlap_accepts_supported_mtp_layer_counts(mtp_num_layers: int | def test_ep_a2a_overlap_rejects_unsupported_mtp_layer_counts(mtp_num_layers: int): with pytest.raises(AssertionError, match="MTP supports at most one layer"): _make_overlap_config(mtp_num_layers) + + +def _make_moonep_config(**overrides) -> TransformerConfig: + kwargs = { + "num_layers": 1, + "hidden_size": 256, + "ffn_hidden_size": 512, + "num_attention_heads": 4, + "num_moe_experts": 8, + "expert_model_parallel_size": 8, + "expert_tensor_parallel_size": 1, + "moe_router_topk": 2, + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "moonep", + "moe_router_dtype": "fp32", + "moe_grouped_gemm": True, + "moe_single_grouped_weight": True, + "use_transformer_engine_op_fuser": True, + "gradient_accumulation_fusion": True, + "add_bias_linear": False, + "bf16": True, + "params_dtype": torch.bfloat16, + "gated_linear_unit": True, + "activation_func": F.silu, + } + kwargs.update(overrides) + return TransformerConfig(**kwargs) + + +@pytest.fixture +def moonep_config(monkeypatch): + """Avoid making config validation depend on the TE version in the unit-test environment.""" + monkeypatch.setattr( + "megatron.core.transformer.transformer_config.is_te_min_version", lambda _: True + ) + return _make_moonep_config + + +@pytest.mark.parametrize( + "activation_overrides", + [ + {"activation_func": F.silu, "gated_linear_unit": True}, + {"activation_func": quick_gelu, "gated_linear_unit": True}, + { + "activation_func": squared_relu, + "gated_linear_unit": False, + "use_fused_weighted_squared_relu": True, + }, + ], +) +def test_moonep_accepts_supported_activations(moonep_config, activation_overrides): + config = moonep_config(**activation_overrides) + + assert config.moe_flex_dispatcher_backend == "moonep" + + +def test_moonep_accepts_latent_moe(moonep_config): + config = moonep_config(moe_latent_size=128) + + assert config.moe_latent_size == 128 + + +def test_moonep_rejects_unaligned_latent_size(moonep_config): + with pytest.raises(ValueError, match="moe_latent_size.*divisible by 128"): + moonep_config(moe_latent_size=64) + + +@pytest.mark.parametrize( + ("override", "requirement"), + [ + ({"moe_token_dispatcher_type": "alltoall"}, "moe_token_dispatcher_type='flex'"), + ({"bf16": False, "params_dtype": torch.float32}, "BF16 execution"), + ({"add_bias_linear": True}, "add_bias_linear=False"), + ({"moe_grouped_gemm": False}, "moe_grouped_gemm=True"), + ({"moe_single_grouped_weight": False}, "moe_single_grouped_weight=True"), + ({"use_transformer_engine_op_fuser": False}, "use_transformer_engine_op_fuser=True"), + ({"gradient_accumulation_fusion": False}, "gradient_accumulation_fusion=True"), + ({"moe_router_dtype": None}, "moe_router_dtype='fp32'"), + ({"expert_tensor_parallel_size": 2}, "expert_tensor_parallel_size=1"), + ({"moe_router_topk": 33}, "moe_router_topk<=32"), + ], +) +def test_moonep_rejects_missing_required_flags(moonep_config, override, requirement): + with pytest.raises(ValueError, match=requirement): + moonep_config(**override) + + +@pytest.mark.parametrize( + "override", + [ + {"fp8": "e4m3", "fp8_recipe": "mxfp8"}, + {"cuda_graph_impl": "local"}, + {"delay_wgrad_compute": True}, + {"overlap_dispatch_backward_with_experts_wgrad": True}, + {"overlap_moe_expert_parallel_comm": True}, + {"moe_shared_expert_overlap": True}, + {"moe_expert_capacity_factor": 1.0}, + {"moe_pad_expert_input_to_capacity": True, "moe_expert_capacity_factor": 1.0}, + {"moe_router_padding_for_quantization": True}, + {"moe_apply_probs_on_input": True}, + ], +) +def test_moonep_rejects_unsupported_features(moonep_config, override): + with pytest.raises(ValueError, match="MoonEP flex dispatcher configuration is unsupported"): + moonep_config(**override) + + +def test_moonep_rejects_unsupported_activation(moonep_config): + with pytest.raises(ValueError, match="weighted squared-ReLU activation"): + moonep_config(activation_func=F.gelu, gated_linear_unit=True) + + +def test_moonep_manager_reports_missing_optional_package(moonep_config, monkeypatch): + config = moonep_config() + monkeypatch.setattr( + "megatron.core.transformer.moe.token_dispatcher.is_moonep_available", lambda: False + ) + + with pytest.raises(ImportError, match="MoonEP is not installed"): + _MoonEPManager( + group=None, + num_local_experts=1, + router_topk=config.moe_router_topk, + num_experts=config.num_moe_experts, + config=config, + )