From 4fd602fdecefecff94b0f0df84453451ff0a6d0a Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 29 Sep 2026 15:12:52 -0700 Subject: [PATCH 01/11] Add grouped HybridStack syntax and EP-overlap scheduling Refresh PR #4942 against main while preserving hash routing, MTP, CP layouts, and ungrouped checkpoint compatibility. Integrate the cross-slot MTP gradient fix and reject unsupported direct-overlap execution modes. Original grouped HybridStack work by @Wohox, @Connor-XY, and @guihong-nv in #4798/#4942. Signed-off-by: Yan Xu --- .../models/common/fine_grained_callables.py | 19 +- .../models/hybrid/fine_grained_callables.py | 416 +++++++++++++ megatron/core/models/hybrid/hybrid_block.py | 296 +++++++-- .../models/hybrid/hybrid_layer_allocation.py | 392 ++++++++++-- megatron/core/models/hybrid/hybrid_model.py | 565 +++++++++++++----- megatron/core/models/hybrid/layers/utils.py | 3 + .../hybrid/model_chunk_schedule_plan.py | 156 +++++ megatron/core/recompute.py | 25 +- megatron/core/ssm/mamba_layer.py | 11 + megatron/core/ssm/mamba_mixer.py | 15 + pretrain_hybrid.py | 38 +- .../a2a_overlap/test_hybrid_schedule_plan.py | 181 ++++++ .../test_schedule_quantization_context.py | 46 ++ .../test_hybrid_fine_grained_callables.py | 295 +++++++++ .../models/test_hybrid_hash_routing.py | 32 +- tests/unit_tests/models/test_hybrid_model.py | 222 ++++++- tests/unit_tests/ssm/test_hybrid_block.py | 284 ++++++++- .../ssm/test_hybrid_layer_allocation.py | 84 +++ .../transformer/moe/test_aux_loss.py | 40 ++ .../test_multi_token_prediction.py | 4 + .../transformer/test_submodule_callables.py | 46 +- 21 files changed, 2897 insertions(+), 273 deletions(-) create mode 100644 megatron/core/models/hybrid/fine_grained_callables.py create mode 100644 megatron/core/models/hybrid/model_chunk_schedule_plan.py create mode 100644 tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py create mode 100644 tests/unit_tests/models/test_hybrid_fine_grained_callables.py diff --git a/megatron/core/models/common/fine_grained_callables.py b/megatron/core/models/common/fine_grained_callables.py index 6fb03cba2fd..aa83dcdd6f0 100644 --- a/megatron/core/models/common/fine_grained_callables.py +++ b/megatron/core/models/common/fine_grained_callables.py @@ -67,8 +67,12 @@ def submodule_mtp_pre_dispatch_forward(node, hidden_states): offset = get_mtp_layer_offset( layer.config, model.vp_stage, pp_rank=model.pg_collection.pp.rank() ) - node.chunk_state.mtp_hidden_states = list(torch.chunk(hidden_states, 1 + offset, dim=0)) - hidden_states = node.chunk_state.mtp_hidden_states[offset] + chunks = list(torch.chunk(hidden_states, 1 + offset, dim=0)) + # These chunks cross from pre-dispatch to MTP post-process. Detach + # their stored views so the slots do not traverse final_norm's graph + # twice; node.backward_impl merges their gradients into this slot. + node.chunk_state.mtp_hidden_states = [node.detach(chunk) for chunk in chunks] + hidden_states = chunks[offset] input_ids, position_ids, padding_mask, mtp_input_mask, decoder_input, hidden_states = ( layer._get_embeddings( @@ -151,9 +155,14 @@ def rng_context_wrapper(func, *args, **kwargs): def get_layer_moe_metadata(layer): """Return ``(is_moe, num_local_experts)`` for schedule-node construction.""" + from megatron.core.models.hybrid.hybrid_block import HybridStack if isinstance(layer, MultiTokenPredictionLayer): return get_layer_moe_metadata(layer.mtp_model_layer) + if isinstance(layer, HybridStack): + from megatron.core.models.hybrid.fine_grained_callables import get_hybrid_stack_moe_metadata + + return get_hybrid_stack_moe_metadata(layer) if isinstance(layer, TransformerLayer): is_moe = isinstance(layer.mlp, MoELayer) num_local_experts = layer.mlp.num_local_experts if is_moe else None @@ -167,9 +176,15 @@ def build_layer_callables(layer): Returns ``(forward_funcs, backward_dw)``. """ + from megatron.core.models.hybrid.hybrid_block import HybridStack if isinstance(layer, MultiTokenPredictionLayer): return build_mtp_layer_callables(layer) + if isinstance(layer, HybridStack): + from megatron.core.models.hybrid.fine_grained_callables import build_hybrid_stack_callables + + forward_funcs, backward_dw, _, _ = build_hybrid_stack_callables(layer) + return forward_funcs, backward_dw if isinstance(layer, TransformerLayer): return build_transformer_layer_callables(layer) diff --git a/megatron/core/models/hybrid/fine_grained_callables.py b/megatron/core/models/hybrid/fine_grained_callables.py new file mode 100644 index 00000000000..1371e3b0b9a --- /dev/null +++ b/megatron/core/models/hybrid/fine_grained_callables.py @@ -0,0 +1,416 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +from contextlib import nullcontext +from functools import partial +from typing import Optional + +import torch +from torch import Tensor + +from megatron.core.enums import Fp8Recipe +from megatron.core.fp4_utils import get_fp4_context +from megatron.core.fp8_utils import get_fp8_context +from megatron.core.models.common.utils import TransformerLayerNode, should_free_input +from megatron.core.models.hybrid.hybrid_block import HybridStack +from megatron.core.models.hybrid.hybrid_layer_allocation import LayerPatternItem +from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols +from megatron.core.models.hybrid.hybrid_layer_allocation import is_layer_group +from megatron.core.pipeline_parallel.utils import ( + ScheduleNode, + StageDispatchBwdGrad, + get_comm_stream, +) +from megatron.core.transformer.transformer_layer import make_viewless_tensor + + +class _MoEBackwardDWWrapper: + """Run delayed MoE wgrad and expose only the parameters owned by its slot. + + Shared experts and the latent down projection run in pre-dispatch; routed + experts and the latent up projection run later. Their gradient-accumulation + hooks must run after their respective wgrad, even if a previous microbatch + has already populated ``param.grad``. + """ + + def __init__(self, mlp, routed_experts: bool): + self.dw_callable = partial( + mlp.backward_dw, routed_experts=routed_experts, shared_experts=not routed_experts + ) + self.wgrad_on_comm_stream = False + self.submodules = [] + if routed_experts: + self.submodules.append(mlp.experts) + elif mlp.use_shared_expert and not mlp.shared_expert_overlap: + self.submodules.append(mlp.shared_experts) + if mlp.config.moe_latent_size and mlp.config.overlap_moe_expert_parallel_comm: + self.submodules.append(mlp.fc2_latent_proj if routed_experts else mlp.fc1_latent_proj) + self.wgrad_on_comm_stream = routed_experts + + def backward_dw(self): + """Run this slot's wgrad after its autograd backward.""" + self.dw_callable() + if self.wgrad_on_comm_stream: + # MoELayer runs the latent up-projection wgrad on the communication + # stream. Its accumulation hooks run on this node's compute stream. + torch.cuda.current_stream().wait_stream(get_comm_stream()) + self.dw_callable = None + + def parameters(self): + """Keep parameter owners alive until the node collects post-wgrad hooks. + + ``TransformerLayerNode.backward_dw`` enumerates these parameters after + calling ``backward_dw``, then releases the wrapper with the slot state. + """ + for module in self.submodules: + yield from module.parameters() + + +class HybridStackNode(TransformerLayerNode): + """Schedule node for HybridStack-built fine-grained callables. + + Subclassed from ``TransformerLayerNode`` so the runtime backbone (forward / + backward / backward_dw plumbing, detach bookkeeping, output-grad release) + is shared. The hybrid path keeps a separate node class so its free-input + policy can diverge from the GPT defaults — for example, the + ``pre_dispatch_computation`` slot here covers the whole pre-dispatch loop + (mamba + attention + …) rather than a single attention block, and + group-level decisions about whether the input is needed in backward may + differ from ``should_free_input`` in ``gpt/fine_grained_callables.py``. + Keep this override thin until a hybrid counter-example forces it to + diverge; the explicit subclass exists so the divergence can be made + surgically without touching the GPT class. + """ + + @staticmethod + def _resolve_free_input(name, is_moe, config, num_local_experts): + """Hybrid free-input policy. + + Currently mirrors the GPT default: dense layers always retain their + input for backward; MoE-only "moe_dispatch", "mlp", and "moe_combine" + slots can free, subject to the dispatcher / cuda-graph constraints + encoded in ``should_free_input``. Hybrid groups have a + "pre_dispatch_computation" slot whose semantics differ (it covers a + loop over Mamba/attention/GDN sub-layers, not a single attention + block), but its policy resolves to ``False`` in + ``should_free_input``, which is correct: pre-layer outputs are needed + for backward through the loop. Override here when a hybrid-specific + rule is needed. + """ + return should_free_input(name, is_moe, config, num_local_experts) + + +def _get_inner_quant_context(layer): + config = layer.config + if config.fp8 and config.fp8_recipe != Fp8Recipe.delayed: + return get_fp8_context(config, layer.layer_number - 1) + if config.fp4: + return get_fp4_context(config, layer.layer_number - 1) + return nullcontext() + + +def _as_hybrid_layers(layer, layer_type: Optional[LayerPatternItem]): + """Return ``(layer_type, layer)`` pairs for a hybrid logical layer.""" + if isinstance(layer, HybridStack): + return list(zip(layer.layer_type_list, layer.layers)) + assert layer_type is not None, "Hybrid layer scheduling requires the layer type symbol." + return [(layer_type, layer)] + + +def _split_hybrid_layers_for_overlap(layer, layer_type: Optional[LayerPatternItem]): + layer_items = _as_hybrid_layers(layer, layer_type) + if any(is_layer_group(item_type) for item_type, _ in layer_items): + raise ValueError("Nested HybridStack groups are not supported in overlap scheduling.") + + terminal_idx = None + for idx, (item_type, _) in enumerate(layer_items): + if item_type in (LayerSymbols.MLP, LayerSymbols.MOE): + terminal_idx = idx + break + + if terminal_idx is not None and terminal_idx != len(layer_items) - 1: + raise ValueError("HybridStack overlap requires MLP/MoE to be the last layer in a group.") + + terminal_type = layer_items[terminal_idx][0] if terminal_idx is not None else None + terminal_layer = layer_items[terminal_idx][1] if terminal_idx is not None else None + pre_layers = layer_items[:terminal_idx] if terminal_idx is not None else layer_items + is_moe = terminal_type == LayerSymbols.MOE + num_local_experts = terminal_layer.mlp.num_local_experts if is_moe else None + return pre_layers, terminal_type, terminal_layer, is_moe, num_local_experts + + +def get_hybrid_stack_moe_metadata(layer, layer_type: Optional[LayerPatternItem] = None): + """Return ``(is_moe, num_local_experts)`` for one HybridStack schedule layer.""" + _, _, _, is_moe, num_local_experts = _split_hybrid_layers_for_overlap(layer, layer_type) + return is_moe, num_local_experts + + +def _maybe_apply_final_norm(node: ScheduleNode, hidden_states: Tensor): + final_norm = getattr(node.chunk_state.model.decoder, "final_norm", None) + final_norm = final_norm or getattr(node.chunk_state.model.decoder, "final_layernorm", None) + if not node.is_mtp and final_norm is not None and node.is_last_layer: + hidden_states = final_norm(hidden_states) + hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + return hidden_states + + +def _get_moe_padding_mask(node: ScheduleNode): + padding_mask = node.chunk_state.padding_mask + if padding_mask is not None: + # MoELayer.forward receives [batch, seq] and transposes before routing. + padding_mask = padding_mask.transpose(0, 1).bool() + return padding_mask + + +def _run_moe_preprocess(layer, node: ScheduleNode, hidden_states: Tensor): + pre_mlp_layernorm_output = layer._forward_pre_mlp_layernorm(hidden_states) + # A subsequent microbatch can reuse this layer before this node combines. + node.layer_state.mlp_norm_manager = layer.mlp_norm_manager + layer.mlp_norm_manager = None + if isinstance(pre_mlp_layernorm_output, tuple): + if len(pre_mlp_layernorm_output) != 2: + raise ValueError( + f"When the output of pre_mlp_layernorm is a tuple, it is expected to have " + f"2 elements (output, residual), but got {len(pre_mlp_layernorm_output)}" + ) + pre_mlp_layernorm_output, residual = pre_mlp_layernorm_output + else: + residual = hidden_states + + if layer.config.fp32_residual_connection: + residual = residual.float() + + shared_expert_output = layer.mlp.shared_experts_compute(pre_mlp_layernorm_output) + probs, routing_map = layer.mlp.route(pre_mlp_layernorm_output, _get_moe_padding_mask(node)) + local_tokens, probs = layer.mlp.preprocess(pre_mlp_layernorm_output, probs, routing_map) + + node.layer_state.residual = node.detach(residual) + if layer.mlp.use_shared_expert and not layer.mlp.shared_expert_overlap: + node.layer_state.shared_expert_output = node.detach(shared_expert_output) + + return local_tokens, probs + + +def _run_moe_experts(layer, node: ScheduleNode, dispatched_tokens: Tensor): + dispatched_probs = node.layer_state.dispatched_probs + enable_hybridep = ( + layer.config.moe_token_dispatcher_type == "flex" + and layer.config.moe_flex_dispatcher_backend == "hybridep" + ) + enable_deepep = ( + layer.config.moe_token_dispatcher_type == "flex" + and layer.config.moe_flex_dispatcher_backend == "deepep" + ) + enable_ncclep = ( + layer.config.moe_token_dispatcher_type == "flex" + and layer.config.moe_flex_dispatcher_backend == "ncclep" + ) + token_dispatcher = layer.mlp.token_dispatcher + if enable_deepep or enable_hybridep or enable_ncclep: + token_dispatcher._comm_manager.dispatched_probs = dispatched_probs + + expert_output, _ = layer.mlp.routed_experts_compute(dispatched_tokens, dispatched_probs) + + if enable_hybridep 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 + + if layer.recompute_pre_mlp_layernorm: + layer.pre_mlp_norm_checkpoint.discard_output_and_register_recompute(expert_output) + + return expert_output + + +def _run_moe_combine(layer, node: ScheduleNode, output: Tensor): + residual = node.layer_state.residual + shared_expert_output = getattr(node.layer_state, 'shared_expert_output', None) + output = layer.mlp.combine(output) + output = layer.mlp.postprocess(output, shared_expert_output) + # Inline bda instead of calling ``layer._forward_post_mlp`` so we can skip + # the redundant ``discard_output_and_register_recompute(mlp_output_with_bias[0])`` + # that ``_forward_post_mlp`` would otherwise issue. The pre_mlp_layernorm recompute + # is already registered on ``expert_output`` inside ``_run_moe_experts``; the second + # hook on the combine-slot ``mlp_output_with_bias[0]`` is not only unnecessary but + # harmful in the bracketed-hybrid case (``[*E]``): it fires during combine_bwd's + # autograd backward and triggers the LN recompute ahead of attention's backward + # in the same pre_dispatch slot, corrupting attention gradients (grad_norm explodes + # from iter 2). GPT's ``submodule_combine_forward`` likewise inlines bda and does + # not call ``_forward_post_mlp`` for the same reason. + mlp_output_with_bias = (output, None) + with layer.bias_dropout_add_exec_handler(): + output = layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( + mlp_output_with_bias, residual, layer.hidden_dropout + ) + mlp_norm_manager = getattr(node.layer_state, "mlp_norm_manager", None) + if mlp_norm_manager is not None: + output = mlp_norm_manager.group_offload(output, forced_released_tensors=[residual]) + node.layer_state.mlp_norm_manager = None + output = make_viewless_tensor(inp=output, requires_grad=output.requires_grad, keep_graph=True) + + node.layer_state.residual.record_stream(torch.cuda.current_stream()) + if shared_expert_output is not None: + shared_expert_output.record_stream(torch.cuda.current_stream()) + + node.layer_state.residual = None + node.layer_state.shared_expert_output = None + + return _maybe_apply_final_norm(node, output) + + +def build_hybrid_stack_callables(layer, layer_type: Optional[LayerPatternItem] = None): + """Create fine-grained callables for one logical HybridStack layer. + + A logical layer may be a bracketed nested ``HybridStack`` (for example ``[M*E]``) + or a single legacy hybrid layer symbol. The split is: + pre-dispatch compute -> dispatch -> MLP/experts -> combine. + """ + pre_layers, terminal_type, terminal_layer, is_moe, num_local_experts = ( + _split_hybrid_layers_for_overlap(layer, layer_type) + ) + + def pre_dispatch_computation(node: ScheduleNode, hidden_states: Tensor): + for item_type, item_layer in pre_layers: + with _get_inner_quant_context(item_layer): + if item_type == LayerSymbols.MAMBA: + hidden_states = item_layer( + hidden_states=hidden_states, + attention_mask=node.chunk_state.attention_mask, + inference_context=getattr(node.chunk_state, "inference_context", None), + packed_seq_params=node.chunk_state.packed_seq_params, + ) + elif item_type in ( + LayerSymbols.ATTENTION, + LayerSymbols.DS_ATTENTION, + LayerSymbols.MLA, + LayerSymbols.GDN, + ): + # Use _forward_attention rather than __call__: an attention half-layer has + # mlp=IdentityOp / mlp_bda=IdentityFuncOp by default, and TransformerLayer's + # __call__ would route through _forward_mlp + mlp_bda, double-applying the + # post-attention residual. + hidden_states, _ = item_layer._forward_attention( + hidden_states=hidden_states, + attention_mask=node.chunk_state.attention_mask, + rotary_pos_emb=node.chunk_state.rotary_pos_emb, + rotary_pos_cos=node.chunk_state.rotary_pos_cos, + rotary_pos_sin=node.chunk_state.rotary_pos_sin, + packed_seq_params=node.chunk_state.packed_seq_params, + sequence_len_offset=node.chunk_state.sequence_len_offset, + ) + # _forward_attention returns the bias_dropout_add output which can be a + # view tensor (the mlp_bda's add into the post-attention residual produces + # a view from a fused/JIT kernel). Downstream cuBLAS matmuls — including + # the terminal MLP/MoE's pre_mlp_layernorm and the next attention's QKV + # projection in a multi-pre-layer group — pick algorithms based on input + # strides; a view's non-canonical strides can lead to different algo + # selection across processes and produce ~1e-5 bit drift on the forward + # output. TransformerLayer's full forward() inserts this exact call at the + # MLP exit (transformer_layer.py:895) for the same reason; the + # _forward_attention shortcut here doesn't get that cleanup, so we add it + # explicitly. Same idea as the make_viewless_tensor in _maybe_apply_final_norm. + hidden_states = make_viewless_tensor( + inp=hidden_states, + requires_grad=hidden_states.requires_grad, + keep_graph=True, + ) + else: + raise ValueError( + f"HybridStack overlap does not support layer type '{item_type}' before " + "the terminal MLP/MoE layer." + ) + + if isinstance(hidden_states, tuple): + hidden_states = hidden_states[0] + + if terminal_type == LayerSymbols.MOE: + with _get_inner_quant_context(terminal_layer): + return _run_moe_preprocess(terminal_layer, node, hidden_states) + + if terminal_type is None: + return _maybe_apply_final_norm(node, hidden_states) + + return hidden_states + + def dispatch(node: ScheduleNode, local_tokens: Tensor, probs: Tensor): + enable_hybridep = ( + terminal_layer.config.moe_token_dispatcher_type == "flex" + and terminal_layer.config.moe_flex_dispatcher_backend == "hybridep" + ) + enable_deepep = ( + terminal_layer.config.moe_token_dispatcher_type == "flex" + and terminal_layer.config.moe_flex_dispatcher_backend == "deepep" + ) + enable_ncclep = ( + terminal_layer.config.moe_token_dispatcher_type == "flex" + and terminal_layer.config.moe_flex_dispatcher_backend == "ncclep" + ) + token_dispatcher = terminal_layer.mlp.token_dispatcher + if enable_deepep or enable_hybridep or enable_ncclep: + token_dispatcher._comm_manager.token_probs = probs + with _get_inner_quant_context(terminal_layer): + dispatched_tokens, dispatched_probs = terminal_layer.mlp.dispatch(local_tokens, probs) + if enable_ncclep and terminal_layer.config.moe_ncclep_zero_copy: + dispatched_tokens = StageDispatchBwdGrad.apply(dispatched_tokens, token_dispatcher) + node.layer_state.dispatched_probs = node.detach(dispatched_probs) + return dispatched_tokens + + def mlp(node: ScheduleNode, hidden_states: Tensor): + if terminal_type == LayerSymbols.MLP: + with _get_inner_quant_context(terminal_layer): + hidden_states = terminal_layer._forward_mlp( + hidden_states, padding_mask=node.chunk_state.padding_mask + ) + return _maybe_apply_final_norm(node, hidden_states) + if terminal_type == LayerSymbols.MOE: + with _get_inner_quant_context(terminal_layer): + return _run_moe_experts(terminal_layer, node, hidden_states) + return hidden_states + + def combine(node: ScheduleNode, output: Tensor): + with _get_inner_quant_context(terminal_layer): + return _run_moe_combine(terminal_layer, node, output) + + def raise_not_implemented(*args): + raise NotImplementedError("This callable is not implemented for non-MoE hybrid layers.") + + backward_dw = {} + pre_bwd_dw = [] + for item_type, item_layer in pre_layers: + if item_type in ( + LayerSymbols.ATTENTION, + LayerSymbols.DS_ATTENTION, + LayerSymbols.MLA, + LayerSymbols.GDN, + ): + # TransformerLayer-backed pre-layers go through the standard + # _BackwardDWWrapper which coordinates attn / shared-expert wgrad + # with cuda-graph replay scopes. + item_layer.init_backward_dw_wrapper() + pre_bwd_dw.append(item_layer.backward_dw_wrapper) + elif item_type == LayerSymbols.MAMBA: + # MambaLayer is not a TransformerLayer, so init_backward_dw_wrapper + # would assert. MambaLayer.backward_dw delegates to its mixer, which + # in turn calls backward_dw on the in_proj / out_proj linears. The + # schedule node iterates this list and calls .backward_dw() on each; + # registering the layer directly is sufficient. + pre_bwd_dw.append(item_layer) + if is_moe: + # Each slot owns the hooks for exactly the parameters whose delayed wgrad + # it computes. Shared-expert dgrad has not run yet in the routed MLP slot. + pre_dispatch_dw = _MoEBackwardDWWrapper(terminal_layer.mlp, routed_experts=False) + if pre_dispatch_dw.submodules: + pre_bwd_dw.append(pre_dispatch_dw) + backward_dw["mlp"] = _MoEBackwardDWWrapper(terminal_layer.mlp, routed_experts=True) + elif terminal_type == LayerSymbols.MLP: + backward_dw["mlp"] = terminal_layer.mlp + + if pre_bwd_dw: + backward_dw["pre_dispatch_computation"] = pre_bwd_dw + + forward_funcs = [ + pre_dispatch_computation, + dispatch if is_moe else raise_not_implemented, + mlp, + combine if is_moe else raise_not_implemented, + None, + ] + return forward_funcs, backward_dw, is_moe, num_local_experts diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index 8c1a4067886..3cfebe63bfd 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -23,7 +23,14 @@ from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode from megatron.core.models.hybrid.hybrid_layer_allocation import ( + LayerConfigItem, + LayerPatternItem, + flatten_layer_type_list, get_layer_type_list_from_layer_config_list, + get_layer_type_physical_count, + is_layer_group, + layer_type_list_to_str, + validate_layer_group, validate_segment_layers, ) from megatron.core.models.hybrid.layers import utils as layer_utils @@ -87,6 +94,17 @@ class HybridStackSubmodules: mtp_stack_submodules: Optional["HybridStackSubmodules"] = None +def _is_layer_type_entry(layer_type) -> bool: + """Return whether a ``layer_type_list`` entry is a layer symbol or a bracketed group.""" + if isinstance(layer_type, str): + return len(layer_type) == 1 + return ( + isinstance(layer_type, tuple) + and len(layer_type) > 0 + and all(isinstance(symbol, str) and len(symbol) == 1 for symbol in layer_type) + ) + + class HybridStack(MegatronModule): """ Constructor for the HybridStack class. @@ -96,15 +114,23 @@ class HybridStack(MegatronModule): submodules (HybridStackSubmodules): the submodules for the stack pre_process (bool, optional): whether to include an embedding layer. Defaults to True. - layer_type_list (list[str], optional): This argument exists for backwards-compatibility - reasons, allowing callers to construct ``HybridStack`` directly with layer symbols. + layer_type_list (list[LayerPatternItem], optional): This argument exists for + backwards-compatibility reasons, allowing callers to construct ``HybridStack`` + directly with layer symbols (bracketed groups as tuples of symbols). It is immediately converted to independent per-layer configs. - layer_config_list (Sequence[TransformerConfig], optional): per-layer configs for this - pipeline segment. When provided by HybridModel, pipeline stage selection has already - been done via '|' separators in the pattern. Exactly one of ``layer_type_list`` or - ``layer_config_list`` must be provided. - pp_layer_offset (int, optional): the global layer offset for this pipeline + layer_config_list (Sequence[LayerConfigItem], optional): per-layer configs for this + pipeline segment. A tuple of configs denotes a bracketed group (e.g. ``[M*E]``) + that is built as one nested ``HybridStack`` logical layer. When provided by + HybridModel, pipeline stage selection has already been done via '|' separators + in the pattern. Exactly one of ``layer_type_list`` or ``layer_config_list`` must + be provided. + pp_layer_offset (int, optional): the global physical layer offset for this pipeline segment. Defaults to 0. + logical_layer_offset (int, optional): the global logical layer offset for this + pipeline segment; bracketed groups count as one logical layer. Used for + checkpoint keys. Defaults to ``pp_layer_offset`` for legacy direct callers. + is_layer_group_stack (bool, optional): whether this stack is the nested stack built + for a bracketed group. Defaults to False. post_layer_norm (bool, optional): whether to include a final layer norm. Defaults to True. post_process (bool, optional): whether to include an output layer. @@ -118,6 +144,9 @@ class HybridStack(MegatronModule): mtp_layer_number (int, optional): enclosing MTP depth for nested MoE metrics. hash_moe_layer_threshold (int, optional): global Hybrid layer-number threshold used to select hash-routed MoE layers. + layer_number_offset (int, optional): global physical layer offset for this + stack's first layer. Defaults to ``pp_layer_offset``. Nested groups use + their own numbering offset while retaining the pipeline's cache offset. """ def __init__( @@ -125,7 +154,7 @@ def __init__( config: TransformerConfig, submodules: HybridStackSubmodules, pre_process: bool = True, - layer_type_list: list[str] | None = None, + layer_type_list: list[LayerPatternItem] | None = None, pp_layer_offset: int = 0, post_layer_norm: bool = True, post_process: bool = True, @@ -136,22 +165,32 @@ def __init__( mtp_layer_number: Optional[int] = None, hash_moe_layer_threshold: Optional[int] = None, name: str | None = None, - layer_config_list: Sequence[TransformerConfig] | None = None, + layer_config_list: Sequence[LayerConfigItem] | None = None, boundary_layout: CPLayout | None = None, + layer_number_offset: int | None = None, + logical_layer_offset: int | None = None, + is_layer_group_stack: bool = False, + transformer_sharded_keys: bool = False, ) -> None: """ Args: + transformer_sharded_keys (bool): emit ``TransformerBlock``-style sharded + checkpoint keys (``final_layernorm`` instead of ``final_norm``) so the + checkpoint is interchangeable with a ``GPTModel`` one. Only set for + bracketed-group patterns, whose logical layers map one-to-one onto + transformer layers; leaving it off keeps the historical hybrid keys so + existing non-grouped hybrid checkpoints stay loadable. name (str | None): module instance name passed top-down from its paranet module """ if (layer_type_list is None) == (layer_config_list is None): raise ValueError("Exactly one of layer_type_list or layer_config_list must be provided") if layer_type_list is not None: - if any( - not isinstance(layer_symbol, str) or len(layer_symbol) != 1 - for layer_symbol in layer_type_list - ): - raise ValueError("Each entry in layer_type_list must be a single layer symbol") - segment = ''.join(layer_type_list) + if any(not _is_layer_type_entry(layer_type) for layer_type in layer_type_list): + raise ValueError( + "Each entry in layer_type_list must be a single layer symbol or a tuple of " + "single layer symbols (a bracketed group)" + ) + segment = layer_type_list_to_str(layer_type_list) warnings.warn( "DEPRECATED(layer_type_list): please use `layer_config_list` instead", DeprecationWarning, @@ -160,6 +199,11 @@ def __init__( layer_config_list = validate_segment_layers(segment, config) for layer_config in layer_config_list: + if is_layer_group(layer_config): + validate_layer_group( + [layer_utils.get_layer_symbol_from_config(config) for config in layer_config] + ) + for layer_config in flatten_layer_type_list(layer_config_list): layer_utils.validate_tp_comm_overlap( layer_config, layer_utils.get_layer_symbol_from_config(layer_config), @@ -178,6 +222,12 @@ def __init__( self.config.linear_cp_layout if boundary_layout is None else boundary_layout ) self.mtp_layer_number = mtp_layer_number + logical_layer_offset = ( + pp_layer_offset if logical_layer_offset is None else logical_layer_offset + ) + self.logical_layer_offset = logical_layer_offset + self.is_layer_group_stack = is_layer_group_stack + self.transformer_sharded_keys = transformer_sharded_keys assert pg_collection is not None, "pg_collection must be provided for HybridStack" @@ -193,19 +243,25 @@ def __init__( self._mhc_block_end_plan: Optional[List[bool]] = None self.layer_config_list = layer_config_list + has_layer_groups = any(is_layer_group(layer_config) for layer_config in layer_config_list) + if has_layer_groups and ( + self.config.enable_mhc_connections + or self.config.wide_residual is not None + or self.config.moe_shortcut_connection + ): + raise NotImplementedError( + "Bracketed HybridStack layer groups are not supported with enable_mhc_connections, " + "wide residuals, or MoE shortcuts." + ) self._has_linear_layer_with_chunkwise_cp = self.cp_group.size() > 1 and any( type(layer_config) is layer_utils.MambaLayerConfig and layer_config.linear_cp_mode == "chunkwise" - for layer_config in self.layer_config_list + for layer_config in flatten_layer_type_list(self.layer_config_list) ) self._cp_layout_manager = None if self.cp_group.size() > 1: layer_layouts = tuple( - ( - layer_config.attention_cp_layout - if type(layer_config) in layer_utils.Symbols.ATTENTION_LAYER_CONFIGS - else layer_config.linear_cp_layout - ) + self._get_layer_cp_layout(layer_config, boundary_layout) for layer_config in self.layer_config_list ) self._cp_layout_manager = ContextParallelLayoutManager( @@ -218,20 +274,55 @@ def __init__( ) # Build layers from the pre-selected segment self.layers = nn.ModuleList() + # ``i`` is the logical layer index within this stack (module-list index and the + # ``name=...layers.{i}`` suffix). ``physical_layer_offset`` is the physical layer + # counter used for ``layer_number`` and the FP8/FP4 contexts; it advances by more + # than one for bracketed groups, which hold several physical layers. + physical_layer_offset = ( + pp_layer_offset if layer_number_offset is None else layer_number_offset + ) for i, layer_config in enumerate(self.layer_config_list): - layer_number = i + 1 + pp_layer_offset - if layer_config.fp8: + layer_number = physical_layer_offset + 1 + if is_layer_group(layer_config): + # The nested stack applies its own per-layer quantization contexts. + quant_init_context = nullcontext() + elif layer_config.fp8: quant_init_context = get_fp8_context( - layer_config, i + pp_layer_offset, is_init=True + layer_config, physical_layer_offset, is_init=True ) elif layer_config.fp4: quant_init_context = get_fp4_context( - layer_config, i + pp_layer_offset, is_init=True + layer_config, physical_layer_offset, is_init=True ) else: quant_init_context = nullcontext() with quant_init_context: - if type(layer_config) is layer_utils.MambaLayerConfig: + if is_layer_group(layer_config): + # A bracketed group (e.g. ``[M*E]``) is one logical layer built as a + # nested HybridStack over its physical layers. It restores the outer + # stack's boundary CP layout on exit. + layer = HybridStack( + config=self.config, + submodules=submodules, + pre_process=True, + layer_config_list=list(layer_config), + pp_layer_offset=pp_layer_offset, + layer_number_offset=physical_layer_offset, + logical_layer_offset=logical_layer_offset + len(self.layers), + is_layer_group_stack=True, + transformer_sharded_keys=transformer_sharded_keys, + post_layer_norm=False, + post_process=False, + device=device, + dtype=dtype, + pg_collection=pg_collection, + is_mtp_layer=is_mtp_layer, + mtp_layer_number=mtp_layer_number, + hash_moe_layer_threshold=hash_moe_layer_threshold, + name=(name + f".layers.{i}") if name is not None else None, + boundary_layout=boundary_layout, + ) + elif type(layer_config) is layer_utils.MambaLayerConfig: layer = build_module( submodules.mamba_layer, config=layer_config, @@ -365,6 +456,7 @@ def __init__( if self.config.enable_mhc_connections: layer = HyperConnectionHybridLayer(config=layer_config, layer=layer) self.layers.append(layer) + physical_layer_offset += get_layer_type_physical_count(layer_config) if self.config.cuda_graph_impl == "local": annotate_first_last_layer(self.layers) @@ -459,6 +551,28 @@ def _uses_hash_routing(layer: torch.nn.Module) -> bool: router = getattr(getattr(inner_layer, "mlp", None), "router", None) return bool(getattr(router, "is_hash_layer", False)) + @property + def final_layernorm(self): + """Alias for ``final_norm`` matching the attribute name on TransformerBlock. + + Lets generic decoder consumers (e.g. ``GPTModel.PostProcessNode``) discover the + final norm via the same attribute name they use for non-hybrid decoders, while + keeping ``final_norm`` as the registered submodule for local state-dict compatibility. + """ + return getattr(self, "final_norm", None) + + @staticmethod + def _get_layer_cp_layout(layer_config: LayerConfigItem, boundary_layout: CPLayout) -> CPLayout: + """Return the CP layout the outer stack must provide for one logical layer.""" + if is_layer_group(layer_config): + # A bracketed group runs as a nested HybridStack that manages the layout + # transitions of its own layers and restores ``boundary_layout`` on exit, so + # the outer stack hands the group its input in the boundary layout. + return boundary_layout + if type(layer_config) in layer_utils.Symbols.ATTENTION_LAYER_CONFIGS: + return layer_config.attention_cp_layout + return layer_config.linear_cp_layout + def set_input_tensor(self, input_tensor: Tensor): """Set input tensor to be used instead of forward()'s input. @@ -490,6 +604,11 @@ def mamba_state_shapes_per_request(self) -> Optional[Tuple[Tuple[int], Tuple[int if this block contains Mamba or GDN layers (this may not be the case with PP > 1). """ for layer_config, layer in zip(self.layer_config_list, self.physical_layers(), strict=True): + if is_layer_group(layer_config): + state_shapes = layer.mamba_state_shapes_per_request() + if state_shapes is not None: + return state_shapes + continue if type(layer_config) is layer_utils.MambaLayerConfig: return layer.mamba_state_shapes_per_request() if type(layer_config) is layer_utils.GDNLayerConfig: @@ -549,6 +668,10 @@ def forward( attention_mask: Tensor, inference_context: Optional[BaseInferenceContext] = None, rotary_pos_emb: Optional[Tensor] = None, + rotary_pos_cos: Optional[Tensor] = None, + rotary_pos_sin: Optional[Tensor] = None, + rotary_pos_cos_sin: Optional[Tensor] = None, + sequence_len_offset: Optional[Tensor] = None, *, inference_params: Optional[BaseInferenceContext] = None, packed_seq_params: Optional[PackedSeqParams] = None, @@ -556,6 +679,7 @@ def forward( packed_seq_params_by_layout: dict[CPLayout, PackedSeqParams | None] | None = None, cp_layout_plan: THDCPLayoutPlan | None = None, input_ids: Optional[Tensor] = None, + _checkpointed_forward_in_parent: bool = False, ): """ Forward function of the HybridStack class. @@ -573,6 +697,14 @@ def forward( Defaults to None. input_ids (Tensor, optional): Token IDs forwarded to hash-routed TransformerLayer instances. Defaults to None. + rotary_pos_cos / rotary_pos_sin / rotary_pos_cos_sin (Tensor, optional): + precomputed rotary embeddings forwarded to transformer layers (flash-decode / + fused-rope inference paths). Defaults to None. + sequence_len_offset (Tensor, optional): precomputed per-sample sequence offsets + for static-batching inference. Computed here when None. + _checkpointed_forward_in_parent (bool): set by ``checkpointed_forward`` when the + enclosing stack already checkpoints this nested group stack, so the group + must not checkpoint its layers a second time. Returns: Tensor: the output tensor. """ @@ -626,7 +758,9 @@ def forward( inference_context.max_seqlen = inference_context.max_sequence_length inference_context.seqlen_offset = inference_context.sequence_len_offset - if ( + if sequence_len_offset is not None: + pass + elif ( (self.config.cuda_graph_impl == "local" or self.config.flash_decode) and inference_context and inference_context.is_static_batching() @@ -688,7 +822,11 @@ def get_inner_quant_context(config, layer_number): ) with outer_fp8_context: - if self.config.recompute_granularity == 'full' and self.training: + if ( + self.config.recompute_granularity == 'full' + and self.training + and not _checkpointed_forward_in_parent + ): hidden_states = checkpointed_forward( self, hidden_states=hidden_states, @@ -703,6 +841,8 @@ def get_inner_quant_context(config, layer_number): use_inner_quantization_context=(use_inner_fp8_context or use_fp4_context), cp_layout_state=cp_layout_state, packed_sequence_cp_metadata=packed_sequence_cp_metadata, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, ) else: for layer_idx, (physical_layer_idx, layer_config, layer) in enumerate( @@ -756,11 +896,29 @@ def get_inner_quant_context(config, layer_number): # Keep both residuals in the layer's layout, inside the CP conversions. residual_accumulator = hidden_states # Layers have 1-indexed layer numbers attribute. - inner_quant_context = get_inner_quant_context( - layer_config, layer.layer_number - 1 + inner_quant_context = ( + nullcontext() + if is_layer_group(layer_config) + else get_inner_quant_context(layer_config, layer.layer_number - 1) ) with inner_quant_context: - if isinstance(layer, (TransformerLayer, HyperConnectionHybridLayer)): + if isinstance(layer, HybridStack): + hidden_states = layer( + hidden_states=hidden_states, + attention_mask=attention_mask, + inference_context=inference_context, + rotary_pos_emb=rotary_pos_emb, + rotary_pos_cos=rotary_pos_cos, + rotary_pos_sin=rotary_pos_sin, + rotary_pos_cos_sin=rotary_pos_cos_sin, + sequence_len_offset=sequence_len_offset, + packed_seq_params=layer_packed_seq_params, + padding_mask=padding_mask, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, + input_ids=input_ids, + ) + elif isinstance(layer, (TransformerLayer, HyperConnectionHybridLayer)): layer_kwargs = dict( hidden_states=hidden_states, attention_mask=attention_mask, @@ -770,6 +928,12 @@ def get_inner_quant_context(config, layer_number): packed_seq_params=layer_packed_seq_params, padding_mask=padding_mask, ) + if rotary_pos_cos is not None: + layer_kwargs["rotary_pos_cos"] = rotary_pos_cos + if rotary_pos_sin is not None: + layer_kwargs["rotary_pos_sin"] = rotary_pos_sin + if rotary_pos_cos_sin is not None: + layer_kwargs["rotary_pos_cos_sin"] = rotary_pos_cos_sin if layer_cp_metadata is not None: layer_kwargs["packed_sequence_cp_metadata"] = layer_cp_metadata if residual_stream_recompute_context is not None: @@ -888,18 +1052,55 @@ def sharded_state_dict( dict: The sharded state dictionary for the current object. """ + return self._sharded_state_dict( + prefix=prefix, + sharded_offsets=sharded_offsets, + metadata=metadata, + sharded_layer_prefix=None, + ) + + def _sharded_state_dict( + self, + prefix: str = '', + sharded_offsets: Optional[tuple] = None, + metadata: Optional[dict] = None, + sharded_layer_prefix: Optional[str] = None, + ) -> ShardedStateDict: + """Build the sharded state dict using logical (bracketed-group aware) layer keys. + + ``sharded_layer_prefix`` is the ``layers.`` prefix of the outermost stack; + a nested group stack publishes all of its physical layers under the outer stack's + logical layer index so a ``[*-]`` group produces the same keys as one transformer + layer (``layers.N.self_attention.*`` and ``layers.N.mlp.*``). + """ sharded_offsets = sharded_offsets or () sharded_state_dict = {} layer_prefix = f'{prefix}layers.' + if sharded_layer_prefix is None: + sharded_layer_prefix = layer_prefix - for local_layer_idx, layer in enumerate(self.layers): - - global_layer_offset = layer.layer_number - 1 # self.layer_number starts at 1 - state_dict_prefix = ( - f'{layer_prefix}{local_layer_idx}.' # module list index in HybridStack + for local_layer_idx, (layer_config, layer) in enumerate( + zip(self.layer_config_list, self.layers, strict=True) + ): + state_dict_prefix = f'{layer_prefix}{local_layer_idx}.' # module list index + logical_layer_idx = ( + self.logical_layer_offset + if self.is_layer_group_stack + else self.logical_layer_offset + local_layer_idx ) - sharded_prefix = f'{layer_prefix}{global_layer_offset}.' + if is_layer_group(layer_config): + sharded_state_dict.update( + layer._sharded_state_dict( + state_dict_prefix, + sharded_offsets, + metadata, + sharded_layer_prefix=sharded_layer_prefix, + ) + ) + continue + + sharded_prefix = f'{sharded_layer_prefix}{logical_layer_idx}.' sharded_pp_offset = [] layer_sharded_state_dict = layer.sharded_state_dict( @@ -913,15 +1114,20 @@ def sharded_state_dict( # Add modules other than self.layers for name, module in self.named_children(): if not module is self.layers: - sharded_state_dict.update( - sharded_state_dict_default( - module, - f'{prefix}{name}.', - sharded_offsets, - metadata, - tp_group=self.tp_group, - ) + module_prefix = f'{prefix}{name}.' + module_sharded_state_dict = sharded_state_dict_default( + module, module_prefix, sharded_offsets, metadata, tp_group=self.tp_group ) + # The registered submodule stays ``final_norm`` (local state-dict keys + # are unchanged), but grouped stacks publish the sharded key under + # TransformerBlock's ``final_layernorm`` name so their checkpoints + # cross-load with GPTModel. Non-grouped stacks keep ``final_norm`` so + # hybrid checkpoints written before this feature still load. + if name == 'final_norm' and self.transformer_sharded_keys: + replace_prefix_for_sharding( + module_sharded_state_dict, module_prefix, f'{prefix}final_layernorm.' + ) + sharded_state_dict.update(module_sharded_state_dict) local_state_dict: dict = {} self._save_to_state_dict(local_state_dict, '', keep_vars=True) diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index e11ea9c1d24..b49b1066db5 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -2,7 +2,7 @@ import logging from dataclasses import dataclass -from typing import Dict, List, Optional, Sequence, Tuple +from typing import Dict, List, Optional, Sequence, Tuple, Union import torch @@ -15,6 +15,95 @@ logger = logging.getLogger(__name__) +# A parsed layer item is either a single layer symbol (e.g. ``'M'``) or a bracketed group of +# symbols (e.g. ``('M', '*')`` for ``[M*]``) that HybridStack builds as one nested logical layer. +LayerPatternItem = Union[str, Tuple[str, ...]] +# The per-layer config projection of ``LayerPatternItem``: one independent config per physical +# layer, with bracketed groups kept together as a tuple of configs. +LayerConfigItem = Union[TransformerConfig, Tuple[TransformerConfig, ...]] + + +def is_layer_group(layer_type) -> bool: + """Return whether a parsed layer item (symbol or config) is a bracketed group.""" + return isinstance(layer_type, tuple) + + +def flatten_layer_type_list(layer_type_list: Sequence) -> list: + """Flatten bracketed layer groups into their physical layer items (symbols or configs).""" + flattened = [] + for layer_type in layer_type_list: + if is_layer_group(layer_type): + flattened.extend(layer_type) + else: + flattened.append(layer_type) + return flattened + + +def get_layer_type_physical_count(layer_type) -> int: + """Return the number of physical layers represented by a parsed layer item.""" + return len(layer_type) if is_layer_group(layer_type) else 1 + + +def get_layer_type_logical_count(layer_type) -> int: + """Return the number of logical layers represented by a parsed layer item.""" + return 1 + + +def get_layer_type_list_physical_count(layer_type_list: Sequence) -> int: + """Return the number of physical layers represented by a parsed layer list.""" + return sum(get_layer_type_physical_count(layer_type) for layer_type in layer_type_list) + + +def get_layer_type_list_logical_count(layer_type_list: Sequence) -> int: + """Return the number of logical layers represented by a parsed layer list.""" + return sum(get_layer_type_logical_count(layer_type) for layer_type in layer_type_list) + + +def layer_type_item_to_str(layer_type: LayerPatternItem) -> str: + """Render one parsed layer item back to pattern syntax.""" + if is_layer_group(layer_type): + return f"{Symbols.GROUP_START}{''.join(layer_type)}{Symbols.GROUP_END}" + return layer_type + + +def layer_type_list_to_str(layer_type_list: Sequence[LayerPatternItem]) -> str: + """Render a parsed layer list back to pattern syntax.""" + return ''.join(layer_type_item_to_str(layer_type) for layer_type in layer_type_list) + + +def validate_layer_group(layer_types: Sequence[str]) -> None: + """Require group members to have distinct sharded checkpoint namespaces. + + A group shares one logical checkpoint layer index, so it can contain at most + one mixer, one attention module (including GDN), and one MLP or MoE module. + """ + if not layer_types: + raise ValueError("Layer groups cannot be empty.") + if Symbols.MOE in layer_types[:-1]: + raise ValueError(f"MoE layer '{Symbols.MOE}' must be the last symbol inside a layer group.") + namespaces = { + Symbols.MAMBA: "mixer", + Symbols.GDN: "self_attention", + Symbols.ATTENTION: "self_attention", + Symbols.DS_ATTENTION: "self_attention", + Symbols.MLA: "self_attention", + Symbols.CSA: "self_attention", + Symbols.HCA: "self_attention", + Symbols.WINDOW: "self_attention", + Symbols.MLP: "mlp", + Symbols.MOE: "mlp", + } + seen = set() + for layer_type in layer_types: + namespace = namespaces[layer_type] + if namespace in seen: + raise ValueError( + f"Layer group '{layer_type_list_to_str([tuple(layer_types)])}' contains " + f"multiple layers in checkpoint namespace '{namespace}'." + ) + seen.add(namespace) + + @dataclass class ParsedHybridPattern: """Result of parsing a unified hybrid pattern string. @@ -106,7 +195,8 @@ def get_hybrid_total_layer_count(pattern: str) -> int: """Returns the total number of main decoder layers in a hybrid layer pattern. Extracts the main pattern (before the first MTP separator '/'), strips - pipeline stage separators '|', and returns the character count. + pipeline stage separators '|', and returns the physical layer count + (bracketed groups contribute one layer per symbol inside the brackets). Args: pattern: Full hybrid layer pattern, possibly including MTP and pipe separators. @@ -116,7 +206,10 @@ def get_hybrid_total_layer_count(pattern: str) -> int: """ main_pattern = pattern.split(Symbols.MTP_SEPARATOR)[0] _validate_pattern(main_pattern, allow_pipe=True) - return len(main_pattern.replace(Symbols.PIPE, '')) + return sum( + get_layer_type_list_physical_count(parse_segment_layers(segment)) + for segment in main_pattern.split(Symbols.PIPE) + ) def get_hybrid_total_pipeline_segment_count(pattern: str) -> int: @@ -139,8 +232,9 @@ def get_hybrid_layer_counts(pattern: str) -> Dict[str, int]: """Count layers by type across the full hybrid pattern (main + MTP). Parses the pattern to extract main and MTP components, then counts - each layer type. Main pattern '|' separators are skipped. MTP layers - are counted once per MTP depth. + each layer type. Main pattern '|' separators are skipped and bracketed + groups are flattened to their physical layers. MTP layers are counted + once per MTP depth. Args: pattern: Full hybrid layer pattern string. @@ -161,15 +255,14 @@ def get_hybrid_layer_counts(pattern: str) -> Dict[str, int]: # Count main decoder layers (skip '|' pipe separators) if parsed.main_pattern: - for char in parsed.main_pattern: - if char in counts: + for segment in parsed.main_pattern.split(Symbols.PIPE): + for char in flatten_layer_type_list(parse_segment_layers(segment)): counts[char] += 1 # Count MTP layers (pattern repeated mtp_num_depths times) if parsed.mtp_pattern and parsed.mtp_num_depths > 0: - for char in parsed.mtp_pattern: - if char in counts: - counts[char] += parsed.mtp_num_depths + for char in flatten_layer_type_list(parse_segment_layers(parsed.mtp_pattern)): + counts[char] += parsed.mtp_num_depths return counts @@ -246,6 +339,16 @@ def parse_hybrid_pattern(pattern: Optional[str]) -> ParsedHybridPattern: # Decoder and MTP share the model's MLA mode and positional embeddings. _validate_pattern(main_pattern + mtp_pattern, allow_pipe=True) + # MTP layers are themselves a fused unit (each MTP depth contains its own attention + # + MLP), so it does not make sense to wrap them in a HybridStack group. Reject + # bracketed groups inside MTP patterns to keep downstream construction simple. + if Symbols.GROUP_START in mtp_pattern or Symbols.GROUP_END in mtp_pattern: + raise ValueError( + f"In MTP pattern, layer groups '{Symbols.GROUP_START}...{Symbols.GROUP_END}' " + f"are not supported because each MTP depth is already a fused unit. " + f"Got MTP pattern: '{mtp_pattern}'." + ) + return ParsedHybridPattern( main_pattern=main_pattern if main_pattern else None, mtp_pattern=mtp_pattern, @@ -253,56 +356,228 @@ def parse_hybrid_pattern(pattern: Optional[str]) -> ParsedHybridPattern: ) +def _invalid_symbol_message(char: str) -> str: + return ( + f"'{char}' is not a valid layer symbol. " + f"Valid symbols are: {Symbols.LAYER_CONFIG_MAP.keys()}" + ) + + def _validate_pattern(pattern: str, allow_pipe: bool = False) -> None: - """Validate that a pattern contains only valid layer symbols. + """Validate that a pattern contains only valid layer symbols and well-formed groups. Args: pattern: Layer pattern string to validate allow_pipe: Whether to allow the pipe '|' separator (for main patterns) Raises: - ValueError: If pattern contains invalid symbols + ValueError: If pattern contains invalid symbols or malformed bracketed groups """ - for char in pattern: - if not layer_utils.is_valid_symbol(char, allow_pipe=allow_pipe): - raise ValueError( - f"'{char}' is not a valid layer symbol. " - f"Valid symbols are: {Symbols.LAYER_CONFIG_MAP.keys()}" - ) + if not allow_pipe and Symbols.PIPE in pattern: + raise ValueError(_invalid_symbol_message(Symbols.PIPE)) + + flat_layers = [] + for segment in pattern.split(Symbols.PIPE): + flat_layers.extend(flatten_layer_type_list(parse_segment_layers(segment))) # MLA variants may coexist, but standard attention cannot share a model with them. - if Symbols.ATTENTION in pattern and any(symbol in pattern for symbol in Symbols.MLA_ATTENTION): + if Symbols.ATTENTION in flat_layers and any( + symbol in flat_layers for symbol in Symbols.MLA_ATTENTION + ): raise ValueError( "Not supported to have both Attention and MLA/DSA/CSA/HCA/Window in one model" ) -def validate_segment_layers(segment: str, config: TransformerConfig) -> List[TransformerConfig]: +def parse_segment_layers(segment: str) -> List[LayerPatternItem]: + """Parse a pipe-free pattern segment into layer symbols and bracketed groups. + + Bracketed groups such as ``[M*E]`` become tuples of symbols (``('M', '*', 'E')``); every + other valid symbol is returned as-is. Groups cannot be empty or nested, and + their layers must have distinct checkpoint namespaces. An MoE layer inside a + group must be its last symbol so that EP-overlap scheduling can split the + group into pre-dispatch compute and the terminal MoE layer. + + Args: + segment: A single pipeline segment pattern string (e.g., "M[M*]-"). + + Returns: + List of layer symbols and symbol tuples in pattern order. + + Raises: + ValueError: If the segment contains invalid symbols or malformed groups. + """ + layer_type_list: List[LayerPatternItem] = [] + flat_layers = [] + i = 0 + while i < len(segment): + layer_char = segment[i] + if layer_char == Symbols.GROUP_START: + group_end = segment.find(Symbols.GROUP_END, i + 1) + if group_end == -1: + raise ValueError( + f"'{Symbols.GROUP_START}' starts a layer group without a matching " + f"'{Symbols.GROUP_END}'." + ) + group = segment[i + 1 : group_end] + if group == "": + raise ValueError("Layer groups cannot be empty.") + if Symbols.GROUP_START in group or Symbols.GROUP_END in group: + raise ValueError("Nested layer groups are not supported.") + for group_char in group: + if not layer_utils.is_valid_symbol(group_char): + raise ValueError(_invalid_symbol_message(group_char)) + validate_layer_group(group) + group_tuple = tuple(group) + layer_type_list.append(group_tuple) + flat_layers.extend(group_tuple) + i = group_end + 1 + continue + if layer_char == Symbols.GROUP_END: + raise ValueError(f"'{Symbols.GROUP_END}' closes a layer group that was not opened.") + if not layer_utils.is_valid_symbol(layer_char): + raise ValueError(_invalid_symbol_message(layer_char)) + layer_type_list.append(layer_char) + flat_layers.append(layer_char) + i += 1 + + # MLA variants may coexist, but standard attention cannot share a model with them. + if Symbols.ATTENTION in flat_layers and any( + symbol in flat_layers for symbol in Symbols.MLA_ATTENTION + ): + raise ValueError( + "Not supported to have both Attention and MLA/DSA/CSA/HCA/Window in one model" + ) + + return layer_type_list + + +def validate_segment_layers(segment: str, config: TransformerConfig) -> List[LayerConfigItem]: """Validate and convert a single pipeline segment pattern to layer configs. This is used after the main pattern has been split by '|' into segments. - Each segment should contain only valid layer symbols (no '|'). + Each segment should contain only valid layer symbols (no '|'), optionally + grouped with brackets (e.g. ``M[M*]-``). Each layer config is copied from the source config without running ``__post_init__`` - a second time. + a second time. Bracketed groups are returned as a tuple of per-layer configs so that + ``HybridStack`` can build them as one nested logical layer. Args: - segment: A single pipeline segment pattern string (e.g., "M-M*-") + segment: A single pipeline segment pattern string (e.g., "M-M*-" or "M[M*]-") config: Normalized stack-level config to copy for each layer. Returns: - List of independent per-layer configs. + List of independent per-layer configs, with groups kept as tuples of configs. Raises: - ValueError: If segment contains invalid layer symbols. + ValueError: If segment contains invalid layer symbols or malformed groups. + """ + layer_config_list: List[LayerConfigItem] = [] + for layer_type in parse_segment_layers(segment): + if is_layer_group(layer_type): + layer_config_list.append( + tuple(layer_utils.create_layer_config(config, symbol) for symbol in layer_type) + ) + else: + layer_config_list.append(layer_utils.create_layer_config(config, layer_type)) + + return layer_config_list + + +def _slice_layer_type_list_by_physical_range( + layer_type_list: List[LayerPatternItem], offset: int, count: int +) -> List[LayerPatternItem]: + """Slice parsed layer items by physical layer range without splitting groups.""" + selected = [] + cursor = 0 + end = offset + count + for layer_type in layer_type_list: + item_count = get_layer_type_physical_count(layer_type) + item_end = cursor + item_count + if item_end <= offset: + cursor = item_end + continue + if cursor >= end: + break + if cursor < offset or item_end > end: + raise ValueError( + "Pipeline splitting would split a bracketed hybrid layer group. " + "Add pipe ('|') separators around bracketed groups to define valid boundaries." + ) + selected.append(layer_type) + cursor = item_end + return selected + + +def _get_logical_offset_from_physical_offset( + layer_type_list: List[LayerPatternItem], offset: int +) -> int: + """Return the logical item count before a physical-layer offset.""" + logical_offset = 0 + cursor = 0 + for layer_type in layer_type_list: + item_count = get_layer_type_physical_count(layer_type) + item_end = cursor + item_count + if item_end <= offset: + logical_offset += get_layer_type_logical_count(layer_type) + cursor = item_end + continue + if cursor == offset: + return logical_offset + raise ValueError( + "Pipeline splitting would split a bracketed hybrid layer group. " + "Add pipe ('|') separators around bracketed groups to define valid boundaries." + ) + if cursor == offset: + return logical_offset + raise ValueError(f"Physical layer offset {offset} is out of range for hybrid layer pattern.") + + +def select_pipeline_segment_with_logical_offset( + main_pattern: str, + config: TransformerConfig, + pp_group: Optional[torch.distributed.ProcessGroup], + vp_stage: Optional[int], + first_stage_layers: Optional[int] = None, + last_stage_layers: Optional[int] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + dp_cp_group: Optional[torch.distributed.ProcessGroup] = None, +) -> Tuple[List[LayerConfigItem], int, int]: + """Select a pipeline segment and return physical and logical offsets. + + The physical offset counts every layer symbol before this segment; the logical + offset counts pattern items, so a bracketed group before this segment adds one. + See :func:`select_pipeline_segment` for the argument semantics. """ - _validate_pattern(segment) + layer_config_list, layer_offset = select_pipeline_segment( + main_pattern, + config, + pp_group, + vp_stage, + first_stage_layers=first_stage_layers, + last_stage_layers=last_stage_layers, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) - layer_configs: list[TransformerConfig] = [] - for layer_symbol in segment: - layer_configs.append(layer_utils.create_layer_config(config, layer_symbol)) + segments = main_pattern.split(Symbols.PIPE) if main_pattern else [''] + if len(segments) == 1: + full_layer_type_list = parse_segment_layers(segments[0]) + logical_layer_offset = _get_logical_offset_from_physical_offset( + full_layer_type_list, layer_offset + ) + else: + pp_rank = torch.distributed.get_rank(pp_group) if pp_group is not None else 0 + pp_size = torch.distributed.get_world_size(pp_group) if pp_group is not None else 1 + vp_rel = vp_stage if vp_stage is not None else 0 + segment_index = vp_rel * pp_size + pp_rank + logical_layer_offset = sum( + get_layer_type_list_logical_count(parse_segment_layers(segments[i])) + for i in range(segment_index) + ) - return layer_configs + return layer_config_list, layer_offset, logical_layer_offset def select_pipeline_segment( @@ -314,7 +589,7 @@ def select_pipeline_segment( last_stage_layers: Optional[int] = None, tp_group: Optional[torch.distributed.ProcessGroup] = None, dp_cp_group: Optional[torch.distributed.ProcessGroup] = None, -) -> Tuple[List[TransformerConfig], int]: +) -> Tuple[List[LayerConfigItem], int]: """Select and validate the pipeline segment for the given PP rank and VP stage. When the main pattern contains '|' pipe separators, splits by '|' into @@ -322,7 +597,8 @@ def select_pipeline_segment( When the pattern has no pipes but pp_size > 1, falls back to runtime layer slicing (for backwards compatibility), supporting both even and uneven PP splits - via first_stage_layers / last_stage_layers. + via first_stage_layers / last_stage_layers. Layer counts are physical layer + counts, and a split may not fall inside a bracketed group. Args: main_pattern: Main decoder pattern (may contain '|' separators). @@ -339,8 +615,9 @@ def select_pipeline_segment( Returns: Tuple of (layer_config_list, layer_offset) where layer_config_list is - the list of independent configs for this segment, and layer_offset - is the sum of layer counts from all preceding segments. + the list of independent configs for this segment (bracketed groups as + tuples of configs), and layer_offset is the sum of physical layer counts + from all preceding segments. Raises: ValueError: If the segment contains invalid layer symbols, if @@ -378,8 +655,8 @@ def select_pipeline_segment( "Example: 'M*M*M*M*' with pp_size=2 should become 'M*M*|M*M*'.", ) full_pattern = segments[0] - _validate_pattern(full_pattern) - num_layers = len(full_pattern) + layer_type_list = parse_segment_layers(full_pattern) + num_layers = get_layer_type_list_physical_count(layer_type_list) if first_stage_layers is not None or last_stage_layers is not None: first = first_stage_layers or 0 @@ -422,14 +699,17 @@ def select_pipeline_segment( offset = pp_rank * layers_per_rank count = layers_per_rank - selected_pattern = full_pattern[offset : offset + count] + selected_pattern = layer_type_list_to_str( + _slice_layer_type_list_by_physical_range(layer_type_list, offset, count) + ) layer_utils.validate_tp_comm_overlap(config, selected_pattern) selected = validate_segment_layers(selected_pattern, config) log_on_each_pipeline_stage( logger, logging.INFO, f"HybridModel: pp_rank={pp_rank}/{pp_size}, vp_stage={vp_stage}, " - f"layers='{selected_pattern}' ({len(selected)} layers), " + f"layers='{selected_pattern}' " + f"({get_layer_type_list_physical_count(selected)} layers), " f"layer_offset={offset} (auto-split)", tp_group=tp_group, dp_cp_group=dp_cp_group, @@ -455,7 +735,10 @@ def select_pipeline_segment( f"the current PP/VPP configuration." ) - layer_offset = sum(len(segments[i]) for i in range(segment_index)) + layer_offset = sum( + get_layer_type_list_physical_count(parse_segment_layers(segments[i])) + for i in range(segment_index) + ) my_segment = segments[segment_index] layer_utils.validate_tp_comm_overlap(config, my_segment) @@ -466,7 +749,8 @@ def select_pipeline_segment( logging.INFO, f"HybridModel: pp_rank={pp_rank}/{pp_size}, vp_stage={vp_rel}, " f"segment_index={segment_index}/{len(segments)}, " - f"layers='{my_segment}' ({len(layer_config_list)} layers), " + f"layers='{my_segment}' " + f"({get_layer_type_list_physical_count(layer_config_list)} layers), " f"layer_offset={layer_offset}", tp_group=tp_group, dp_cp_group=dp_cp_group, @@ -475,14 +759,17 @@ def select_pipeline_segment( return layer_config_list, layer_offset -def get_layer_maps_from_layer_type_list(layer_type_list: list[str]) -> dict[str, dict[int, int]]: +def get_layer_maps_from_layer_type_list( + layer_type_list: list[LayerPatternItem], +) -> dict[str, dict[int, int]]: """ Returns maps from global layer index to the corresponding layer index for each valid layer type (the keys of Symbols.LAYER_CONFIG_MAP) given a layer type list. + Bracketed groups are flattened to their physical layers first. """ layer_types = [symbol for symbol in Symbols.name_sorted_valid_layer_symbols()] layer_maps = {layer_type: {} for layer_type in layer_types} - for global_layer_idx, layer_type in enumerate(layer_type_list): + for global_layer_idx, layer_type in enumerate(flatten_layer_type_list(layer_type_list)): layer_map = layer_maps[layer_type] local_layer_idx = len(layer_map) layer_map[global_layer_idx] = local_layer_idx @@ -490,19 +777,26 @@ def get_layer_maps_from_layer_type_list(layer_type_list: list[str]) -> dict[str, def get_layer_type_list_from_layer_config_list( - layer_config_list: Sequence[TransformerConfig], -) -> list[str]: + layer_config_list: Sequence[LayerConfigItem], +) -> list[LayerPatternItem]: """Return the layer symbols corresponding to a sequence of layer configs. This compatibility projection keeps ``layer_config_list`` as the source of truth while - supporting callers that still read ``HybridStack.layer_type_list``. + supporting callers that still read ``HybridStack.layer_type_list``. Bracketed groups + (tuples of configs) map to tuples of symbols. Args: layer_config_list: Per-layer configs in layer order. Returns: - The canonical layer symbol for each config. + The canonical layer symbol (or symbol tuple) for each config item. """ - return [ - layer_utils.get_layer_symbol_from_config(layer_config) for layer_config in layer_config_list - ] + layer_type_list: list[LayerPatternItem] = [] + for layer_config in layer_config_list: + if is_layer_group(layer_config): + layer_type_list.append( + tuple(layer_utils.get_layer_symbol_from_config(config) for config in layer_config) + ) + else: + layer_type_list.append(layer_utils.get_layer_symbol_from_config(layer_config)) + return layer_type_list diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index d846d5aaa2a..2e5c7a32f4a 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -2,7 +2,7 @@ import logging from contextlib import nullcontext -from typing import Literal, Optional +from typing import Any, Callable, Dict, Literal, Optional import torch from torch import Tensor @@ -10,6 +10,7 @@ from megatron.core import tensor_parallel from megatron.core.config_logger import has_config_logger_enabled, log_config_to_disk from megatron.core.context_parallel import ContextParallelBatch +from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding @@ -62,9 +63,15 @@ def _get_hash_moe_layer_threshold(main_pattern: str | None, n_hash_layers: int) if n_hash_layers <= 0: return 0 - from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols + from megatron.core.models.hybrid.hybrid_layer_allocation import ( + Symbols, + flatten_layer_type_list, + parse_segment_layers, + ) - global_layer_pattern = (main_pattern or '').replace(Symbols.PIPE, '') + global_layer_pattern = flatten_layer_type_list( + parse_segment_layers((main_pattern or '').replace(Symbols.PIPE, '')) + ) moe_layer_numbers = [ layer_number for layer_number, layer_type in enumerate(global_layer_pattern, start=1) @@ -79,17 +86,22 @@ def _get_hash_moe_layer_threshold(main_pattern: str | None, n_hash_layers: int) def _validate_hash_moe_pipeline_placement( - layer_type_list: list[str], layer_offset: int, hash_moe_layer_threshold: int, pre_process: bool + layer_type_list: list[str | tuple[str, ...]], + layer_offset: int, + hash_moe_layer_threshold: int, + pre_process: bool, ) -> None: """Reject local hash-MoE layers on a stage that does not own the token IDs.""" if hash_moe_layer_threshold <= 0 or pre_process: return - from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols + from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols, flatten_layer_type_list local_hash_layer_numbers = [ layer_offset + local_layer_number - for local_layer_number, layer_type in enumerate(layer_type_list, start=1) + for local_layer_number, layer_type in enumerate( + flatten_layer_type_list(layer_type_list), start=1 + ) if layer_type == Symbols.MOE and layer_offset + local_layer_number <= hash_moe_layer_threshold ] @@ -244,9 +256,10 @@ def __init__( # Parse unified pattern to extract main and MTP components. from megatron.core.models.hybrid.hybrid_layer_allocation import ( + Symbols, get_layer_type_list_from_layer_config_list, parse_hybrid_pattern, - select_pipeline_segment, + select_pipeline_segment_with_logical_offset, ) parsed = parse_hybrid_pattern(self.hybrid_layer_pattern) @@ -318,16 +331,27 @@ def __init__( # before constructing the decoder or MTP modules. layer_utils.validate_tp_comm_overlap(self.config, '', has_mtp=self.mtp_process) + # Bracketed-group patterns give every logical layer the structure of a + # transformer layer, so their checkpoints are made key-compatible with + # GPTModel. Derived from the full pattern rather than this rank's segment so + # every PP stage agrees on the naming. Non-grouped patterns keep the + # historical hybrid keys, which existing hybrid checkpoints were saved with. + transformer_sharded_keys = Symbols.GROUP_START in (parsed.main_pattern or '') + logging_pg_kwargs = _hybrid_logging_pg_kwargs(self.pg_collection) - layer_config_list, layer_offset = select_pipeline_segment( - parsed.main_pattern or '', - self.config, - self.pg_collection.pp, - vp_stage, - first_stage_layers=self.config.num_layers_in_first_pipeline_stage, - last_stage_layers=self.config.num_layers_in_last_pipeline_stage, - **logging_pg_kwargs, + # Bracketed groups (e.g. ``[M*E]``) count as one logical layer for checkpoint + # keys but several physical layers for layer numbering, so track both offsets. + layer_config_list, layer_offset, logical_layer_offset = ( + select_pipeline_segment_with_logical_offset( + parsed.main_pattern or '', + self.config, + self.pg_collection.pp, + vp_stage, + first_stage_layers=self.config.num_layers_in_first_pipeline_stage, + last_stage_layers=self.config.num_layers_in_last_pipeline_stage, + **logging_pg_kwargs, + ) ) _validate_hash_moe_pipeline_placement( get_layer_type_list_from_layer_config_list(layer_config_list), @@ -388,6 +412,8 @@ def __init__( pre_process=self.pre_process, layer_config_list=layer_config_list, pp_layer_offset=layer_offset, + logical_layer_offset=logical_layer_offset, + transformer_sharded_keys=transformer_sharded_keys, post_process=self.post_process, dtype=config.params_dtype, pg_collection=self.pg_collection, @@ -532,64 +558,36 @@ def create_mcore_cudagraph_manager(self, config): self.cudagraph_manager = CudaGraphManager(config) - def forward( + def _preprocess( self, input_ids: Tensor, position_ids: Tensor, - attention_mask: Tensor, decoder_input: Tensor = None, - labels: Tensor = None, inference_context: BaseInferenceContext = None, - runtime_gather_output: Optional[bool] = None, - *, - inference_params: Optional[BaseInferenceContext] = None, - loss_mask: Optional[Tensor] = None, - mtp_input_mask: Optional[Tensor] = None, - packed_seq_params: Optional[PackedSeqParams] = None, + packed_seq_params: PackedSeqParams = None, padding_mask: Optional[Tensor] = None, - compute_mtp_loss: bool = True, - cp_batch: ContextParallelBatch | None = None, - ) -> Tensor: - """Forward function of the Hybrid model. This function passes the input tensors - through the embedding layer, and then the decoder and finally into the post - processing layer (optional). - - It either returns the Loss values if labels are given or the final hidden units - - Args: - compute_mtp_loss (bool): Whether to compute the non-inference MTP auxiliary - objective. Disabling it skips the MTP branch while leaving its parameters - loaded. This does not control speculative decoding. On post-process stages, - ``labels`` still determine whether the model returns loss or logits. - Defaults to True. - cp_batch: Input tensors and packed metadata keyed by CP layout. + ): + """Preprocess inputs for HybridStack or combined-1F1B scheduling. + + Mirrors ``GPTModel._preprocess`` so the eager forward and the + EP-overlap ``PreProcessNode`` see the same embedding / rotary / + padding-mask code. Returns the canonical 6-tuple ``(decoder_input, + rotary_pos_emb, rotary_pos_cos, rotary_pos_sin, sequence_len_offset, + padding_mask)`` -- slots HybridModel does not compute (rotary cos/sin) + come back as ``None``. """ - # If decoder_input is provided (not None), then input_ids and position_ids are ignored. - # Otherwise, apply embedding layer on input_ids and position_ids to get decoder_input. - - if self.config.fine_grained_activation_offloading: - self.preprocess_for_fine_grained_offloading() - - if self.config.moe_paged_stash: - self.preprocess_for_paged_stash() - - inference_context = deprecate_inference_params(inference_context, inference_params) - in_inference_mode = InferenceMode.is_active() - if in_inference_mode: - assert runtime_gather_output, "Inference must always gather TP logits" - - use_precomputed_mtp_embeddings = decoder_input is not None - - # Decoder embedding. + # If decoder_input is provided, input_ids and position_ids are ignored; + # otherwise apply the embedding layer to get decoder_input. if decoder_input is not None: pass elif self.pre_process: + # Decoder embedding. decoder_input = self.embedding(input_ids=input_ids, position_ids=position_ids) - # Clear the outputs for padding tokens when using dynamic batching with - # quantization scales to avoid corrupting amax calculations + # Clear the outputs for padding tokens when using dynamic batching + # with quantization scales to avoid corrupting amax calculations. if ( in_inference_mode and inference_context is not None @@ -605,29 +603,10 @@ def forward( decoder_input, group=self.pg_collection.tp ) else: - # intermediate stage of pipeline - # decoder will get hidden_states from encoder.input_tensor + # Intermediate stage of pipeline parallelism -- the decoder will get + # hidden_states from encoder.input_tensor. decoder_input = None - # Hash routing consumes batch-major token IDs. Under sequence parallelism, - # shard them with decoder activations so each TP rank hashes its local tokens. - hash_input_ids = None - if self.config.moe_num_hash_layers > 0: - hash_input_ids = input_ids - if ( - self.config.sequence_parallel - and decoder_input is not None - and hash_input_ids is not None - and hash_input_ids.shape[1] != decoder_input.shape[0] - ): - hash_input_ids = ( - tensor_parallel.scatter_to_sequence_parallel_region( - hash_input_ids.transpose(0, 1).contiguous(), group=self.pg_collection.tp - ) - .transpose(0, 1) - .contiguous() - ) - # TODO: Apply the same later-stage SP mask alignment in GPTModel. # Later pipeline stages receive activations through set_input_tensor. decoder_reference = decoder_input @@ -660,82 +639,112 @@ def forward( rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( inference_context, self.decoder, decoder_input, self.config, packed_seq_params ) - # YarnRotaryEmbedding.forward returns (emb, mscale); discard mscale here + # YarnRotaryEmbedding.forward returns (emb, mscale); discard mscale here. rotary_pos_emb, _ = self.rotary_pos_emb( rotary_seq_len, packed_seq=packed_seq_params is not None and packed_seq_params.qkv_format == 'thd', ) - # Wrap decoder_input to allow the decoder (HybridStack) to delete the - # reference held by this caller function, enabling early garbage collection - # for inference. - if in_inference_mode: - decoder_input = WrappedTensor(decoder_input) - - # The following assert will currently fail when running inference. - # Commented out for now. - # TODO (duncan/rwaleffe): (1) confirm that the externally-generated - # attention mask is not needed and is ignored by the model in - # inference mode, (2) reduce the size of the externally-generated - # attention mask to prevent CPU OOM (as we did for training), (3) - # force the attention mask passed to the model in inference mode to - # be None, so this assert will succeed. - # assert attention_mask is None, "The attention mask is ignored and should be set to None" - - packed_seq_params_by_layout = ( - cp_batch.packed_seq_params_by_layout if cp_batch is not None else None - ) - cp_layout_plan = cp_batch.thd_plan if cp_batch is not None else None - - # Run decoder. - backbone_context = ( - torch.no_grad() - if self.config.freeze_base_model_for_mtp and self.training - else nullcontext() - ) - with backbone_context: - decoder_output = self.decoder( - hidden_states=decoder_input, - attention_mask=attention_mask, - inference_context=inference_context, - rotary_pos_emb=rotary_pos_emb, - packed_seq_params=packed_seq_params, - padding_mask=padding_mask, - packed_seq_params_by_layout=packed_seq_params_by_layout, - cp_layout_plan=cp_layout_plan, - input_ids=hash_input_ids, + # ``sequence_len_offset`` is only needed for flash-decode / local-cudagraph + # static-batching inference; otherwise leave it as ``None``. + if ( + in_inference_mode + and inference_context is not None + and (self.config.cuda_graph_impl == "local" or self.config.flash_decode) + and inference_context.is_static_batching() + ): + current_batch_size = input_ids.shape[0] + sequence_len_offset = torch.tensor( + [inference_context.sequence_len_offset] * current_batch_size, + dtype=torch.int32, + device='cuda', ) - if isinstance(decoder_output, tuple): - hidden_states, mhc_multistream = decoder_output else: - hidden_states = decoder_output - mhc_multistream = None + sequence_len_offset = None + + return decoder_input, rotary_pos_emb, None, None, sequence_len_offset, padding_mask + + def _postprocess( + self, + hidden_states, + input_ids, + position_ids, + labels, + rotary_pos_emb, + rotary_pos_cos=None, + rotary_pos_sin=None, + mtp_in_postprocess=None, + loss_mask=None, + mtp_input_mask=None, + decoder_input=None, + attention_mask=None, + padding_mask=None, + inference_params=None, + packed_seq_params=None, + sequence_len_offset=None, + runtime_gather_output=None, + extra_block_kwargs=None, + inference_context=None, + is_spec_decode=None, + output_processor=None, + output_processor_context=None, + compute_mtp_loss=True, + cp_batch=None, + mhc_multistream=None, + mtp_decoder_input=None, + ): + """Postprocess HybridStack hidden states into logits or language-model loss. + + Mirrors ``GPTModel._postprocess`` so the eager forward and the EP-overlap + ``PostProcessNode`` produce the same logits / loss / MTP outputs. + ``mtp_in_postprocess`` lets the EP-overlap path skip the inline MTP block + (it schedules MTP as separate layer nodes and hands over the concatenated + main / MTP hidden states); the eager forward leaves it ``True`` so the regular + MTP forward runs here. ``compute_mtp_loss`` disables the MTP auxiliary + objective (see ``forward``); ``cp_batch`` / ``mhc_multistream`` feed the MTP + block's CP-layout preparation on the eager path. + ``output_processor`` replaces the default logits / loss computation with a + caller-supplied hook (used by RL and other custom output paths); it is + forwarded by ``PostProcessNode`` and handled here exactly as in + ``GPTModel._postprocess``. + """ + in_inference_mode = InferenceMode.is_active() + if in_inference_mode: + assert runtime_gather_output, "Inference must always gather TP logits" output_weight = None if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() - # Check if speculative decoding is active. When it is, MTP must be - # computed *after* verification so that it is conditioned on verified - # tokens rather than stale speculative tokens from the previous step. - is_spec_decode = ( - in_inference_mode - and inference_context is not None - and inference_context.is_dynamic_batching() - and inference_context.num_speculative_tokens > 0 - ) + # Speculative decoding: when active, MTP must run *after* verification so + # it conditions on verified tokens rather than stale speculative ones. + if is_spec_decode is None: + is_spec_decode = ( + in_inference_mode + and inference_context is not None + and inference_context.is_dynamic_batching() + and inference_context.num_speculative_tokens > 0 + ) - mtp_forward_ran = ( + # Whether the MTP auxiliary objective is computed in this call. + # ``self.mtp_process`` guards against models built without an MTP block. + mtp_loss_active = ( self.mtp_process and not (in_inference_mode or is_spec_decode) and compute_mtp_loss ) + # MTP forward inline (skipped when the EP-overlap plan schedules MTP separately). + mtp_forward_ran = bool(mtp_in_postprocess) and mtp_loss_active mtp_hidden_states = hidden_states mtp_inputs = None if mtp_forward_ran: + packed_seq_params_by_layout = ( + cp_batch.packed_seq_params_by_layout if cp_batch is not None else None + ) + cp_layout_plan = cp_batch.thd_plan if cp_batch is not None else None mtp_inputs = self.mtp.prepare_cp_layout( input_ids=input_ids, position_ids=position_ids, hidden_states=hidden_states, - decoder_input=decoder_input if use_precomputed_mtp_embeddings else None, + decoder_input=mtp_decoder_input, mhc_multistream=mhc_multistream, labels=labels, loss_mask=loss_mask, @@ -770,11 +779,13 @@ def forward( if self.config.mtp_num_layers is not None and self.mtp_process: assert self.config.mtp_num_layers > 0 if is_spec_decode: + # Cache decoder hidden states for serial MTP computation after + # speculative token verification. assert inference_context is not None if self.config.inference_cuda_graph_scope == InferenceCudaGraphScope.block: - # Block-scope CUDA graph mode: copy_() into the - # pre-allocated buffer so every graph replay writes to - # the same fixed GPU address regardless of batch size. + # Block-scope CUDA graph mode: copy_() into the pre-allocated buffer so + # every graph replay writes to the same fixed GPU address regardless of + # batch size. assert inference_context.mtp_decoder_hidden_states is not None inference_context.mtp_decoder_hidden_states[: hidden_states.shape[0]].copy_( hidden_states @@ -783,14 +794,33 @@ def forward( # Non-block scope: direct assignment; the controller will set # this back to None after reading to allow GC. inference_context.mtp_decoder_hidden_states = hidden_states - elif mtp_forward_ran: - assert mtp_inputs is not None + elif mtp_loss_active: + # In training/eval, fold the MTP loss into hidden_states. When the MTP block + # ran inline above, use its CP-layout-prepared inputs; on the EP-overlap path + # the MTP layer nodes already produced the concatenated hidden states and the + # raw inputs apply (mirrors ``GPTModel._postprocess``). # For RL (labels is None), process_mtp_loss derives labels from # input_ids to match the SFT label format. + if mtp_inputs is not None: + mtp_loss_inputs = dict( + hidden_states=mtp_hidden_states, + labels=mtp_inputs.labels, + loss_mask=mtp_inputs.loss_mask, + packed_seq_params=mtp_inputs.packed_seq_params, + input_ids=mtp_inputs.input_ids, + mtp_input_mask=mtp_inputs.mtp_input_mask, + main_hidden_states=hidden_states, + ) + else: + mtp_loss_inputs = dict( + hidden_states=hidden_states, + labels=labels, + loss_mask=loss_mask, + packed_seq_params=packed_seq_params, + input_ids=input_ids, + mtp_input_mask=mtp_input_mask, + ) hidden_states = process_mtp_loss( - hidden_states=mtp_hidden_states, - labels=mtp_inputs.labels, - loss_mask=mtp_inputs.loss_mask, output_layer=self.output_layer, output_weight=output_weight, runtime_gather_output=runtime_gather_output, @@ -799,17 +829,36 @@ def forward( config=self.config, cp_group=self.pg_collection.cp, tp_group=self.tp_group, - packed_seq_params=mtp_inputs.packed_seq_params, scale_logits_fn=self._scale_logits if self.config.use_mup else None, - input_ids=mtp_inputs.input_ids, - mtp_input_mask=mtp_inputs.mtp_input_mask, metric_avg_group=( getattr(self.pg_collection, 'dp_cp_gtp_remat', None) or self.pg_collection.dp_cp ), - main_hidden_states=hidden_states, + **mtp_loss_inputs, ) + sequence_parallel_override = False + + if output_processor is not None: + return output_processor( + hidden_states=hidden_states, + output_layer=self.output_layer, + output_weight=output_weight, + labels=labels, + loss_mask=loss_mask, + input_ids=input_ids, + position_ids=position_ids, + attention_mask=attention_mask, + decoder_input=decoder_input, + inference_context=inference_context, + packed_seq_params=packed_seq_params, + runtime_gather_output=runtime_gather_output, + context=output_processor_context, + compute_language_model_loss=self.compute_language_model_loss, + scale_logits=self._scale_logits, + config=self.config, + ) + if ( in_inference_mode and inference_context is not None @@ -819,9 +868,9 @@ def forward( hidden_states = hidden_states[-1:, :, :] else: if self.output_layer.sequence_parallel: - # Perform the sequence parallel gather here instead of after the output layer - # because we need to slice the last token logits from the full view of the - # packed logits across all requests. + # Perform the sequence-parallel gather here instead of after + # the output layer so we can slice the last-token logits from + # the full view of the packed logits across all requests. hidden_states = gather_from_sequence_parallel_region( hidden_states, group=self.pg_collection.tp ) @@ -868,3 +917,221 @@ def forward( loss = self.compute_language_model_loss(labels, logits) return loss + + def build_schedule_plan( + self, + input_ids: Tensor, + position_ids: Tensor, + attention_mask: Tensor, + decoder_input: Tensor = None, + labels: Tensor = None, + inference_context: BaseInferenceContext = None, + packed_seq_params: PackedSeqParams = None, + extra_block_kwargs: dict = None, + runtime_gather_output: Optional[bool] = None, + inference_params: Optional[BaseInferenceContext] = None, + loss_mask: Optional[Tensor] = None, + padding_mask: Optional[Tensor] = None, + *, + mtp_input_mask: Optional[Tensor] = None, + output_processor: Optional[Callable[..., Any]] = None, + output_processor_context: Optional[Any] = None, + ): + """Build the HybridModel combined-1F1B schedule plan. + + Mirrors ``GPTModel.build_schedule_plan``; ``mtp_input_mask``, ``output_processor`` + and ``output_processor_context`` are stored on the chunk state so the schedule + plan's MTP layer nodes and ``PostProcessNode`` apply them. + """ + if self.config.fine_grained_activation_offloading: + self.preprocess_for_fine_grained_offloading() + if self.config.moe_paged_stash: + self.preprocess_for_paged_stash() + + from .model_chunk_schedule_plan import HybridStackModelChunkSchedulePlan + + return HybridStackModelChunkSchedulePlan( + self, + input_ids, + position_ids, + attention_mask, + decoder_input, + labels, + packed_seq_params, + extra_block_kwargs, + runtime_gather_output, + loss_mask, + padding_mask, + mtp_input_mask=mtp_input_mask, + output_processor=output_processor, + output_processor_context=output_processor_context, + ) + + def sharded_state_dict( + self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[Dict] = None + ) -> ShardedStateDict: + """Return a Transformer-compatible sharded state dict for HybridModel.""" + sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata) + output_layer_extra_state_key = f'{prefix}output_layer._extra_state' + + # Match GPTModel checkpoint compatibility: old GPT checkpoints do not include + # output layer extra state, and the TE extra state should be empty. + output_extra_state = sharded_state_dict.pop(output_layer_extra_state_key, None) + assert not ( + output_extra_state and output_extra_state.data + ), f'Expected output layer extra state to be empty, got: {output_extra_state}' + + return sharded_state_dict + + def forward( + self, + input_ids: Tensor, + position_ids: Tensor, + attention_mask: Tensor, + decoder_input: Tensor = None, + labels: Tensor = None, + inference_context: BaseInferenceContext = None, + runtime_gather_output: Optional[bool] = None, + *, + inference_params: Optional[BaseInferenceContext] = None, + loss_mask: Optional[Tensor] = None, + mtp_input_mask: Optional[Tensor] = None, + packed_seq_params: Optional[PackedSeqParams] = None, + padding_mask: Optional[Tensor] = None, + compute_mtp_loss: bool = True, + cp_batch: ContextParallelBatch | None = None, + ) -> Tensor: + """Forward function of the Hybrid model. This function passes the input tensors + through the embedding layer, and then the decoder and finally into the post + processing layer (optional). + + It either returns the Loss values if labels are given or the final hidden units + + Args: + compute_mtp_loss (bool): Whether to compute the non-inference MTP auxiliary + objective. Disabling it skips the MTP branch while leaving its parameters + loaded. This does not control speculative decoding. On post-process stages, + ``labels`` still determine whether the model returns loss or logits. + Defaults to True. + cp_batch: Input tensors and packed metadata keyed by CP layout. + """ + if self.config.fine_grained_activation_offloading: + self.preprocess_for_fine_grained_offloading() + + if self.config.moe_paged_stash: + self.preprocess_for_paged_stash() + + inference_context = deprecate_inference_params(inference_context, inference_params) + + in_inference_mode = InferenceMode.is_active() + if in_inference_mode: + assert runtime_gather_output, "Inference must always gather TP logits" + + mtp_decoder_input = decoder_input + + # Mirror GPTModel.forward: delegate the embedding / rotary computation and + # the output-layer / MTP / loss computation to the same hooks the + # EP-overlap PreProcessNode / PostProcessNode call. Keeps the eager and + # combined-1F1B paths on the same code. + ( + decoder_input, + rotary_pos_emb, + rotary_pos_cos, + rotary_pos_sin, + sequence_len_offset, + padding_mask, + ) = self._preprocess( + input_ids=input_ids, + position_ids=position_ids, + decoder_input=decoder_input, + inference_context=inference_context, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + ) + # Hash routing consumes batch-major token IDs. Under sequence parallelism, + # shard them with decoder activations so each TP rank hashes its local tokens. + hash_input_ids = None + if self.config.moe_num_hash_layers > 0: + hash_input_ids = input_ids + if ( + self.config.sequence_parallel + and decoder_input is not None + and hash_input_ids is not None + and hash_input_ids.shape[1] != decoder_input.shape[0] + ): + hash_input_ids = ( + tensor_parallel.scatter_to_sequence_parallel_region( + hash_input_ids.transpose(0, 1).contiguous(), group=self.pg_collection.tp + ) + .transpose(0, 1) + .contiguous() + ) + + # Wrap decoder_input so the decoder (HybridStack) can drop its caller's + # reference for early garbage collection during inference. + if in_inference_mode: + decoder_input = WrappedTensor(decoder_input) + + # The following assert will currently fail when running inference. + # Commented out for now. + # TODO (duncan/rwaleffe): (1) confirm that the externally-generated + # attention mask is not needed and is ignored by the model in + # inference mode, (2) reduce the size of the externally-generated + # attention mask to prevent CPU OOM (as we did for training), (3) + # force the attention mask passed to the model in inference mode to + # be None, so this assert will succeed. + # assert attention_mask is None, "The attention mask is ignored and should be set to None" + + packed_seq_params_by_layout = ( + cp_batch.packed_seq_params_by_layout if cp_batch is not None else None + ) + cp_layout_plan = cp_batch.thd_plan if cp_batch is not None else None + + # Run decoder. + backbone_context = ( + torch.no_grad() + if self.config.freeze_base_model_for_mtp and self.training + else nullcontext() + ) + with backbone_context: + decoder_output = self.decoder( + hidden_states=decoder_input, + attention_mask=attention_mask, + inference_context=inference_context, + rotary_pos_emb=rotary_pos_emb, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, + input_ids=hash_input_ids, + ) + if isinstance(decoder_output, tuple): + hidden_states, mhc_multistream = decoder_output + else: + hidden_states = decoder_output + mhc_multistream = None + + return self._postprocess( + hidden_states=hidden_states, + input_ids=input_ids, + position_ids=position_ids, + labels=labels, + rotary_pos_emb=rotary_pos_emb, + rotary_pos_cos=rotary_pos_cos, + rotary_pos_sin=rotary_pos_sin, + mtp_in_postprocess=True, + loss_mask=loss_mask, + mtp_input_mask=mtp_input_mask, + decoder_input=decoder_input, + attention_mask=attention_mask, + padding_mask=padding_mask, + inference_params=inference_params, + packed_seq_params=packed_seq_params, + sequence_len_offset=sequence_len_offset, + runtime_gather_output=runtime_gather_output, + inference_context=inference_context, + compute_mtp_loss=compute_mtp_loss, + cp_batch=cp_batch, + mhc_multistream=mhc_multistream, + mtp_decoder_input=mtp_decoder_input, + ) diff --git a/megatron/core/models/hybrid/layers/utils.py b/megatron/core/models/hybrid/layers/utils.py index ecc4b15a0b9..5d065fb03a7 100644 --- a/megatron/core/models/hybrid/layers/utils.py +++ b/megatron/core/models/hybrid/layers/utils.py @@ -28,6 +28,9 @@ class Symbols: MOE = 'E' PIPE = '|' MTP_SEPARATOR = "/" + # Bracketed groups (e.g. ``[M*E]``) build one nested HybridStack logical layer. + GROUP_START = "[" + GROUP_END = "]" LAYER_CONFIG_MAP = { MAMBA: MambaLayerConfig, GDN: GDNLayerConfig, diff --git a/megatron/core/models/hybrid/model_chunk_schedule_plan.py b/megatron/core/models/hybrid/model_chunk_schedule_plan.py new file mode 100644 index 00000000000..73ab5bc6c8a --- /dev/null +++ b/megatron/core/models/hybrid/model_chunk_schedule_plan.py @@ -0,0 +1,156 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Schedule-plan classes for HybridStack-based decoders. + +These extend the GPT-side ``TransformerLayerSchedulePlan`` / +``TransformerModelChunkSchedulePlan`` with the per-layer ``layer_type`` symbol +that HybridStack assigns to each entry of its ``layer_type_list`` (including +bracketed groups like ``[*-]``). The base classes remain GPT-only; this module +adds the hybrid-specific dispatch into ``build_hybrid_stack_callables`` and +uses ``HybridStackNode`` so the schedule node's free-input policy can diverge +from the GPT default. The pre/post-process nodes from +``core.models.common.utils`` are reused as-is — they already call +``model._preprocess`` / ``model._postprocess`` which work on a HybridModel. +""" + +from contextlib import nullcontext + +from megatron.core.models.common.model_chunk_schedule_plan import ( + TransformerLayerSchedulePlan, + TransformerModelChunkSchedulePlan, +) + + +class HybridStackSchedulePlan(TransformerLayerSchedulePlan): + """Per-layer schedule plan for HybridStack decoders. + + Adds the ``layer_type`` extra-arg propagation; routes through + ``build_hybrid_stack_callables`` when ``layer_type`` is set (i.e. the layer + is a HybridStack entry, possibly a bracketed group); falls back to the GPT + path for plain TransformerLayer / MTP layers when ``layer_type`` is None. + """ + + def __init__(self, layer, event, chunk_state, comp_stream, comm_stream, extra_args=None): + if extra_args is None: + extra_args = {} + self.layer_type = extra_args.get("layer_type", None) + super().__init__(layer, event, chunk_state, comp_stream, comm_stream, extra_args) + + def _build_callable_nodes(self, event, comp_stream, comm_stream, extra_args): + if self.layer_type is None: + return super()._build_callable_nodes(event, comp_stream, comm_stream, extra_args) + + # Hybrid grouped path. Imports are local because hybrid pulls in TE / SSM + # extensions that we don't want to load when only the GPT path is used. + from megatron.core.models.hybrid.fine_grained_callables import ( + HybridStackNode, + build_hybrid_stack_callables, + ) + from megatron.core.pipeline_parallel.utils import NoopScheduleNode + + fwd_callables, bwd_dw_callable_map, is_moe, num_local_experts = ( + build_hybrid_stack_callables(self.layer, layer_type=self.layer_type) + ) + + extra_args["config"] = self.layer.config + extra_args["is_moe"] = is_moe + extra_args["num_local_experts"] = num_local_experts + extra_args["delay_wgrad_compute"] = self.layer.config.delay_wgrad_compute + extra_args["is_mtp"] = False + + def create_node(stream, module, name): + bwd_dw_callables = bwd_dw_callable_map.get(name, None) + node_extra_args = dict(extra_args) + if bwd_dw_callables is None: + node_extra_args["delay_wgrad_compute"] = False + return HybridStackNode( + stream, + event, + self.layer_state, + self.chunk_state, + module, + name=name, + bwd_dw_callables=bwd_dw_callables, + extra_args=node_extra_args, + ) + + ( + pre_dispatch_module, + moe_dispatch_module, + mlp_module, + moe_combine_module, + mtp_post_process_module, + ) = fwd_callables + + self.pre_dispatch_computation = create_node( + comp_stream, pre_dispatch_module, "pre_dispatch_computation" + ) + self.mlp = create_node(comp_stream, mlp_module, "mlp") + if is_moe: + self.moe_dispatch = create_node(comm_stream, moe_dispatch_module, "moe_dispatch") + self.moe_combine = create_node(comm_stream, moe_combine_module, "moe_combine") + else: + self.moe_dispatch = NoopScheduleNode() + self.moe_combine = NoopScheduleNode() + + # HybridStack groups never carry an MTP terminal, so mtp_post_process is + # always a no-op here. + self.mtp_post_process = NoopScheduleNode() + + def get_low_precision_context(self): + """Let hybrid callables manage each physical layer's quantization context.""" + # HybridStack and MambaLayer do not expose the TransformerLayer context + # hook. Their callables enter the appropriate context for each inner layer. + if self.layer_type is not None or not hasattr(self.layer, "get_inner_quantization_context"): + return nullcontext() + return super().get_low_precision_context() + + +class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): + """Model-chunk schedule plan that builds ``HybridStackSchedulePlan`` layer plans. + + Threads HybridStack's ``layer_type_list[layer_idx]`` symbol into each + layer plan's ``extra_args`` so the per-layer plan can dispatch grouped + layers correctly. Ordinary GPT/MTP layers (no ``layer_type_list``) + default to ``layer_type=None`` and follow the GPT path. The pre/post + process nodes inherit from the GPT base class — they already dispatch + on ``model._preprocess`` / ``model._postprocess`` which a HybridModel + implements. + """ + + LAYER_SCHEDULE_PLAN_CLASS = HybridStackSchedulePlan + + def __init__(self, model, *args, **kwargs): + """Initialize the hybrid chunk plan after validating cuda graph support.""" + assert model.config.cuda_graph_impl == "none", ( + "EP A2A overlap with grouped HybridStack patterns (e.g. '[*E]') does not " + "support cuda graphs yet. Set cuda_graph_impl='none' or use an ungrouped pattern." + ) + if getattr(model.config, "moe_num_hash_layers", 0): + raise ValueError("HybridStack EP overlap does not support hash-routed MoE layers.") + if any( + getattr(model.config, option, False) + for option in ("enable_mhc_connections", "wide_residual", "moe_shortcut_connection") + ): + raise ValueError( + "HybridStack EP overlap does not support mHC, wide residuals, or MoE shortcuts." + ) + # The schedule plan calls the layer callables directly and bypasses + # ``HybridStack.forward``, which is where per-layer context-parallel layout + # conversion happens; mixed linear/attention CP layouts are therefore unsupported. + # A group's outer layout is the boundary layout even when its inner + # layers need conversion, so inspect the nested stacks as well. + for module in model.modules(): + cp_layout_manager = getattr(module, "_cp_layout_manager", None) + assert cp_layout_manager is None or not cp_layout_manager.requires_conversion, ( + "EP A2A overlap with HybridStack does not support mixed context-parallel layouts " + "(linear_cp_layout != attention_cp_layout with context_parallel_size > 1)." + ) + super().__init__(model, *args, **kwargs) + + def _extra_args_for_layer(self, module, layer_idx, num_layers): + extra_args = super()._extra_args_for_layer(module, layer_idx, num_layers) + extra_args["layer_type"] = ( + module.layer_type_list[layer_idx] if hasattr(module, "layer_type_list") else None + ) + return extra_args diff --git a/megatron/core/recompute.py b/megatron/core/recompute.py index 5d14314a6db..a8c73785113 100644 --- a/megatron/core/recompute.py +++ b/megatron/core/recompute.py @@ -37,6 +37,8 @@ def checkpointed_forward( cp_layout_state: Optional[ContextParallelLayoutState] = None, packed_sequence_cp_metadata: object | None = None, input_ids: Optional[Tensor] = None, + packed_seq_params_by_layout: Optional[dict] = None, + cp_layout_plan: object | None = None, ) -> Union[Tensor, Tuple[Tensor, Tensor]]: """Forward method with activation checkpointing. @@ -50,6 +52,11 @@ def checkpointed_forward( cp_layout_state (ContextParallelLayoutState, optional): CP layout state for this forward. packed_sequence_cp_metadata (optional): Packed-sequence CP metadata for Mamba layers. input_ids (Tensor, optional): Token IDs forwarded to hash-routed MoE layers. + packed_seq_params_by_layout (dict, optional): Prebuilt packed metadata per CP layout, + forwarded to nested HybridStack group layers so they can build their own layout + state while being checkpointed by the enclosing stack. + cp_layout_plan (optional): Prebuilt packed contiguous/zigzag route, forwarded to nested + HybridStack group layers together with ``packed_seq_params_by_layout``. Returns: If extract_layer_indices is empty: hidden_states tensor @@ -92,9 +99,10 @@ def custom_forward( ) # Keep both residuals in the layer's layout, inside the CP conversions. residual_accumulator = hidden_states + is_hybrid_group = getattr(layer, "is_layer_group_stack", False) # Get appropriate inner quantization context - if use_inner_quantization_context: + if use_inner_quantization_context and not is_hybrid_group: if self.config.fp8: inner_quantization_context = get_fp8_context( self.config, layer.layer_number - 1 @@ -136,6 +144,21 @@ def custom_forward( for k in ("context", "context_mask", "attention_bias"): layer_kwargs.pop(k, None) hidden_states, context = layer(**layer_kwargs) + elif is_hybrid_group: + # Nested HybridStack group: run its physical layers inside this + # checkpoint segment (it must not checkpoint them again) and let it + # build its own CP layout state from the prebuilt per-layout metadata. + for k in ("context", "context_mask", "attention_bias"): + layer_kwargs.pop(k, None) + if input_ids is not None: + layer_kwargs["input_ids"] = input_ids + if packed_seq_params_by_layout is not None or cp_layout_plan is not None: + layer_kwargs["packed_seq_params_by_layout"] = ( + packed_seq_params_by_layout + ) + layer_kwargs["cp_layout_plan"] = cp_layout_plan + hidden_states = layer(**layer_kwargs, _checkpointed_forward_in_parent=True) + context = None else: # MambaLayer (HybridStack `M` slot) for k in ( "context", diff --git a/megatron/core/ssm/mamba_layer.py b/megatron/core/ssm/mamba_layer.py index 7f83d919668..61ee94249a5 100644 --- a/megatron/core/ssm/mamba_layer.py +++ b/megatron/core/ssm/mamba_layer.py @@ -319,6 +319,17 @@ def forward( return hidden_states + def backward_dw(self): + """Compute weight gradients for the layer's linear projections. + + Delegates to the mixer; lets the hybrid EP-overlap schedule plan + register a Mamba pre-layer's wgrad alongside attention/GDN pre-layers + so the schedule node iterates a uniform set of callables. No-op when + the linears in the spec do not support delayed wgrad. + """ + if hasattr(self.mixer, "backward_dw"): + self.mixer.backward_dw() + def sharded_state_dict( self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: diff --git a/megatron/core/ssm/mamba_mixer.py b/megatron/core/ssm/mamba_mixer.py index 26801fc3a26..729e9e04db1 100644 --- a/megatron/core/ssm/mamba_mixer.py +++ b/megatron/core/ssm/mamba_mixer.py @@ -1325,6 +1325,21 @@ def ssm_inference_chunk_size(self) -> int: """ return self.chunk_size + def backward_dw(self): + """Compute weight gradients for the linear layers wrapped by this mixer. + + Mirrors ``GatedDeltaNet.backward_dw``. The selective-scan kernel is a + single autograd function whose wgrad runs in the regular backward pass, + so only the input/output projections need delayed wgrad here. Each + ``backward_dw`` call is a no-op unless the underlying linear is built + from a TE primitive that supports delayed wgrad; if the spec uses + non-TE linears, ``backward_dw`` simply does nothing. + """ + if hasattr(self.in_proj, "backward_dw"): + self.in_proj.backward_dw() + if hasattr(self.out_proj, "backward_dw"): + self.out_proj.backward_dw() + def _get_states_from_cache(self, inference_context, batch_size, *, inference_params=None): """Initializes or retrieves the SSM state tensors from the cache. diff --git a/pretrain_hybrid.py b/pretrain_hybrid.py index da6d4419b25..bd6b6e2c0f3 100644 --- a/pretrain_hybrid.py +++ b/pretrain_hybrid.py @@ -315,13 +315,15 @@ def loss_func( return loss, num_tokens, report -def forward_step(data_iterator, model: HybridModel): +def forward_step(data_iterator, model: HybridModel, return_schedule_plan: bool = False): """Forward training step. Args: data_iterator : Input data iterator model (HybridModel): The Hybrid Model + return_schedule_plan (bool): Whether to return the schedule plan instead of output tensor. """ + args = get_args() timers = get_timers() # Get the batch. @@ -345,16 +347,30 @@ def forward_step(data_iterator, model: HybridModel): timers('batch-generator').stop() with stimer: - output_tensor = model( - tokens, - position_ids, - attention_mask, - labels=labels, - packed_seq_params=packed_seq_params, - loss_mask=loss_mask, - padding_mask=padding_mask, - cp_batch=cp_batch, - ) + if return_schedule_plan: + assert ( + args.overlap_moe_expert_parallel_comm + ), "overlap_moe_expert_parallel_comm must be enabled to return the schedule plan" + output_tensor = model.build_schedule_plan( + tokens, + position_ids, + attention_mask, + labels=labels, + packed_seq_params=packed_seq_params, + loss_mask=loss_mask, + padding_mask=padding_mask, + ) + else: + output_tensor = model( + tokens, + position_ids, + attention_mask, + labels=labels, + packed_seq_params=packed_seq_params, + loss_mask=loss_mask, + padding_mask=padding_mask, + cp_batch=cp_batch, + ) # [ModelOpt]: model is needed to access ModelOpt distillation losses return output_tensor, partial(loss_func, loss_mask, model=model) diff --git a/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py new file mode 100644 index 00000000000..d13def4f44c --- /dev/null +++ b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py @@ -0,0 +1,181 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch + +from megatron.core.context_parallel.layout import ContextParallelLayoutManager +from megatron.core.models.common.model_chunk_schedule_plan import ( + TransformerLayerSchedulePlan, + TransformerModelChunkSchedulePlan, +) +from megatron.core.models.hybrid.hybrid_block import HybridStack +from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.model_chunk_schedule_plan import ( + HybridStackModelChunkSchedulePlan, + HybridStackSchedulePlan, +) +from megatron.core.pipeline_parallel.utils import get_comm_stream, get_comp_stream, set_streams +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.attention_layer_config import AttentionLayerConfig +from megatron.core.transformer.transformer_config import MLATransformerConfig +from tests.unit_tests.a2a_overlap.utils import DummyState +from tests.unit_tests.test_utilities import Utils + + +@pytest.mark.parametrize("location", ["decoder", "group", "mtp"]) +@pytest.mark.parametrize("needs_conversion", [True, False]) +def test_hybrid_schedule_checks_nested_cp_layouts(location, needs_conversion): + """An outer group's boundary layout can hide an inner attention conversion.""" + config = AttentionLayerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + config.attention_cp_layout = "zigzag" if needs_conversion else "contiguous" + cp_group = SimpleNamespace(size=lambda: 2) + + def layout_manager(layer_configs): + return ContextParallelLayoutManager( + layer_layouts=tuple( + HybridStack._get_layer_cp_layout(layer_config, "contiguous") + for layer_config in layer_configs + ), + boundary_layout="contiguous", + sequence_parallel=False, + cp_group=cp_group, + tp_group=None, + tp_cp_group=None, + ) + + model = torch.nn.Module() + model.config = SimpleNamespace(cuda_graph_impl="none") + model.decoder = torch.nn.Module() + model.decoder._cp_layout_manager = layout_manager([(config,)]) + assert not model.decoder._cp_layout_manager.requires_conversion + if location == "group": + model.decoder.group = torch.nn.Module() + target = model.decoder.group + elif location == "mtp": + model.mtp = torch.nn.Module() + target = model.mtp + else: + target = model.decoder + target._cp_layout_manager = layout_manager([config]) + + with patch.object(TransformerModelChunkSchedulePlan, "__init__", return_value=None) as build: + if needs_conversion: + with pytest.raises(AssertionError, match="mixed context-parallel layouts"): + HybridStackModelChunkSchedulePlan(model) + build.assert_not_called() + else: + HybridStackModelChunkSchedulePlan(model) + build.assert_called_once() + + +@pytest.mark.parametrize( + "option", + ["moe_num_hash_layers", "enable_mhc_connections", "wide_residual", "moe_shortcut_connection"], +) +def test_hybrid_schedule_rejects_unsupported_layer_execution(option): + """Direct overlap callables must not bypass token routing or residual wrappers.""" + model = torch.nn.Module() + model.config = SimpleNamespace(cuda_graph_impl="none", **{option: 1}) + with patch.object(TransformerModelChunkSchedulePlan, "__init__", return_value=None) as build: + with pytest.raises(ValueError, match="HybridStack EP overlap does not support"): + HybridStackModelChunkSchedulePlan(model) + build.assert_not_called() + + +@pytest.mark.parametrize("pattern", ["[*E]", "[+E]"]) +def test_grouped_overlap_matches_eager_outputs_and_gradients(pattern): + """Exercise grouped attention, MLA and shared experts through real A2A nodes.""" + if Utils.world_size < 2: + pytest.skip("Expert-parallel overlap requires at least two ranks") + Utils.initialize_model_parallel(expert_model_parallel_size=2) + try: + model_parallel_cuda_manual_seed(123) + config = MLATransformerConfig( + num_layers=2, + hidden_size=256, + num_attention_heads=4, + ffn_hidden_size=256, + bf16=True, + params_dtype=torch.bfloat16, + use_cpu_initialization=True, + hidden_dropout=0.0, + attention_dropout=0.0, + add_bias_linear=False, + num_moe_experts=4, + moe_grouped_gemm=True, + moe_router_topk=2, + moe_router_dtype="fp32", + moe_shared_expert_intermediate_size=256, + expert_model_parallel_size=2, + moe_token_dispatcher_type="alltoall", + overlap_moe_expert_parallel_comm=True, + multi_latent_attention="+" in pattern, + q_lora_rank=64, + kv_lora_rank=64, + qk_head_dim=64, + qk_pos_emb_head_dim=32, + v_head_dim=64, + ) + block = HybridStack( + config, + hybrid_stack_spec.submodules, + layer_config_list=validate_segment_layers(pattern, config), + pg_collection=ProcessGroupCollection.use_mpu_process_groups(), + ).cuda() + inputs = [torch.randn(16, 1, 256, device="cuda", dtype=torch.bfloat16) for _ in range(3)] + references = [] + for hidden_states in inputs: + output = block(hidden_states.clone().requires_grad_(), attention_mask=None) + references.append(output.detach().clone()) + output.backward(torch.ones_like(output)) + reference_grads = { + name: param.grad.detach().clone() + for name, param in block.named_parameters() + if param.grad is not None + } + block.zero_grad(set_to_none=True) + + set_streams() + plans = [] + for _ in inputs: + state = DummyState() + state.model = SimpleNamespace(decoder=block) + plans.append( + HybridStackSchedulePlan( + block.layers[0], + torch.cuda.Event(), + state, + get_comp_stream, + get_comm_stream, + extra_args={"layer_type": block.layer_type_list[0], "is_last_layer": True}, + ) + ) + + outputs = [] + previous = None + for index, plan in enumerate(plans): + output, _ = TransformerLayerSchedulePlan.run( + plan, + previous, + f_input=inputs[index].clone().requires_grad_(), + b_grad=None if previous is None else torch.ones_like(outputs[-1]), + ) + torch.cuda.synchronize() + outputs.append(output.detach().clone()) + previous = plan + TransformerLayerSchedulePlan.run(None, previous, b_grad=torch.ones_like(outputs[-1])) + torch.cuda.synchronize() + + for output, reference in zip(outputs, references): + torch.testing.assert_close(output, reference, rtol=1e-2, atol=1e-2) + for name, param in block.named_parameters(): + if name in reference_grads: + assert param.grad is not None, name + torch.testing.assert_close(param.grad, reference_grads[name], rtol=2e-2, atol=2e-2) + finally: + Utils.destroy_model_parallel() diff --git a/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py b/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py index a8ece4846a0..19e1032f28f 100644 --- a/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py +++ b/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py @@ -4,10 +4,12 @@ from types import SimpleNamespace from unittest.mock import Mock, patch +import pytest import torch from megatron.core.enums import Fp8Recipe from megatron.core.models.common.model_chunk_schedule_plan import TransformerLayerSchedulePlan +from megatron.core.models.hybrid.model_chunk_schedule_plan import HybridStackSchedulePlan from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer from megatron.core.transformer.transformer_layer import TransformerLayer @@ -28,6 +30,50 @@ def test_schedule_uses_layer_quantization_context(): layer.get_inner_quantization_context.assert_called_once_with() +@pytest.mark.parametrize("layer_type", [("*", "E"), "M", None]) +def test_hybrid_schedule_runs_without_a_transformer_context_hook(layer_type): + """The common scheduler must dispatch to the hybrid quantization hook.""" + plan = HybridStackSchedulePlan.__new__(HybridStackSchedulePlan) + plan.layer = SimpleNamespace() + plan.layer_type = layer_type + visited = [] + for name in ( + "pre_dispatch_computation", + "moe_dispatch", + "mlp", + "moe_combine", + "mtp_post_process", + ): + setattr( + plan, + name, + SimpleNamespace(forward=lambda value, name=name: visited.append(name) or value), + ) + + value = torch.ones(1) + output, _ = TransformerLayerSchedulePlan.run(plan, None, f_input=value) + + assert output is value + assert visited == [ + "pre_dispatch_computation", + "moe_dispatch", + "mlp", + "moe_combine", + "mtp_post_process", + ] + + +def test_hybrid_schedule_preserves_plain_layer_quantization_context(): + expected_context = nullcontext() + layer = SimpleNamespace(get_inner_quantization_context=Mock(return_value=expected_context)) + plan = HybridStackSchedulePlan.__new__(HybridStackSchedulePlan) + plan.layer = layer + plan.layer_type = None + + assert plan.get_low_precision_context() is expected_context + layer.get_inner_quantization_context.assert_called_once_with() + + def test_transformer_layer_uses_fp4_context(): config = SimpleNamespace(fp8=None, fp8_recipe=Fp8Recipe.delayed, fp4="e2m1") layer = TransformerLayer.__new__(TransformerLayer) diff --git a/tests/unit_tests/models/test_hybrid_fine_grained_callables.py b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py new file mode 100644 index 00000000000..a909b765ade --- /dev/null +++ b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py @@ -0,0 +1,295 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +from contextlib import nullcontext +from functools import partial +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +import megatron.core.models.common.utils as node_utils +import megatron.core.models.hybrid.fine_grained_callables as hybrid_callables +import megatron.core.pipeline_parallel.utils as schedule_utils +from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols +from megatron.core.transformer.moe.moe_layer import MoELayer + + +@pytest.fixture +def cpu_streams(monkeypatch): + """Exercise real schedule-node autograd while replacing only CUDA stream operations.""" + streams = SimpleNamespace(comp=Mock(), comm=Mock()) + monkeypatch.setattr(torch.cuda, "stream", lambda stream: nullcontext()) + monkeypatch.setattr(torch.cuda, "current_stream", lambda: streams.comp) + monkeypatch.setattr(torch.Tensor, "record_stream", lambda self, stream: None) + monkeypatch.setattr(schedule_utils, "get_comm_stream", lambda: streams.comm) + monkeypatch.setattr(hybrid_callables, "get_comm_stream", lambda: streams.comm) + for module in (node_utils, schedule_utils): + monkeypatch.setattr(module, "nvtx_range_push", lambda *args, **kwargs: None) + monkeypatch.setattr(module, "nvtx_range_pop", lambda *args, **kwargs: None) + return streams + + +def _config(**kwargs): + values = dict( + fp8=None, + fp4=None, + fp32_residual_connection=False, + moe_token_dispatcher_type="flex", + moe_flex_dispatcher_backend="ncclep", + moe_ncclep_zero_copy=False, + moe_latent_size=None, + overlap_moe_expert_parallel_comm=True, + cuda_graph_modules=[], + bias_dropout_fusion=False, + ) + values.update(kwargs) + return SimpleNamespace(**values) + + +def _moe_layer(config, *, shared=False): + mlp = SimpleNamespace( + config=config, + num_local_experts=1, + use_shared_expert=shared, + shared_expert_overlap=False, + experts=torch.nn.Linear(2, 2, bias=False), + shared_experts=torch.nn.Linear(2, 2, bias=False) if shared else None, + backward_dw=Mock(), + ) + return SimpleNamespace(config=config, mlp=mlp, recompute_pre_mlp_layernorm=False) + + +def _chunk_state(): + return SimpleNamespace( + model=SimpleNamespace(decoder=SimpleNamespace(final_norm=None)), + padding_mask=None, + attention_mask=None, + rotary_pos_emb=None, + rotary_pos_cos=None, + rotary_pos_sin=None, + packed_seq_params=None, + sequence_len_offset=None, + ) + + +def _node(layer, layer_state, function, name): + return hybrid_callables.HybridStackNode( + stream=object(), + event=Mock(), + layer_state=layer_state, + chunk_state=_chunk_state(), + submodule=function, + name=name, + extra_args={"config": layer.config, "is_moe": True, "num_local_experts": 1}, + ) + + +@pytest.mark.parametrize("shared", [False, True]) +@pytest.mark.parametrize("latent", [False, True]) +def test_delayed_wgrad_hooks_belong_to_the_slot_that_computes_them(cpu_streams, shared, latent): + """Stale gradients must not cause the routed slot to accumulate shared-expert gradients.""" + events = [] + cpu_streams.comp.wait_stream.side_effect = lambda stream: events.append("wait_comm") + + class Expert(torch.nn.Module): + def __init__(self, name): + super().__init__() + self.name = name + self.weight = torch.nn.Parameter(torch.ones(1)) + # A prior microbatch can leave a gradient before this slot runs. + self.weight.grad = torch.ones_like(self.weight) + self.weight.post_wgrad_grad_acc_hook = lambda: events.append(f"{name}:hook") + + def backward_dw(self): + events.append(f"{self.name}:wgrad") + + layer = _moe_layer(_config(moe_latent_size=2 if latent else None), shared=shared) + mlp = layer.mlp + mlp.experts = Expert("routed") + mlp.shared_experts = Expert("shared") if shared else None + if latent: + mlp.fc1_latent_proj = Expert("down") + mlp.fc2_latent_proj = Expert("up") + mlp.backward_dw = partial(MoELayer.backward_dw, mlp) + _, backward_dw, _, _ = hybrid_callables.build_hybrid_stack_callables(layer, Symbols.MOE) + + def run_wgrad(modules): + node = SimpleNamespace( + delay_wgrad_compute=True, + stream=object(), + name="test slot", + bwd_dw_callables=modules, + post_wgrad_grad_acc_hooks=None, + is_layer_first_node=False, + ) + node_utils.TransformerLayerNode.backward_dw(node) + assert node.bwd_dw_callables is None + + run_wgrad([backward_dw["mlp"]]) + routed_names = ["routed", "up"] if latent else ["routed"] + assert events == ( + [f"{name}:wgrad" for name in routed_names] + + (["wait_comm"] if latent else []) + + [f"{name}:hook" for name in routed_names] + ) + if latent: + cpu_streams.comp.wait_stream.assert_called_once_with(cpu_streams.comm) + else: + cpu_streams.comp.wait_stream.assert_not_called() + + events.clear() + pre_names = (["shared"] if shared else []) + (["down"] if latent else []) + if pre_names: + run_wgrad(backward_dw["pre_dispatch_computation"]) + else: + assert "pre_dispatch_computation" not in backward_dw + assert events == [f"{name}:wgrad" for name in pre_names] + [ + f"{name}:hook" for name in pre_names + ] + + +@pytest.mark.parametrize( + "symbol", [Symbols.ATTENTION, Symbols.DS_ATTENTION, Symbols.MLA, Symbols.GDN] +) +def test_attention_half_layer_forward_and_wgrad_are_scheduled(symbol): + """All attention symbols, including MLA '+', preserve their forward and backward work.""" + backward_dw_wrapper = Mock() + layer = SimpleNamespace( + config=_config(), + _forward_attention=Mock( + side_effect=lambda hidden_states, **kwargs: (hidden_states.sin(), None) + ), + backward_dw_wrapper=backward_dw_wrapper, + init_backward_dw_wrapper=Mock(), + ) + forward, backward_dw, is_moe, _ = hybrid_callables.build_hybrid_stack_callables(layer, symbol) + node = SimpleNamespace(chunk_state=_chunk_state(), is_mtp=False, is_last_layer=False) + hidden_states = torch.tensor([0.2, 0.4], requires_grad=True) + output = forward[0](node, hidden_states) + output.sum().backward() + + torch.testing.assert_close(output, hidden_states.sin()) + torch.testing.assert_close(hidden_states.grad, hidden_states.cos()) + layer.init_backward_dw_wrapper.assert_called_once_with() + assert backward_dw["pre_dispatch_computation"] == [backward_dw_wrapper] + assert not is_moe + + +@pytest.mark.parametrize("zero_copy", [False, True]) +def test_ncclep_probabilities_do_not_reconnect_schedule_graphs(cpu_streams, zero_copy): + """Backward across all three real schedule nodes must match the unsplit computation.""" + layer = _moe_layer(_config(moe_ncclep_zero_copy=zero_copy)) + manager = SimpleNamespace( + token_probs=None, + dispatched_probs=None, + get_number_of_tokens_per_expert=Mock(return_value=torch.tensor([2])), + _zc_bwd_token_buf=torch.empty(2), + ) + dispatch_grad_ptrs = [] + + class Dispatch(torch.autograd.Function): + @staticmethod + def forward(ctx, tokens, probs): + return tokens * 3, probs * 4 + + @staticmethod + def backward(ctx, token_grad, prob_grad): + dispatch_grad_ptrs.append(token_grad.data_ptr()) + return token_grad * 3, prob_grad * 4 + + def dispatch(tokens, probs): + # Flex dispatch consumes the saved manager state, not its explicit probs argument. + output, manager.dispatched_probs = Dispatch.apply(tokens, manager.token_probs) + return output, manager.dispatched_probs + + layer.mlp.token_dispatcher = SimpleNamespace(_comm_manager=manager) + layer.mlp.dispatch = dispatch + layer.mlp.routed_experts_compute = lambda tokens, probs: ( + tokens * manager.dispatched_probs, + None, + ) + forward, _, _, _ = hybrid_callables.build_hybrid_stack_callables(layer, Symbols.MOE) + + def preprocess(node, hidden_states): + manager.token_probs = hidden_states.square() + return hidden_states * 2, manager.token_probs + + state = SimpleNamespace() + pre_node = _node(layer, state, preprocess, "pre_dispatch_computation") + dispatch_node = _node(layer, state, forward[1], "moe_dispatch") + expert_node = _node(layer, state, forward[2], "mlp") + hidden_states = torch.tensor([0.2, 0.4], requires_grad=True) + output = expert_node.forward(dispatch_node.forward(pre_node.forward(hidden_states))) + grad = expert_node.backward(torch.ones_like(output)) + grad = dispatch_node.backward(grad) + grad = pre_node.backward(grad) + + torch.testing.assert_close(output, 24 * hidden_states.pow(3)) + torch.testing.assert_close(grad, 72 * hidden_states.square()) + assert state.tokens_per_expert is manager.get_number_of_tokens_per_expert.return_value + if zero_copy: + assert dispatch_grad_ptrs == [manager._zc_bwd_token_buf.data_ptr()] + + +@pytest.mark.parametrize("recompute", [False, True]) +def test_norm_offload_uses_its_microbatch_manager_after_bda(cpu_streams, recompute): + """Two in-flight microbatches keep distinct offload managers and single recompute hooks.""" + layer = _moe_layer(_config()) + layer.recompute_pre_mlp_layernorm = recompute + layer.offload_mlp_norm = True + layer.training = True + layer.hidden_dropout = 0.0 + layer.bias_dropout_add_exec_handler = nullcontext + layer.mlp_bda = lambda *args: lambda output, residual, dropout: output[0] + residual + layer.mlp.shared_experts_compute = lambda hidden: None + layer.mlp.route = lambda hidden, mask: (hidden, None) + layer.mlp.preprocess = lambda hidden, probs, routing: (hidden, probs) + layer.mlp.routed_experts_compute = lambda hidden, probs: (hidden * 3, None) + layer.mlp.combine = lambda output: output + layer.mlp.postprocess = lambda output, shared: output + layer.mlp.token_dispatcher = SimpleNamespace( + _comm_manager=SimpleNamespace(get_number_of_tokens_per_expert=lambda: torch.tensor([2])) + ) + managers, checkpoints = [], [] + + def normalize(hidden_states): + manager = SimpleNamespace(group_offload=Mock(side_effect=lambda output, **kwargs: output)) + layer.mlp_norm_manager = manager + managers.append(manager) + layer.pre_mlp_norm_checkpoint = Mock() + checkpoints.append(layer.pre_mlp_norm_checkpoint) + return hidden_states * 2 + + layer._forward_pre_mlp_layernorm = normalize + microbatches = [] + for value in (1.0, 2.0): + hidden_states = torch.full((2,), value, requires_grad=True) + node = SimpleNamespace( + layer_state=SimpleNamespace(), + chunk_state=_chunk_state(), + detach=lambda tensor: tensor.detach(), + is_mtp=False, + is_last_layer=False, + ) + tokens, probs = hybrid_callables._run_moe_preprocess(layer, node, hidden_states) + node.layer_state.dispatched_probs = probs + expert_output = hybrid_callables._run_moe_experts(layer, node, tokens) + microbatches.append((node, hidden_states, expert_output)) + assert layer.mlp_norm_manager is None + + for index, (node, hidden_states, expert_output) in enumerate(microbatches): + residual = node.layer_state.residual + output = hybrid_callables._run_moe_combine(layer, node, expert_output) + torch.testing.assert_close(output, hidden_states * 7) + managers[index].group_offload.assert_called_once_with( + output, forced_released_tensors=[residual] + ) + assert node.layer_state.mlp_norm_manager is None + assert node.layer_state.residual is None + if recompute: + checkpoints[index].discard_output_and_register_recompute.assert_called_once_with( + expert_output + ) + else: + checkpoints[index].discard_output_and_register_recompute.assert_not_called() diff --git a/tests/unit_tests/models/test_hybrid_hash_routing.py b/tests/unit_tests/models/test_hybrid_hash_routing.py index 73a6e9e8d24..e46d1028250 100644 --- a/tests/unit_tests/models/test_hybrid_hash_routing.py +++ b/tests/unit_tests/models/test_hybrid_hash_routing.py @@ -7,6 +7,7 @@ import torch from megatron.core import recompute as recompute_module +from megatron.core.inference.utils import InferenceMode from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers @@ -25,6 +26,7 @@ ) from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayer +from megatron.core.utils import WrappedTensor class RecordingTransformerLayer(TransformerLayer): @@ -123,7 +125,8 @@ def __init__(self): def __call__(self, **kwargs): self.kwargs = kwargs - return kwargs['hidden_states'] + hidden_states = kwargs['hidden_states'] + return hidden_states.unwrap() if isinstance(hidden_states, WrappedTensor) else hidden_states class RecordingMoE(torch.nn.Module): @@ -316,6 +319,8 @@ def test_hybrid_model_passes_ids_to_decoder_only_for_hash_routing( mtp_process=False, vocab_size=128, ) + model._preprocess = HybridModel._preprocess.__get__(model) + model._postprocess = HybridModel._postprocess.__get__(model) input_ids = torch.arange(8).reshape(2, 4) hidden_states = torch.randn(4, 2, 8) @@ -334,7 +339,10 @@ def test_hybrid_model_passes_ids_to_decoder_only_for_hash_routing( @pytest.mark.parametrize("pre_process", [True, False]) -def test_hybrid_model_sequence_shards_hash_ids_with_decoder_input(monkeypatch, pre_process): +@pytest.mark.parametrize("inference", [False, True]) +def test_hybrid_model_sequence_shards_hash_ids_with_decoder_input( + monkeypatch, pre_process, inference +): decoder = RecordingDecoder() tp_group = object() scattered = [] @@ -367,6 +375,9 @@ def fake_scatter(tensor, group): vocab_size=128, pg_collection=SimpleNamespace(tp=tp_group), ) + model._preprocess = HybridModel._preprocess.__get__(model) + model._postprocess = HybridModel._postprocess.__get__(model) + monkeypatch.setattr(InferenceMode, 'is_active', lambda: inference) input_ids = torch.arange(8).reshape(2, 4) padding_mask = torch.tensor([[False, True, False, True], [True, False, True, False]]) hidden_states = torch.randn(2, 2, 8) @@ -379,6 +390,7 @@ def fake_scatter(tensor, group): attention_mask=None, decoder_input=hidden_states if pre_process else None, padding_mask=padding_mask, + runtime_gather_output=inference, ) assert len(scattered) == (2 if pre_process else 1) @@ -479,6 +491,22 @@ def test_hash_moe_threshold_counts_only_moe_positions(): assert _get_hash_moe_layer_threshold("-E-E-E-E", 3) == 6 +def test_hash_moe_threshold_counts_group_members_as_physical_layers(): + assert _get_hash_moe_layer_threshold("[M*E]|[M*E]", 1) == 3 + assert _get_hash_moe_layer_threshold("[M*E]|[M*E]", 2) == 6 + + +def test_hash_moe_pipeline_placement_checks_grouped_moe_members(): + grouped_layers = [(LayerSymbols.MAMBA, LayerSymbols.ATTENTION, LayerSymbols.MOE)] + with pytest.raises(ValueError, match=r"hash MoE layer\(s\) \[6\]"): + _validate_hash_moe_pipeline_placement( + grouped_layers, layer_offset=3, hash_moe_layer_threshold=6, pre_process=False + ) + _validate_hash_moe_pipeline_placement( + grouped_layers, layer_offset=3, hash_moe_layer_threshold=3, pre_process=False + ) + + def test_hash_moe_threshold_rejects_count_larger_than_pattern(): with pytest.raises(ValueError, match="exceeds the 2 MoE layers"): _get_hash_moe_layer_threshold("-E-E", 3) diff --git a/tests/unit_tests/models/test_hybrid_model.py b/tests/unit_tests/models/test_hybrid_model.py index a6150336b37..5f6dc0fccc5 100644 --- a/tests/unit_tests/models/test_hybrid_model.py +++ b/tests/unit_tests/models/test_hybrid_model.py @@ -33,7 +33,7 @@ from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import MLATransformerConfig, TransformerConfig from megatron.core.transformer.attention_layer_config import AttentionLayerConfig -from megatron.core.transformer.enums import AttnBackend +from megatron.core.transformer.enums import AttnBackend, InferenceCudaGraphScope from megatron.core.transformer.module import Float16Module from megatron.core.utils import divide, is_fa_min_version, is_torch_min_version from tests.unit_tests.test_utilities import Utils @@ -57,6 +57,184 @@ def _mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tenso return x * scale +def _make_postprocess_stub(config, training): + """Build the minimal model surface needed by ``HybridModel._postprocess``.""" + + def output_layer(hidden_states, weight=None, runtime_gather_output=None): + return hidden_states, None + + output_layer.sequence_parallel = False + pg_collection = SimpleNamespace(cp=object(), tp=object(), dp_cp=object()) + return SimpleNamespace( + config=config, + training=training, + post_process=True, + mtp_process=True, + share_embeddings_and_output_weights=False, + output_layer=output_layer, + pg_collection=pg_collection, + tp_group=pg_collection.tp, + _scale_logits=lambda logits: logits, + compute_language_model_loss=lambda labels, logits: logits, + ) + + +@pytest.mark.parametrize( + "cuda_graph_scope", [InferenceCudaGraphScope.none, InferenceCudaGraphScope.block] +) +def test_hybrid_postprocess_caches_spec_decode_hidden_states(cuda_graph_scope): + """Speculative decoding publishes decoder states through the inference context.""" + hidden_states = torch.randn(3, 2, 8) + context_buffer = ( + torch.full((5, 2, 8), -1.0) if cuda_graph_scope == InferenceCudaGraphScope.block else None + ) + inference_context = SimpleNamespace( + config=SimpleNamespace(materialize_only_last_token_logits=False), + num_speculative_tokens=1, + mtp_decoder_hidden_states=context_buffer, + is_dynamic_batching=lambda: True, + is_static_batching=lambda: False, + ) + model = _make_postprocess_stub( + SimpleNamespace( + mtp_num_layers=1, inference_cuda_graph_scope=cuda_graph_scope, use_mup=False + ), + training=False, + ) + + with InferenceMode.active(): + output = HybridModel._postprocess( + model, + hidden_states=hidden_states, + input_ids=torch.zeros(2, 3, dtype=torch.long), + position_ids=torch.zeros(2, 3, dtype=torch.long), + labels=None, + rotary_pos_emb=None, + mtp_in_postprocess=False, + runtime_gather_output=True, + inference_context=inference_context, + ) + + torch.testing.assert_close(output, hidden_states.transpose(0, 1), rtol=0, atol=0) + if cuda_graph_scope == InferenceCudaGraphScope.block: + assert inference_context.mtp_decoder_hidden_states is context_buffer + torch.testing.assert_close(context_buffer[:3], hidden_states, rtol=0, atol=0) + assert torch.all(context_buffer[3:] == -1) + else: + assert inference_context.mtp_decoder_hidden_states is hidden_states + assert not hasattr(model, "_decoder_hidden_states_cache") + + +def test_hybrid_postprocess_does_not_cache_regular_inference_hidden_states(): + """Regular inference must not make the controller run serial MTP decoding.""" + hidden_states = torch.randn(3, 2, 8) + inference_context = SimpleNamespace( + config=SimpleNamespace(materialize_only_last_token_logits=False), + num_speculative_tokens=0, + mtp_decoder_hidden_states=None, + is_dynamic_batching=lambda: True, + is_static_batching=lambda: False, + ) + model = _make_postprocess_stub( + SimpleNamespace( + mtp_num_layers=1, inference_cuda_graph_scope=InferenceCudaGraphScope.none, use_mup=False + ), + training=False, + ) + + with InferenceMode.active(): + HybridModel._postprocess( + model, + hidden_states=hidden_states, + input_ids=torch.zeros(2, 3, dtype=torch.long), + position_ids=torch.zeros(2, 3, dtype=torch.long), + labels=None, + rotary_pos_emb=None, + mtp_in_postprocess=False, + runtime_gather_output=True, + inference_context=inference_context, + ) + + assert inference_context.mtp_decoder_hidden_states is None + assert not hasattr(model, "_decoder_hidden_states_cache") + + +def test_hybrid_postprocess_forwards_rl_mtp_inputs(monkeypatch): + """RL MTP loss receives token IDs and the TP group needed to derive labels.""" + captured_kwargs = {} + + def fake_process_mtp_loss(**kwargs): + captured_kwargs.update(kwargs) + return kwargs["hidden_states"][:2] + + monkeypatch.setattr( + "megatron.core.models.hybrid.hybrid_model.process_mtp_loss", fake_process_mtp_loss + ) + model = _make_postprocess_stub( + SimpleNamespace( + mtp_num_layers=1, inference_cuda_graph_scope=InferenceCudaGraphScope.none, use_mup=False + ), + training=True, + ) + input_ids = torch.arange(4, dtype=torch.long).reshape(1, 4) + + HybridModel._postprocess( + model, + hidden_states=torch.randn(4, 1, 8), + input_ids=input_ids, + position_ids=torch.arange(4, dtype=torch.long).reshape(1, 4), + labels=None, + rotary_pos_emb=None, + mtp_in_postprocess=False, + runtime_gather_output=False, + inference_context=None, + ) + + assert captured_kwargs["input_ids"] is input_ids + assert captured_kwargs["tp_group"] is model.tp_group + + +def test_hybrid_postprocess_uses_output_processor_hook(): + """A caller-supplied output processor replaces the default logits / loss path.""" + captured_kwargs = {} + sentinel = torch.randn(2, 4, 8) + + def output_processor(**kwargs): + captured_kwargs.update(kwargs) + return sentinel + + model = _make_postprocess_stub( + SimpleNamespace( + mtp_num_layers=None, inference_cuda_graph_scope=InferenceCudaGraphScope.none + ), + training=True, + ) + hidden_states = torch.randn(4, 2, 8) + labels = torch.zeros(2, 4, dtype=torch.long) + context = object() + + output = HybridModel._postprocess( + model, + hidden_states=hidden_states, + input_ids=torch.zeros(2, 4, dtype=torch.long), + position_ids=torch.zeros(2, 4, dtype=torch.long), + labels=labels, + rotary_pos_emb=None, + mtp_in_postprocess=False, + runtime_gather_output=False, + inference_context=None, + output_processor=output_processor, + output_processor_context=context, + ) + + assert output is sentinel + assert captured_kwargs["hidden_states"] is hidden_states + assert captured_kwargs["labels"] is labels + assert captured_kwargs["context"] is context + assert captured_kwargs["output_layer"] is model.output_layer + assert captured_kwargs["compute_language_model_loss"] is model.compute_language_model_loss + + def test_hybrid_logging_process_groups_are_paired(): tp_group = object() dp_cp_group = object() @@ -636,6 +814,48 @@ def test_save_load(self, tmp_path): self.model.load_state_dict(torch.load(path)) + def test_grouped_sharded_state_dict_uses_transformer_checkpoint_keys(self): + """Grouped HybridModel checkpoints should be load-compatible with GPTModel keys.""" + model_config = TransformerConfig( + num_layers=2, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + model = HybridModel( + config=model_config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern="[*-]", + ) + + sharded_state_dict = model.sharded_state_dict() + sharded_keys = {value.key for value in sharded_state_dict.values() if hasattr(value, "key")} + + assert "decoder.layers.0.self_attention.linear_qkv.weight" in sharded_keys + assert "decoder.layers.0.mlp.linear_fc1.weight" in sharded_keys + assert "decoder.layers.1.mlp.linear_fc1.weight" not in sharded_keys + assert "decoder.final_layernorm.weight" in sharded_keys + assert "decoder.final_norm.weight" not in sharded_keys + assert "output_layer._extra_state" not in sharded_state_dict + + def test_ungrouped_sharded_state_dict_keeps_hybrid_final_norm_key(self): + """Non-grouped patterns keep ``final_norm`` so older hybrid checkpoints load.""" + model_config = TransformerConfig( + num_layers=2, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + model = HybridModel( + config=model_config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern="*-", + ) + + sharded_state_dict = model.sharded_state_dict() + sharded_keys = {value.key for value in sharded_state_dict.values() if hasattr(value, "key")} + + assert "decoder.final_norm.weight" in sharded_keys + assert "decoder.final_layernorm.weight" not in sharded_keys + def test_layer_numbers(self): """ The layer numbers should start at one (for the embedding # layer) and go up diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index 98c262b028a..9e5279f8959 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -9,6 +9,7 @@ import megatron.core.models.hybrid.hybrid_block as hybrid_block_module import megatron.core.transformer.utils as transformer_utils from megatron.core.extensions.transformer_engine import TEDotProductAttention +from megatron.core.models.hybrid.fine_grained_callables import build_hybrid_stack_callables from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols, validate_segment_layers from megatron.core.models.hybrid.hybrid_layer_specs import ( @@ -51,6 +52,11 @@ def _make_pg_collection(): return SimpleNamespace(pp=None, tp=None, cp=SimpleNamespace(size=lambda: 1), tp_cp=None) +def _num_physical_layers(layer_pattern: str) -> int: + """Number of physical layers in a segment pattern (bracketed groups flattened).""" + return len(layer_pattern.replace(Symbols.GROUP_START, '').replace(Symbols.GROUP_END, '')) + + def test_wide_residual_spec_preserves_unmodified_stack_submodules(): """Wide specialization changes layer classes without dropping stack-level specs.""" @@ -192,6 +198,69 @@ def __init__(self, **kwargs): assert layout_manager_kwargs["boundary_layout"] == "contiguous" +@pytest.mark.parametrize("pp_layer_offset", [0, 5]) +@pytest.mark.parametrize("layer_pattern", ["[*-][*-]", "M[*]", "[M*E][M*E]", "[M+E][M+E]"]) +def test_group_inference_offsets_match_flat_layers(monkeypatch, layer_pattern, pp_layer_offset): + """Grouping keeps physical layer numbers and pipeline-local cache indices unchanged.""" + + class BuiltLayer(torch.nn.Module): + + def __init__(self, config, layer_number): + super().__init__() + self.config = config + self.layer_number = layer_number + + build_calls = [] + + def fake_build_module(module_spec, **kwargs): + build_calls.append(kwargs) + return BuiltLayer(kwargs["config"], kwargs["layer_number"]) + + monkeypatch.setattr(hybrid_block_module, "build_module", fake_build_module) + flat_pattern = layer_pattern.replace('[', '').replace(']', '') + config = MLATransformerConfig( + num_layers=pp_layer_offset + len(flat_pattern), hidden_size=64, num_attention_heads=4 + ) + offsets = [] + for pattern in (flat_pattern, layer_pattern): + build_calls.clear() + HybridStack( + config, + hybrid_stack_spec.submodules, + layer_config_list=validate_segment_layers(pattern, config), + pp_layer_offset=pp_layer_offset, + post_process=False, + pg_collection=_make_pg_collection(), + ) + assert [call["layer_number"] for call in build_calls] == list( + range(pp_layer_offset + 1, pp_layer_offset + len(flat_pattern) + 1) + ) + offsets.append( + [ + call["layer_number"] - call["pp_layer_offset"] + for call in build_calls + if "pp_layer_offset" in call + ] + ) + + assert offsets[1] == offsets[0] + assert len(set(offsets[1])) == len(offsets[1]) + + +@pytest.mark.parametrize("group", ["MM", "--", "**", "G*", "D+", "-E"]) +def test_explicit_group_configs_reject_checkpoint_namespace_collisions(group): + """Direct config tuples enforce the same checkpoint constraints as parsed patterns.""" + config = MLATransformerConfig(num_layers=2, hidden_size=64, num_attention_heads=4) + group_configs = tuple(layer_utils.create_layer_config(config, symbol) for symbol in group) + with pytest.raises(ValueError, match="multiple layers in checkpoint namespace"): + HybridStack( + config, + hybrid_stack_spec.submodules, + layer_config_list=[group_configs], + pg_collection=_make_pg_collection(), + ) + + def test_hybrid_stack_rejects_layer_config_subclasses(monkeypatch): """Layer config subclasses must be registered as distinct layer types.""" @@ -734,8 +803,10 @@ def setup_method(self, method): Utils.initialize_model_parallel(1, 1) model_parallel_cuda_manual_seed(123) - def get_pg_collection(self): - return ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'pp', 'cp']) + def get_pg_collection(self, required_pgs=None): + if required_pgs is None: + required_pgs = ['tp', 'pp', 'cp'] + return ProcessGroupCollection.use_mpu_process_groups(required_pgs=required_pgs) def test_hybrid_mtp_rejects_expert_parallel_overlap_before_build(self, monkeypatch): """Reject overlap before constructing any HybridModel submodule.""" @@ -765,7 +836,7 @@ def get_hybrid_block(self, layer_pattern, *, stack_spec=hybrid_stack_spec, **con hidden_size=256, # The Mamba layer places several constraints on this # Need to specify num_attention_heads and num_layers or TransformerConfig # will generate errors. - num_layers=len(layer_pattern), + num_layers=_num_physical_layers(layer_pattern), num_attention_heads=4, use_cpu_initialization=True, **config_kwargs, @@ -858,6 +929,60 @@ def get_gdn2_hybrid_block( **config_kwargs, ) + def get_attention_mlp_block(self, layer_pattern): + transformer_config = TransformerConfig( + hidden_size=256, + num_layers=_num_physical_layers(layer_pattern), + num_attention_heads=4, + hidden_dropout=0.0, + attention_dropout=0.0, + use_cpu_initialization=True, + ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) + return HybridStack( + transformer_config, + hybrid_stack_spec.submodules, + layer_config_list=layer_config_list, + pp_layer_offset=0, + pg_collection=self.get_pg_collection(), + ) + + def get_attention_moe_block(self, layer_pattern): + transformer_config = TransformerConfig( + hidden_size=256, + num_layers=_num_physical_layers(layer_pattern), + num_attention_heads=4, + ffn_hidden_size=256, + num_moe_experts=8, + expert_model_parallel_size=1, + moe_router_topk=2, + moe_grouped_gemm=True, + moe_token_dispatcher_type="alltoall", + hidden_dropout=0.0, + attention_dropout=0.0, + use_cpu_initialization=True, + ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) + return HybridStack( + transformer_config, + hybrid_stack_spec.submodules, + layer_config_list=layer_config_list, + pp_layer_offset=0, + pg_collection=self.get_pg_collection( + required_pgs=[ + 'tp', + 'pp', + 'cp', + 'tp_cp', + 'tp_dp_cp', + 'ep', + 'expt_tp', + 'tp_ep', + 'expt_dp', + ] + ), + ) + def teardown_method(self, method): Utils.destroy_model_parallel() @@ -909,6 +1034,7 @@ def _run_forward(self, block, sequence_length=32, micro_batch_size=2): Symbols.MLP * 5, Symbols.ATTENTION + Symbols.MLP + Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP, Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP, + "[*-]", ], ) def test_recompute(self, recompute_kwargs: dict, layer_pattern: str): @@ -1344,6 +1470,158 @@ def test_shortcut_pair_supports_a_residual_returning_pre_mlp_norm(self): assert hidden_states.grad is not None assert torch.isfinite(hidden_states.grad).all() + def test_group_layer_type_builds_nested_hybrid_stack(self): + """Bracketed groups build an inner HybridStack with physical layer numbering.""" + layer_pattern = "M[M*]-" + transformer_config = TransformerConfig( + hidden_size=256, + num_layers=_num_physical_layers(layer_pattern), + num_attention_heads=4, + use_cpu_initialization=True, + ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) + block = HybridStack( + transformer_config, + hybrid_stack_spec.submodules, + layer_config_list=layer_config_list, + pp_layer_offset=0, + pg_collection=self.get_pg_collection(), + ) + assert block.layer_type_list == ['M', ('M', '*'), '-'] + assert isinstance(block.layers[0], MambaLayer) + assert isinstance(block.layers[1], HybridStack) + assert block.layers[1].is_layer_group_stack + assert block.layers[1].layer_type_list == ['M', '*'] + assert isinstance(block.layers[1].layers[0], MambaLayer) + assert isinstance(block.layers[1].layers[1], TransformerLayer) + assert isinstance(block.layers[2], TransformerLayer) + assert [layer.layer_number for layer in block.layers[1].layers] == [2, 3] + assert block.layers[2].layer_number == 4 + # Each physical layer owns the independent per-layer config it was built from. + assert block.layers[1].layers[0].config is layer_config_list[1][0] + assert block.layers[1].layers[1].config is layer_config_list[1][1] + + def test_group_layer_type_list_builds_nested_hybrid_stack(self): + """The deprecated ``layer_type_list`` path accepts symbol tuples for groups.""" + transformer_config = TransformerConfig( + hidden_size=256, num_layers=3, num_attention_heads=4, use_cpu_initialization=True + ) + with pytest.warns(DeprecationWarning, match=r"DEPRECATED\(layer_type_list\)"): + block = HybridStack( + transformer_config, + hybrid_stack_spec.submodules, + layer_type_list=[Symbols.MAMBA, (Symbols.ATTENTION, Symbols.MLP)], + pp_layer_offset=0, + pg_collection=self.get_pg_collection(), + ) + assert block.layer_type_list == ['M', ('*', '-')] + assert isinstance(block.layers[1], HybridStack) + assert [layer.layer_number for layer in block.layers[1].layers] == [2, 3] + + @pytest.mark.parametrize("stack_spec", [hybrid_stack_spec, hybrid_inference_stack_spec]) + def test_group_sharded_state_dict_uses_logical_layer_keys(self, stack_spec): + """Grouped attention+MLP layers share one Transformer-compatible checkpoint key.""" + layer_pattern = "[*-]" + transformer_config = TransformerConfig( + hidden_size=256, + num_layers=_num_physical_layers(layer_pattern), + num_attention_heads=4, + use_cpu_initialization=True, + ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) + block = HybridStack( + transformer_config, + stack_spec.submodules, + layer_config_list=layer_config_list, + pp_layer_offset=0, + logical_layer_offset=0, + # HybridModel sets this from the full layer pattern; a directly + # constructed stack has to opt in itself. + transformer_sharded_keys=True, + pg_collection=self.get_pg_collection(), + ) + + sharded_state_dict = block.sharded_state_dict(prefix="decoder.") + sharded_keys = {value.key for value in sharded_state_dict.values() if hasattr(value, "key")} + + assert "decoder.layers.0.self_attention.linear_qkv.weight" in sharded_keys + assert "decoder.layers.0.mlp.linear_fc1.weight" in sharded_keys + assert "decoder.layers.1.mlp.linear_fc1.weight" not in sharded_keys + assert "decoder.final_layernorm.weight" in sharded_keys + assert "decoder.final_norm.weight" not in sharded_keys + + @pytest.mark.parametrize("pp_layer_offset", [0, 5]) + def test_sharded_state_dict_keeps_historical_keys(self, pp_layer_offset): + """Direct callers retain physical layer indices and the historical final norm key.""" + layer_pattern = "*-" + transformer_config = TransformerConfig( + hidden_size=256, + num_layers=pp_layer_offset + _num_physical_layers(layer_pattern), + num_attention_heads=4, + use_cpu_initialization=True, + ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) + block = HybridStack( + transformer_config, + hybrid_stack_spec.submodules, + layer_config_list=layer_config_list, + pp_layer_offset=pp_layer_offset, + pg_collection=self.get_pg_collection(), + ) + + sharded_state_dict = block.sharded_state_dict(prefix="decoder.") + sharded_keys = {value.key for value in sharded_state_dict.values() if hasattr(value, "key")} + + assert "decoder.final_norm.weight" in sharded_keys + assert "decoder.final_layernorm.weight" not in sharded_keys + assert f"decoder.layers.{pp_layer_offset}.self_attention.linear_qkv.weight" in sharded_keys + assert f"decoder.layers.{pp_layer_offset + 1}.mlp.linear_fc1.weight" in sharded_keys + + def test_group_forward_matches_equivalent_flat_layers(self): + """A bracket group is only a scheduling/checkpoint boundary, not new math.""" + flat_block = self.get_attention_mlp_block("*-") + group_block = self.get_attention_mlp_block("[*-]") + + group_block.layers[0].layers[0].load_state_dict(flat_block.layers[0].state_dict()) + group_block.layers[0].layers[1].load_state_dict(flat_block.layers[1].state_dict()) + group_block.final_norm.load_state_dict(flat_block.final_norm.state_dict()) + + flat_block.cuda().eval() + group_block.cuda().eval() + sequence_length = 16 + micro_batch_size = 2 + hidden_states = torch.randn( + sequence_length, micro_batch_size, flat_block.config.hidden_size, device="cuda" + ) + attention_mask = torch.ones( + (micro_batch_size, 1, sequence_length, sequence_length), dtype=bool, device="cuda" + ) + + with torch.no_grad(): + flat_output = flat_block(hidden_states.clone(), attention_mask=attention_mask) + group_output = group_block(hidden_states.clone(), attention_mask=attention_mask) + + torch.testing.assert_close(group_output, flat_output, rtol=0, atol=0) + + def test_group_overlap_callables_keep_ep_moe_split_visible(self): + """EP-overlap scheduling still sees dispatch/experts/combine inside a group.""" + block = self.get_attention_moe_block("[*E]") + + forward_callables, bwd_dw_callable_map, is_moe, num_local_experts = ( + build_hybrid_stack_callables(block.layers[0], layer_type=block.layer_type_list[0]) + ) + + pre_dispatch, dispatch, experts, combine, mtp_post_process = forward_callables + assert callable(pre_dispatch) + assert callable(dispatch) + assert callable(experts) + assert callable(combine) + assert mtp_post_process is None + assert is_moe + assert num_local_experts == 8 + assert "pre_dispatch_computation" in bwd_dw_callable_map + assert "mlp" in bwd_dw_callable_map + def test_invalid_layer_types_cause_failure(self): invalid_pattern_char = 'X' assert not layer_utils.is_valid_symbol(invalid_pattern_char) # sanity check. diff --git a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py index ac648b0aa36..3dc085a7c70 100644 --- a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py +++ b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py @@ -13,9 +13,12 @@ get_hybrid_total_layer_count, get_hybrid_total_pipeline_segment_count, get_layer_maps_from_layer_type_list, + get_layer_type_list_from_layer_config_list, parse_hybrid_pattern, + parse_segment_layers, pattern_from_ratios, select_pipeline_segment, + select_pipeline_segment_with_logical_offset, validate_segment_layers, ) from megatron.core.models.hybrid.layers import utils as layer_utils @@ -125,6 +128,29 @@ class TestValidateSegmentLayers: def setup_method(self): self.config = _make_transformer_config() + def test_group_patterns(self): + """Bracketed groups parse to symbol tuples and build tuples of per-layer configs.""" + assert parse_segment_layers("M[M*]-") == ['M', ('M', '*'), '-'] + assert parse_segment_layers("[M*E]") == [('M', '*', 'E')] + + result = validate_segment_layers("M[M*]-", self.config) + assert type(result[0]) is MambaLayerConfig + assert isinstance(result[1], tuple) + assert [type(config) for config in result[1]] == [MambaLayerConfig, AttentionLayerConfig] + assert type(result[2]) is MLPLayerConfig + assert get_layer_type_list_from_layer_config_list(result) == ['M', ('M', '*'), '-'] + flat_configs = [result[0], *result[1], result[2]] + assert len({id(config) for config in flat_configs}) == len(flat_configs) + + @pytest.mark.parametrize("group", ["MM", "--", "**", "GG", "G*", "D+", "-E", "CH", "CW", "+C"]) + def test_groups_reject_checkpoint_namespace_collisions(self, group): + with pytest.raises(ValueError, match="multiple layers in checkpoint namespace"): + validate_segment_layers(f"[{group}]", self.config) + + @pytest.mark.parametrize("group", ["M*E", "M*-", "MGE", "M+E", "MD-", "*", "MCE", "MHE", "MWE"]) + def test_groups_with_distinct_checkpoint_namespaces(self, group): + assert parse_segment_layers(f"[{group}]") == [tuple(group)] + def test_valid_patterns(self): """Test that valid segment patterns produce configs in the correct order.""" for pattern in [ @@ -216,6 +242,10 @@ def test_invalid_symbols_cause_failure(self): validate_segment_layers("M|M", self.config) # pipe not valid in a segment with pytest.raises(ValueError): validate_segment_layers("M/M", self.config) # MTP separator not valid in a segment + with pytest.raises(ValueError): + validate_segment_layers("M[[M]]", self.config) # nested groups are not valid + with pytest.raises(ValueError): + validate_segment_layers("M[EM]", self.config) # MoE must be last in a group with pytest.raises(ValueError): # Not allowed to have both standard Attention and MLA/DSA validate_segment_layers("MDM*-", self.config) @@ -244,6 +274,8 @@ def test_simple_patterns(self): assert get_hybrid_total_layer_count("M*M*") == 4 assert get_hybrid_total_layer_count("MMMM") == 4 assert get_hybrid_total_layer_count("M") == 1 + assert get_hybrid_total_layer_count("[M*E]") == 3 + assert get_hybrid_total_layer_count("M[M*]-") == 4 def test_with_pipe_separators(self): assert get_hybrid_total_layer_count("M-M-|M-M*-") == 9 @@ -289,6 +321,8 @@ def test_main_pattern_only(self): """Test patterns without MTP (no / separator).""" test_cases = [ ("M*M*", "M*M*"), + ("[M*E]", "[M*E]"), + ("M[M*]-", "M[M*]-"), ("MMMM", "MMMM"), ("*M*M", "*M*M"), ("MM-*", "MM-*"), @@ -373,11 +407,22 @@ def test_invalid_symbols_in_main_pattern(self): "M*X*", # X is not valid "MaMM", # a is not valid "M*M*1", # 1 is not valid + "M[M*]X", # X is not valid after a group ] for pattern in invalid_patterns: with pytest.raises(ValueError, match="not a valid layer symbol"): parse_hybrid_pattern(pattern) + def test_invalid_group_syntax(self): + with pytest.raises(ValueError, match="without a matching"): + parse_hybrid_pattern("M[M*") + with pytest.raises(ValueError, match="not supported"): + parse_hybrid_pattern("M[M[*]]") + with pytest.raises(ValueError, match="cannot be empty"): + parse_hybrid_pattern("M[]") + with pytest.raises(ValueError, match="must be the last"): + parse_hybrid_pattern("M[EM]") + def test_invalid_symbols_in_mtp_pattern(self): """Test that invalid symbols in MTP pattern raise ValueError.""" # Single MTP depth with invalid symbol - should raise "not a valid layer symbol" @@ -573,6 +618,17 @@ def test_moe_pattern(self): 'W': 0, } + def test_group_pattern(self): + assert get_hybrid_layer_counts("M[M*]E") == { + '*': 1, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 0, + '-': 0, + 'E': 1, + } + def test_mtp_with_attention(self): # MTP pattern "*M" repeated 3 depths -> 3 attn + 3 mamba from MTP assert get_hybrid_layer_counts("MMMM/*M/*M/*M") == { @@ -733,6 +789,25 @@ def test_four_segments(self, mock_log): _assert_layer_config_types(layer_configs, expected_pattern) assert offset == expected_offset, f"Failed for vp_stage={vp_stage}" + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') + def test_group_segment_offsets(self, mock_log): + layer_configs, offset = select_pipeline_segment( + "[M*E]|M-", self.config, pp_group=None, vp_stage=1 + ) + _assert_layer_config_types(layer_configs, "M-") + assert offset == 3 + + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') + def test_group_segment_logical_offsets(self, mock_log): + layer_configs, physical_offset, logical_offset = ( + select_pipeline_segment_with_logical_offset( + "[*-][*-]|[*E][*E]", self.config, pp_group=None, vp_stage=1 + ) + ) + assert get_layer_type_list_from_layer_config_list(layer_configs) == [('*', 'E'), ('*', 'E')] + assert physical_offset == 4 + assert logical_offset == 2 + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_empty_segment(self, mock_log): """Empty segments are allowed for pipeline balancing.""" @@ -1095,3 +1170,12 @@ def test_mixed_dsa_and_mla(self): assert mamba_map == {2: 0} assert mlp_map == {3: 0} assert moe_map == {} + + def test_grouped_layers_are_flattened(self): + maps = get_layer_maps_from_layer_type_list([("M", "*", "E"), "M"]) + attention_map, mamba_map, moe_map = operator.itemgetter( + Symbols.ATTENTION, Symbols.MAMBA, Symbols.MOE + )(maps) + assert attention_map == {1: 0} + assert mamba_map == {0: 0, 3: 1} + assert moe_map == {2: 0} diff --git a/tests/unit_tests/transformer/moe/test_aux_loss.py b/tests/unit_tests/transformer/moe/test_aux_loss.py index d631dad35d2..a88358ff42a 100644 --- a/tests/unit_tests/transformer/moe/test_aux_loss.py +++ b/tests/unit_tests/transformer/moe/test_aux_loss.py @@ -11,6 +11,7 @@ get_cuda_rng_tracker, model_parallel_cuda_manual_seed, ) +from megatron.core.transformer.moe.moe_logging import destroy_moe_metrics_tracker from megatron.core.transformer.moe.moe_utils import ( clear_aux_losses_tracker, get_default_pg_collection, @@ -133,6 +134,45 @@ def test_a2a_dispatcher(self, tp_size, ep_size, cp_size): container.aux_loss_test(self.input, self.baseline_grad, "load_balancing_loss") +@pytest.mark.internal +def test_z_loss_uses_explicit_hybrid_mtp_depth_for_tracker_slot(): + """Hybrid MTP routers must record z-loss in one of the configured MTP slots.""" + + class DummyGroup: + @staticmethod + def size(): + return 1 + + class DummyConfig: + moe_z_loss_coeff = 1.0 + mtp_num_layers = 2 + mtp_use_repeated_layer = False + num_layers = 8 + + class DummyRouter: + config = DummyConfig() + tp_cp_group = DummyGroup() + tp_dp_cp_group = DummyGroup() + training = True + calculate_per_token_loss = False + is_mtp_layer = True + layer_number = 5 + mtp_layer_number = 1 + _get_metric_layer_number = TopKRouter._get_metric_layer_number + + destroy_moe_metrics_tracker() + try: + logits = torch.randn(4, 3, requires_grad=True) + TopKRouter.apply_z_loss(DummyRouter(), logits) + + values = get_moe_layer_wise_logging_tracker()["z_loss"]["values"] + assert values.shape == (10,) + assert values[8] > 0 + assert torch.count_nonzero(values) == 1 + finally: + destroy_moe_metrics_tracker() + + class TestSeqAuxLoss: def setup_method(self, method): baseline_container = AuxlossTestContainer( diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index 66ee60a51c5..7cce9afad44 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -3620,6 +3620,8 @@ def compute_language_model_loss(labels, logits): use_mup=False, inference_cuda_graph_scope=None, sequence_parallel=False, + cuda_graph_impl='none', + flash_decode=False, ), pre_process=False, post_process=True, @@ -3638,6 +3640,8 @@ def compute_language_model_loss(labels, logits): tp_group=None, _scale_logits=lambda logits: logits, ) + model._preprocess = types.MethodType(HybridModel._preprocess, model) + model._postprocess = types.MethodType(HybridModel._postprocess, model) return model, hidden_states, call_counts, metric_avg_group @pytest.mark.parametrize( diff --git a/tests/unit_tests/transformer/test_submodule_callables.py b/tests/unit_tests/transformer/test_submodule_callables.py index 029fcb88865..33af4ff6bcc 100644 --- a/tests/unit_tests/transformer/test_submodule_callables.py +++ b/tests/unit_tests/transformer/test_submodule_callables.py @@ -103,9 +103,11 @@ def run_model_submodules_with_capture(model, input_tensors, microbatches): return capture -def test_mtp_pre_dispatch_applies_hybrid_empty_decoder_final_norm(monkeypatch): - """Covers the HybridModel empty-decoder MTP pre-dispatch final_norm path.""" +@pytest.mark.parametrize("mtp_offset", [0, 1]) +def test_mtp_pre_dispatch_applies_hybrid_empty_decoder_final_norm(monkeypatch, mtp_offset): + """MTP slot boundaries preserve final-norm values and gradients.""" + from megatron.core.models.common.utils import TransformerLayerNode from megatron.core.models.hybrid.hybrid_model import HybridModel def inner_pre_dispatch(_node, hidden_states): @@ -149,35 +151,59 @@ def _postprocess(self, hidden_states): monkeypatch.setattr(common_callables, "build_layer_callables", fake_build_layer_callables) monkeypatch.setattr(common_callables, "get_layer_moe_metadata", lambda _layer: (True, 1)) monkeypatch.setattr( - common_callables, "get_mtp_layer_offset", lambda _config, _vp_stage, pp_rank=None: 0 + common_callables, + "get_mtp_layer_offset", + lambda _config, _vp_stage, pp_rank=None: mtp_offset, ) model = HybridModel.__new__(HybridModel) torch.nn.Module.__init__(model) model.decoder = DummyState() model.decoder.layers = [] - model.decoder.final_norm = lambda hidden_states: hidden_states + 4.0 + model.decoder.final_norm = torch.nn.LayerNorm(3) + reference_norm = torch.nn.LayerNorm(3) + reference_norm.load_state_dict(model.decoder.final_norm.state_dict()) model.embedding = object() model.vp_stage = None model.pg_collection = DummyState() model.pg_collection.pp = DummyState() model.pg_collection.pp.rank = lambda: 0 - node = DummyNode() + node = TransformerLayerNode.__new__(TransformerLayerNode) + node.before_detached = () + node.detached = () + node.default_backward_func = torch.autograd.backward node.chunk_state = DummyState() node.chunk_state.model = model node.chunk_state.context = None node.chunk_state.packed_seq_params = None node.is_first_layer = True - hidden_states = torch.arange(6, dtype=torch.float32).reshape(3, 1, 2).requires_grad_() - expected = hidden_states + 4.0 + hidden_states = torch.arange(12, dtype=torch.float32).reshape(4, 1, 3).requires_grad_() + reference_input = hidden_states.detach().clone().requires_grad_() + reference_chunks = list(torch.chunk(reference_norm(reference_input), 1 + mtp_offset, dim=0)) forward_funcs, _ = common_callables.build_mtp_layer_callables(FakeMTPLayer()) output = forward_funcs[0](node, hidden_states) - - torch.testing.assert_close(output, expected) - torch.testing.assert_close(node.chunk_state.mtp_hidden_states[0], expected) + torch.testing.assert_close(output, reference_chunks[mtp_offset]) + for actual, expected in zip(node.chunk_state.mtp_hidden_states, reference_chunks): + torch.testing.assert_close(actual, expected) + assert actual.is_leaf + + # MTP post-process runs its backward before pre-dispatch. Its copy of the + # pre-dispatch output is detached by ScheduleNode at the slot boundary. + scheduled_output = output.detach().requires_grad_() + node.is_last_layer = True + combined = forward_funcs[4](node, scheduled_output) + coefficients = torch.arange(1, combined.numel() + 1, dtype=combined.dtype).reshape_as(combined) + (combined * coefficients).square().sum().backward() + node.backward_impl((output,), (scheduled_output.grad,)) + + reference_combined = torch.cat(reference_chunks + [reference_chunks[mtp_offset]], dim=0) + (reference_combined * coefficients).square().sum().backward() + torch.testing.assert_close(hidden_states.grad, reference_input.grad) + torch.testing.assert_close(model.decoder.final_norm.weight.grad, reference_norm.weight.grad) + torch.testing.assert_close(model.decoder.final_norm.bias.grad, reference_norm.bias.grad) class TestTransformerLayerSubmoduleCallables: From b38bebfcf668d6b9d7982cdb48584dbf65260775 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 29 Sep 2026 15:28:50 -0700 Subject: [PATCH 02/11] Cover delayed Mamba projection weight gradients in deterministic replay Signed-off-by: Yan Xu --- .../correctness/test_ssm_conv1d.py | 46 +++++++++++++++---- 1 file changed, 38 insertions(+), 8 deletions(-) diff --git a/tests/unit_tests/determinism/correctness/test_ssm_conv1d.py b/tests/unit_tests/determinism/correctness/test_ssm_conv1d.py index a59e4cadb69..b81141b9d3a 100644 --- a/tests/unit_tests/determinism/correctness/test_ssm_conv1d.py +++ b/tests/unit_tests/determinism/correctness/test_ssm_conv1d.py @@ -18,6 +18,7 @@ from megatron.core.ssm.mamba_mixer import MambaMixer from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig +from megatron.core.utils import is_te_min_version from tests.unit_tests.determinism.configs import hybrid_base from tests.unit_tests.determinism.utils import ( assert_bit_exact, @@ -77,13 +78,26 @@ def _conv_backward(x, weight, bias, grad): return torch.autograd.grad(out, (x, weight, bias), grad_outputs=grad) -def _build_mixer(deterministic_mode=True): +def _build_mixer(deterministic_mode=True, delay_wgrad_compute=False): """A bare MambaMixer on the suite's shared hybrid config.""" - Utils.initialize_model_parallel() - model_parallel_cuda_manual_seed(123) - config = TransformerConfig( - **(hybrid_base() | {"num_layers": 1, "deterministic_mode": deterministic_mode}) + overrides = {"num_layers": 1, "deterministic_mode": deterministic_mode} + if delay_wgrad_compute: + # Delayed projection wgrad is enabled by the EP-overlap configuration, + # even though this test isolates a single dense Mamba mixer. It runs + # no EP communication or overlap streams, so it needs no connection-count override. + overrides.update( + delay_wgrad_compute=True, + overlap_moe_expert_parallel_comm=True, + expert_model_parallel_size=2, + num_moe_experts=2, + moe_token_dispatcher_type="alltoall", + add_bias_linear=False, + ) + Utils.initialize_model_parallel( + expert_model_parallel_size=overrides.get("expert_model_parallel_size", 1) ) + model_parallel_cuda_manual_seed(123) + config = TransformerConfig(**(hybrid_base() | overrides)) mixer = MambaMixer( config, hybrid_stack_spec.submodules.mamba_layer.submodules.mixer.submodules, @@ -170,15 +184,21 @@ def teardown_method(self, method): @requires_deterministic_conv1d @pytest.mark.parametrize("packed_layout", ["none", "unpadded", "padded"]) - def test_mixer_replays_bit_exactly(self, monkeypatch, packed_layout): + @pytest.mark.parametrize("delay_wgrad_compute", [False, True]) + def test_mixer_replays_bit_exactly(self, monkeypatch, packed_layout, delay_wgrad_compute): """Two runs of one mixer agree bitwise under the deterministic conv reduction. ``MAMBA_DETERMINISTIC`` pins the SSD scan, whose nondeterminism would otherwise reach - the conv weight gradient, so a failure points at the convolution. + the conv weight gradient. The delayed arm also covers projection gradients completed + through ``MambaMixer.backward_dw`` after the regular backward pass. """ + if delay_wgrad_compute and (Utils.world_size < 2 or Utils.world_size % 2): + pytest.skip("Delayed projection wgrad config requires a world size divisible by EP=2") + if delay_wgrad_compute and not is_te_min_version("2.3.0"): + pytest.skip("Delayed projection wgrad requires TE >= 2.3.0") monkeypatch.setenv("MAMBA_DETERMINISTIC", "1") monkeypatch.setenv("CAUSAL_CONV1D_DETERMINISTIC", "1") - mixer = _build_mixer() + mixer = _build_mixer(delay_wgrad_compute=delay_wgrad_compute) hidden_size = mixer.config.hidden_size micro_batch = _MICRO_BATCH if packed_layout == "none" else 1 packed_seq_params = None @@ -207,6 +227,16 @@ def test_mixer_replays_bit_exactly(self, monkeypatch, packed_layout): def fwd_bwd(): output, _ = mixer(hidden_states, packed_seq_params=packed_seq_params) output.backward(grad) + if delay_wgrad_compute: + # Prove that the callback completes deferred work, so replay + # cannot pass merely because both runs omitted projection grads. + for projection in (mixer.in_proj, mixer.out_proj): + assert projection.weight.grad is None + mixer.backward_dw() + for projection in (mixer.in_proj, mixer.out_proj): + assert projection.weight.grad is not None + assert torch.isfinite(projection.weight.grad).all() + assert torch.count_nonzero(projection.weight.grad) > 0 return output.detach().clone(), collect_grads([mixer]) state = capture_rng_state() From c0d9ae2cd67bb291c78dae0821c17c327064d347 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 29 Sep 2026 15:35:30 -0700 Subject: [PATCH 03/11] Fix grouped hybrid layer-count expectation for DSv4 symbols Signed-off-by: Yan Xu --- tests/unit_tests/ssm/test_hybrid_layer_allocation.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py index 3dc085a7c70..eceb16fd52f 100644 --- a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py +++ b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py @@ -621,12 +621,15 @@ def test_moe_pattern(self): def test_group_pattern(self): assert get_hybrid_layer_counts("M[M*]E") == { '*': 1, + 'C': 0, 'D': 0, 'G': 0, + 'H': 0, 'M': 2, '+': 0, '-': 0, 'E': 1, + 'W': 0, } def test_mtp_with_attention(self): From 2c144e2b67d9fe22ef3d168db58de098f20536f5 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 29 Sep 2026 15:41:59 -0700 Subject: [PATCH 04/11] Preserve shortcut checkpoint keys in grouped HybridStack support Signed-off-by: Yan Xu --- megatron/core/models/hybrid/hybrid_block.py | 14 ++++++-- tests/unit_tests/ssm/test_hybrid_block.py | 36 +++++++++++++++++++++ 2 files changed, 47 insertions(+), 3 deletions(-) diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index 3cfebe63bfd..bfd368668fe 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -1079,14 +1079,22 @@ def _sharded_state_dict( if sharded_layer_prefix is None: sharded_layer_prefix = layer_prefix - for local_layer_idx, (layer_config, layer) in enumerate( - zip(self.layer_config_list, self.layers, strict=True) + for local_layer_idx, (source_layer_idx, layer_config, layer) in enumerate( + zip( + self._execution_layer_indices, + self._execution_layer_config_list, + self.layers, + strict=True, + ) ): state_dict_prefix = f'{layer_prefix}{local_layer_idx}.' # module list index + # Shortcut blocks collapse adjacent physical layers, while bracketed groups + # already occupy one logical slot. Keep the index from before shortcut + # grouping so both the shortcut and subsequent layers retain their old keys. logical_layer_idx = ( self.logical_layer_offset if self.is_layer_group_stack - else self.logical_layer_offset + local_layer_idx + else self.logical_layer_offset + source_layer_idx ) if is_layer_group(layer_config): diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index 9e5279f8959..adbad0fe949 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -8,6 +8,7 @@ import megatron.core.models.hybrid.hybrid_block as hybrid_block_module import megatron.core.transformer.utils as transformer_utils +from megatron.core.dist_checkpointing.mapping import ShardedTensor from megatron.core.extensions.transformer_engine import TEDotProductAttention from megatron.core.models.hybrid.fine_grained_callables import build_hybrid_stack_callables from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules @@ -57,6 +58,41 @@ def _num_physical_layers(layer_pattern: str) -> int: return len(layer_pattern.replace(Symbols.GROUP_START, '').replace(Symbols.GROUP_END, '')) +@pytest.mark.parametrize("pp_layer_offset", [0, 5]) +def test_shortcut_checkpoint_keys_preserve_physical_layer_offsets(pp_layer_offset): + """Collapsed shortcut pairs must not renumber following layers or truncate saving.""" + + class CheckpointLayer(torch.nn.Module): + def __init__(self, layer_number): + super().__init__() + self.layer_number = layer_number + self.weight = torch.nn.Parameter(torch.ones(1)) + + def sharded_state_dict(self, prefix, sharded_offsets, metadata): + key = f'{prefix}weight' + return {key: ShardedTensor.from_rank_offsets(key, self.weight)} + + # Two shortcut pairs followed by a normal layer: five physical layers become + # three executable modules, whose storage positions remain 0, 2, and 4. + block = HybridStack.__new__(HybridStack) + torch.nn.Module.__init__(block) + block.layers = torch.nn.ModuleList( + [CheckpointLayer(pp_layer_offset + index + 1) for index in (0, 2, 4)] + ) + block.layer_config_list = [object() for _ in range(5)] + block._execution_layer_indices = [0, 2, 4] + block._execution_layer_config_list = [block.layer_config_list[index] for index in (0, 2, 4)] + block.logical_layer_offset = pp_layer_offset + block.is_layer_group_stack = False + block.transformer_sharded_keys = False + + sharded_state_dict = block.sharded_state_dict(prefix='decoder.') + assert set(sharded_state_dict) == {f'decoder.layers.{index}.weight' for index in range(3)} + assert [value.key for value in sharded_state_dict.values()] == [ + f'decoder.layers.{pp_layer_offset + index}.weight' for index in (0, 2, 4) + ] + + def test_wide_residual_spec_preserves_unmodified_stack_submodules(): """Wide specialization changes layer classes without dropping stack-level specs.""" From a9abddf2a6a6132271a5c38157f8d129221ebd0a Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 29 Sep 2026 15:42:41 -0700 Subject: [PATCH 05/11] Preserve static hybrid inference with caller-supplied embeddings Signed-off-by: Yan Xu --- megatron/core/models/hybrid/hybrid_model.py | 12 +++++--- .../models/test_hybrid_hash_routing.py | 28 +++++++++++++++++++ 2 files changed, 36 insertions(+), 4 deletions(-) diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 2e5c7a32f4a..07449cde82e 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -653,11 +653,15 @@ def _preprocess( and (self.config.cuda_graph_impl == "local" or self.config.flash_decode) and inference_context.is_static_batching() ): - current_batch_size = input_ids.shape[0] - sequence_len_offset = torch.tensor( - [inference_context.sequence_len_offset] * current_batch_size, + # Caller-supplied embeddings and later PP stages may omit token IDs. + hidden_states = ( + decoder_input if decoder_input is not None else self.decoder.input_tensor + ) + sequence_len_offset = torch.full( + (hidden_states.shape[1],), + inference_context.sequence_len_offset, dtype=torch.int32, - device='cuda', + device=hidden_states.device, ) else: sequence_len_offset = None diff --git a/tests/unit_tests/models/test_hybrid_hash_routing.py b/tests/unit_tests/models/test_hybrid_hash_routing.py index e46d1028250..c93cb4b9b13 100644 --- a/tests/unit_tests/models/test_hybrid_hash_routing.py +++ b/tests/unit_tests/models/test_hybrid_hash_routing.py @@ -401,6 +401,34 @@ def fake_scatter(tensor, group): assert torch.equal(decoder.kwargs["padding_mask"], padding_mask[:, :2]) +@pytest.mark.parametrize("pre_process", [True, False]) +@pytest.mark.parametrize("cuda_graph_impl,flash_decode", [("local", False), ("none", True)]) +def test_hybrid_preprocess_static_inference_without_token_ids( + pre_process, cuda_graph_impl, flash_decode +): + hidden_states = torch.randn(4, 2, 8) + model = SimpleNamespace( + config=SimpleNamespace( + sequence_parallel=False, cuda_graph_impl=cuda_graph_impl, flash_decode=flash_decode + ), + pre_process=pre_process, + position_embedding_type="none", + decoder=SimpleNamespace(input_tensor=hidden_states), + ) + inference_context = SimpleNamespace(is_static_batching=lambda: True, sequence_len_offset=7) + + with InferenceMode.active(): + *_, sequence_len_offset, _ = HybridModel._preprocess( + model, + input_ids=None, + position_ids=None, + decoder_input=hidden_states if pre_process else None, + inference_context=inference_context, + ) + + torch.testing.assert_close(sequence_len_offset, torch.tensor([7, 7], dtype=torch.int32)) + + def test_chunked_hash_moe_keeps_ids_and_padding_aligned(): moe = RecordingMoE() layer = SimpleNamespace( From 14e75a89d6e4bf2e059049d4919df069a9fa2f62 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Wed, 30 Sep 2026 15:16:52 -0700 Subject: [PATCH 06/11] Make hybrid stack comments describe the current contract Rewrite comments that only made sense relative to the diff or to a past incident: drop a line-number reference in favor of the function name, state the recompute-hook invariant instead of the symptom it caused, and describe the checkpoint keys each pattern publishes instead of calling them historical or pre-existing. Signed-off-by: Yan Xu --- .../models/hybrid/fine_grained_callables.py | 44 +++++++------------ megatron/core/models/hybrid/hybrid_block.py | 20 ++++----- megatron/core/models/hybrid/hybrid_model.py | 8 ++-- 3 files changed, 30 insertions(+), 42 deletions(-) diff --git a/megatron/core/models/hybrid/fine_grained_callables.py b/megatron/core/models/hybrid/fine_grained_callables.py index 1371e3b0b9a..af2ccf7bb98 100644 --- a/megatron/core/models/hybrid/fine_grained_callables.py +++ b/megatron/core/models/hybrid/fine_grained_callables.py @@ -70,22 +70,16 @@ class HybridStackNode(TransformerLayerNode): Subclassed from ``TransformerLayerNode`` so the runtime backbone (forward / backward / backward_dw plumbing, detach bookkeeping, output-grad release) - is shared. The hybrid path keeps a separate node class so its free-input - policy can diverge from the GPT defaults — for example, the + is shared. The subclass is where the hybrid free-input policy lives: the ``pre_dispatch_computation`` slot here covers the whole pre-dispatch loop - (mamba + attention + …) rather than a single attention block, and - group-level decisions about whether the input is needed in backward may - differ from ``should_free_input`` in ``gpt/fine_grained_callables.py``. - Keep this override thin until a hybrid counter-example forces it to - diverge; the explicit subclass exists so the divergence can be made - surgically without touching the GPT class. + (mamba + attention + …) rather than a single attention block. """ @staticmethod def _resolve_free_input(name, is_moe, config, num_local_experts): """Hybrid free-input policy. - Currently mirrors the GPT default: dense layers always retain their + Same as the GPT default: dense layers always retain their input for backward; MoE-only "moe_dispatch", "mlp", and "moe_combine" slots can free, subject to the dispatcher / cuda-graph constraints encoded in ``should_free_input``. Hybrid groups have a @@ -93,8 +87,7 @@ def _resolve_free_input(name, is_moe, config, num_local_experts): loop over Mamba/attention/GDN sub-layers, not a single attention block), but its policy resolves to ``False`` in ``should_free_input``, which is correct: pre-layer outputs are needed - for backward through the loop. Override here when a hybrid-specific - rule is needed. + for backward through the loop. """ return should_free_input(name, is_moe, config, num_local_experts) @@ -225,16 +218,12 @@ def _run_moe_combine(layer, node: ScheduleNode, output: Tensor): shared_expert_output = getattr(node.layer_state, 'shared_expert_output', None) output = layer.mlp.combine(output) output = layer.mlp.postprocess(output, shared_expert_output) - # Inline bda instead of calling ``layer._forward_post_mlp`` so we can skip - # the redundant ``discard_output_and_register_recompute(mlp_output_with_bias[0])`` - # that ``_forward_post_mlp`` would otherwise issue. The pre_mlp_layernorm recompute - # is already registered on ``expert_output`` inside ``_run_moe_experts``; the second - # hook on the combine-slot ``mlp_output_with_bias[0]`` is not only unnecessary but - # harmful in the bracketed-hybrid case (``[*E]``): it fires during combine_bwd's - # autograd backward and triggers the LN recompute ahead of attention's backward - # in the same pre_dispatch slot, corrupting attention gradients (grad_norm explodes - # from iter 2). GPT's ``submodule_combine_forward`` likewise inlines bda and does - # not call ``_forward_post_mlp`` for the same reason. + # Inline bda instead of calling ``layer._forward_post_mlp``, which would register + # a second pre_mlp_layernorm recompute hook on the combine output. The hook is + # already registered on ``expert_output`` in ``_run_moe_experts``. A second one + # would run the layernorm recompute before attention's backward in the same + # pre-dispatch slot of a bracketed group (``[*E]``), which corrupts the attention + # gradients. GPT's ``submodule_combine_forward`` inlines bda for the same reason. mlp_output_with_bias = (output, None) with layer.bias_dropout_add_exec_handler(): output = layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( @@ -260,7 +249,7 @@ def build_hybrid_stack_callables(layer, layer_type: Optional[LayerPatternItem] = """Create fine-grained callables for one logical HybridStack layer. A logical layer may be a bracketed nested ``HybridStack`` (for example ``[M*E]``) - or a single legacy hybrid layer symbol. The split is: + or a single ungrouped hybrid layer symbol. The split is: pre-dispatch compute -> dispatch -> MLP/experts -> combine. """ pre_layers, terminal_type, terminal_layer, is_moe, num_local_experts = ( @@ -301,12 +290,11 @@ def pre_dispatch_computation(node: ScheduleNode, hidden_states: Tensor): # a view from a fused/JIT kernel). Downstream cuBLAS matmuls — including # the terminal MLP/MoE's pre_mlp_layernorm and the next attention's QKV # projection in a multi-pre-layer group — pick algorithms based on input - # strides; a view's non-canonical strides can lead to different algo - # selection across processes and produce ~1e-5 bit drift on the forward - # output. TransformerLayer's full forward() inserts this exact call at the - # MLP exit (transformer_layer.py:895) for the same reason; the - # _forward_attention shortcut here doesn't get that cleanup, so we add it - # explicitly. Same idea as the make_viewless_tensor in _maybe_apply_final_norm. + # strides, so a view can make the forward output differ slightly between + # processes. TransformerLayer makes its MLP output viewless for the same + # reason (``TransformerLayer._apply_mlp_bda_step``); the _forward_attention + # shortcut here does not, so do it explicitly. Same idea as the + # make_viewless_tensor in _maybe_apply_final_norm. hidden_states = make_viewless_tensor( inp=hidden_states, requires_grad=hidden_states.requires_grad, diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index bfd368668fe..d3be1ef3a42 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -128,7 +128,7 @@ class HybridStack(MegatronModule): segment. Defaults to 0. logical_layer_offset (int, optional): the global logical layer offset for this pipeline segment; bracketed groups count as one logical layer. Used for - checkpoint keys. Defaults to ``pp_layer_offset`` for legacy direct callers. + checkpoint keys. Defaults to ``pp_layer_offset``. is_layer_group_stack (bool, optional): whether this stack is the nested stack built for a bracketed group. Defaults to False. post_layer_norm (bool, optional): whether to include a final layer norm. @@ -178,8 +178,8 @@ def __init__( checkpoint keys (``final_layernorm`` instead of ``final_norm``) so the checkpoint is interchangeable with a ``GPTModel`` one. Only set for bracketed-group patterns, whose logical layers map one-to-one onto - transformer layers; leaving it off keeps the historical hybrid keys so - existing non-grouped hybrid checkpoints stay loadable. + transformer layers; when off, the final norm is published as + ``final_norm``, the HybridModel checkpoint key. name (str | None): module instance name passed top-down from its paranet module """ if (layer_type_list is None) == (layer_config_list is None): @@ -1089,8 +1089,9 @@ def _sharded_state_dict( ): state_dict_prefix = f'{layer_prefix}{local_layer_idx}.' # module list index # Shortcut blocks collapse adjacent physical layers, while bracketed groups - # already occupy one logical slot. Keep the index from before shortcut - # grouping so both the shortcut and subsequent layers retain their old keys. + # already occupy one logical slot. Use the index from before shortcut + # grouping so a shortcut block and the layers after it are keyed by their + # logical layer index. logical_layer_idx = ( self.logical_layer_offset if self.is_layer_group_stack @@ -1126,11 +1127,10 @@ def _sharded_state_dict( module_sharded_state_dict = sharded_state_dict_default( module, module_prefix, sharded_offsets, metadata, tp_group=self.tp_group ) - # The registered submodule stays ``final_norm`` (local state-dict keys - # are unchanged), but grouped stacks publish the sharded key under - # TransformerBlock's ``final_layernorm`` name so their checkpoints - # cross-load with GPTModel. Non-grouped stacks keep ``final_norm`` so - # hybrid checkpoints written before this feature still load. + # Ungrouped stacks publish the final norm as ``final_norm``; grouped + # stacks publish it as ``final_layernorm``, matching TransformerBlock, so + # their checkpoints cross-load with GPTModel. The registered submodule + # (and so the local state-dict key) is ``final_norm`` in both cases. if name == 'final_norm' and self.transformer_sharded_keys: replace_prefix_for_sharding( module_sharded_state_dict, module_prefix, f'{prefix}final_layernorm.' diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 07449cde82e..c1484ca2bbb 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -334,8 +334,8 @@ def __init__( # Bracketed-group patterns give every logical layer the structure of a # transformer layer, so their checkpoints are made key-compatible with # GPTModel. Derived from the full pattern rather than this rank's segment so - # every PP stage agrees on the naming. Non-grouped patterns keep the - # historical hybrid keys, which existing hybrid checkpoints were saved with. + # every PP stage agrees on the naming. Ungrouped patterns use the HybridModel + # checkpoint keys. transformer_sharded_keys = Symbols.GROUP_START in (parsed.main_pattern or '') logging_pg_kwargs = _hybrid_logging_pg_kwargs(self.pg_collection) @@ -978,8 +978,8 @@ def sharded_state_dict( sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata) output_layer_extra_state_key = f'{prefix}output_layer._extra_state' - # Match GPTModel checkpoint compatibility: old GPT checkpoints do not include - # output layer extra state, and the TE extra state should be empty. + # Like GPTModel, publish only the output layer weight: drop its TE extra state, + # which is expected to be empty. output_extra_state = sharded_state_dict.pop(output_layer_extra_state_key, None) assert not ( output_extra_state and output_extra_state.data From 74f1397416f157d54806901c70c763cf0249889a Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Wed, 30 Sep 2026 15:17:04 -0700 Subject: [PATCH 07/11] Validate hybrid EP-overlap support once in HybridModel The schedule plan checked static config on every build_schedule_plan() call, i.e. every microbatch, walked all modules each time, and mixed assert with ValueError. Move the checks to HybridModel.__init__ and raise ValueError: - CUDA graphs (the message no longer claims the check is specific to grouped patterns; it applies to any HybridModel with EP overlap), - mHC connections and MoE shortcut connections, - mixed context-parallel layouts, checked on every HybridStack including nested bracketed groups. Drop the hash-routed MoE and wide-residual checks: TransformerConfig already rejects both with overlap_moe_expert_parallel_comm. Signed-off-by: Yan Xu --- megatron/core/models/hybrid/hybrid_model.py | 36 +++++++++ .../hybrid/model_chunk_schedule_plan.py | 28 ------- .../a2a_overlap/test_hybrid_schedule_plan.py | 73 ++++++++++--------- 3 files changed, 76 insertions(+), 61 deletions(-) diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index c1484ca2bbb..123755f60ba 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -17,6 +17,7 @@ from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule +from megatron.core.models.hybrid.hybrid_block import HybridStack from megatron.core.models.hybrid.layers import utils as layer_utils from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( @@ -473,11 +474,46 @@ def __init__( if self.pre_process or self.post_process or self.mtp_process: self.setup_embeddings_and_output_layer() + if self.config.overlap_moe_expert_parallel_comm: + self._validate_ep_overlap_support() + for name, module in self.named_modules(): if hasattr(module, 'finish_init'): quant_config = get_quant_config_or_none(name, self.config.quant_recipe) module.finish_init(quant_config) + def _validate_ep_overlap_support(self) -> None: + """Reject features that the EP-overlap schedule plan cannot run. + + Hash-routed MoE layers and wide residuals are rejected by TransformerConfig. + """ + if self.config.cuda_graph_impl != "none": + raise ValueError( + "overlap_moe_expert_parallel_comm with HybridModel does not support CUDA graphs " + "yet. Set cuda_graph_impl='none'." + ) + if self.config.enable_mhc_connections or self.config.moe_shortcut_connection: + raise ValueError( + "overlap_moe_expert_parallel_comm with HybridModel does not support " + "enable_mhc_connections or moe_shortcut_connection." + ) + # The schedule plan calls the layer callables directly and bypasses + # ``HybridStack.forward``, which is where per-layer context-parallel layout + # conversion happens. A bracketed group presents the boundary layout to its + # enclosing stack even when its own layers need conversion, so check every + # stack, including the nested ones. + for module in self.modules(): + if ( + isinstance(module, HybridStack) + and module._cp_layout_manager is not None + and module._cp_layout_manager.requires_conversion + ): + raise ValueError( + "overlap_moe_expert_parallel_comm with HybridModel does not support mixed " + "context-parallel layouts (linear_cp_layout != attention_cp_layout with " + "context_parallel_size > 1)." + ) + def set_input_tensor(self, input_tensor: Tensor) -> None: """Sets input tensor to the model. diff --git a/megatron/core/models/hybrid/model_chunk_schedule_plan.py b/megatron/core/models/hybrid/model_chunk_schedule_plan.py index 73ab5bc6c8a..3aebbc4590b 100644 --- a/megatron/core/models/hybrid/model_chunk_schedule_plan.py +++ b/megatron/core/models/hybrid/model_chunk_schedule_plan.py @@ -120,34 +120,6 @@ class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): LAYER_SCHEDULE_PLAN_CLASS = HybridStackSchedulePlan - def __init__(self, model, *args, **kwargs): - """Initialize the hybrid chunk plan after validating cuda graph support.""" - assert model.config.cuda_graph_impl == "none", ( - "EP A2A overlap with grouped HybridStack patterns (e.g. '[*E]') does not " - "support cuda graphs yet. Set cuda_graph_impl='none' or use an ungrouped pattern." - ) - if getattr(model.config, "moe_num_hash_layers", 0): - raise ValueError("HybridStack EP overlap does not support hash-routed MoE layers.") - if any( - getattr(model.config, option, False) - for option in ("enable_mhc_connections", "wide_residual", "moe_shortcut_connection") - ): - raise ValueError( - "HybridStack EP overlap does not support mHC, wide residuals, or MoE shortcuts." - ) - # The schedule plan calls the layer callables directly and bypasses - # ``HybridStack.forward``, which is where per-layer context-parallel layout - # conversion happens; mixed linear/attention CP layouts are therefore unsupported. - # A group's outer layout is the boundary layout even when its inner - # layers need conversion, so inspect the nested stacks as well. - for module in model.modules(): - cp_layout_manager = getattr(module, "_cp_layout_manager", None) - assert cp_layout_manager is None or not cp_layout_manager.requires_conversion, ( - "EP A2A overlap with HybridStack does not support mixed context-parallel layouts " - "(linear_cp_layout != attention_cp_layout with context_parallel_size > 1)." - ) - super().__init__(model, *args, **kwargs) - def _extra_args_for_layer(self, module, layer_idx, num_layers): extra_args = super()._extra_args_for_layer(module, layer_idx, num_layers) extra_args["layer_type"] = ( diff --git a/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py index d13def4f44c..19f04d8766a 100644 --- a/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py +++ b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py @@ -1,23 +1,17 @@ # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. from types import SimpleNamespace -from unittest.mock import patch import pytest import torch from megatron.core.context_parallel.layout import ContextParallelLayoutManager -from megatron.core.models.common.model_chunk_schedule_plan import ( - TransformerLayerSchedulePlan, - TransformerModelChunkSchedulePlan, -) +from megatron.core.models.common.model_chunk_schedule_plan import TransformerLayerSchedulePlan from megatron.core.models.hybrid.hybrid_block import HybridStack from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec -from megatron.core.models.hybrid.model_chunk_schedule_plan import ( - HybridStackModelChunkSchedulePlan, - HybridStackSchedulePlan, -) +from megatron.core.models.hybrid.hybrid_model import HybridModel +from megatron.core.models.hybrid.model_chunk_schedule_plan import HybridStackSchedulePlan from megatron.core.pipeline_parallel.utils import get_comm_stream, get_comp_stream, set_streams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed @@ -27,9 +21,27 @@ from tests.unit_tests.test_utilities import Utils +def _hybrid_stack_stub(cp_layout_manager=None): + stack = HybridStack.__new__(HybridStack) + torch.nn.Module.__init__(stack) + stack._cp_layout_manager = cp_layout_manager + return stack + + +def _overlap_model_stub(**config_overrides): + config = dict( + cuda_graph_impl="none", enable_mhc_connections=False, moe_shortcut_connection=False + ) + config.update(config_overrides) + model = torch.nn.Module() + model.config = SimpleNamespace(**config) + model.decoder = _hybrid_stack_stub() + return model + + @pytest.mark.parametrize("location", ["decoder", "group", "mtp"]) @pytest.mark.parametrize("needs_conversion", [True, False]) -def test_hybrid_schedule_checks_nested_cp_layouts(location, needs_conversion): +def test_hybrid_ep_overlap_checks_nested_cp_layouts(location, needs_conversion): """An outer group's boundary layout can hide an inner attention conversion.""" config = AttentionLayerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) config.attention_cp_layout = "zigzag" if needs_conversion else "contiguous" @@ -48,43 +60,38 @@ def layout_manager(layer_configs): tp_cp_group=None, ) - model = torch.nn.Module() - model.config = SimpleNamespace(cuda_graph_impl="none") - model.decoder = torch.nn.Module() + model = _overlap_model_stub() model.decoder._cp_layout_manager = layout_manager([(config,)]) assert not model.decoder._cp_layout_manager.requires_conversion if location == "group": - model.decoder.group = torch.nn.Module() + model.decoder.group = _hybrid_stack_stub() target = model.decoder.group elif location == "mtp": - model.mtp = torch.nn.Module() + model.mtp = _hybrid_stack_stub() target = model.mtp else: target = model.decoder target._cp_layout_manager = layout_manager([config]) - with patch.object(TransformerModelChunkSchedulePlan, "__init__", return_value=None) as build: - if needs_conversion: - with pytest.raises(AssertionError, match="mixed context-parallel layouts"): - HybridStackModelChunkSchedulePlan(model) - build.assert_not_called() - else: - HybridStackModelChunkSchedulePlan(model) - build.assert_called_once() + if needs_conversion: + with pytest.raises(ValueError, match="mixed context-parallel layouts"): + HybridModel._validate_ep_overlap_support(model) + else: + HybridModel._validate_ep_overlap_support(model) @pytest.mark.parametrize( - "option", - ["moe_num_hash_layers", "enable_mhc_connections", "wide_residual", "moe_shortcut_connection"], + "config_overrides, message", + [ + (dict(cuda_graph_impl="full_iteration"), "does not support CUDA graphs"), + (dict(enable_mhc_connections=True), "does not support enable_mhc_connections"), + (dict(moe_shortcut_connection=True), "moe_shortcut_connection"), + ], ) -def test_hybrid_schedule_rejects_unsupported_layer_execution(option): - """Direct overlap callables must not bypass token routing or residual wrappers.""" - model = torch.nn.Module() - model.config = SimpleNamespace(cuda_graph_impl="none", **{option: 1}) - with patch.object(TransformerModelChunkSchedulePlan, "__init__", return_value=None) as build: - with pytest.raises(ValueError, match="HybridStack EP overlap does not support"): - HybridStackModelChunkSchedulePlan(model) - build.assert_not_called() +def test_hybrid_ep_overlap_rejects_unsupported_features(config_overrides, message): + """Direct overlap callables must not bypass CUDA graphs or residual wrappers.""" + with pytest.raises(ValueError, match=message): + HybridModel._validate_ep_overlap_support(_overlap_model_stub(**config_overrides)) @pytest.mark.parametrize("pattern", ["[*E]", "[+E]"]) From c7a69dd4f5700f214b34f067b78cdddd912251e9 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Wed, 30 Sep 2026 15:17:27 -0700 Subject: [PATCH 08/11] Drop attribute probes and optional kwargs from hybrid overlap code Read known attributes directly and pass arguments unconditionally: - HybridStack sets final_norm to None when it has no final norm (like TransformerBlock.final_layernorm), so the final_layernorm alias and the overlap callables read it directly. The flextron hook checks for a None final norm instead of a missing attribute. - The schedule plan dispatches on type: HybridStack for the layer-type symbol, and TransformerLayer / MultiTokenPredictionLayer for the layer-level quantization context. Its imports move to module level. - The MoE preprocess slot always sets shared_expert_output; the combine slot reads it and mlp_norm_manager directly. The Mamba pre-layer call no longer probes the chunk state for an inference context the training schedule never sets. - HybridStack.forward drops rotary_pos_cos / rotary_pos_sin / rotary_pos_cos_sin, which nothing passes in, and HybridModel._postprocess drops is_spec_decode, which no caller passes. - A bracketed group stack never applies full recomputation itself; gate on the static is_layer_group_stack instead of threading the private _checkpointed_forward_in_parent flag through forward. checkpointed_forward identifies group layers with isinstance(layer, HybridStack) and passes input_ids, packed_seq_params_by_layout and cp_layout_plan unconditionally. - MambaLayer.backward_dw and MambaMixer.backward_dw call their delayed wgrad unconditionally. Building the overlap callables with delay_wgrad_compute now raises for a Mamba layer whose mixer cannot defer its wgrad (e.g. GatedDeltaProductMixer) instead of silently skipping those weight gradients. Signed-off-by: Yan Xu --- .../models/hybrid/fine_grained_callables.py | 17 ++++++--- megatron/core/models/hybrid/hybrid_block.py | 31 ++++----------- megatron/core/models/hybrid/hybrid_model.py | 14 +++---- .../hybrid/model_chunk_schedule_plan.py | 38 ++++++++++--------- megatron/core/recompute.py | 27 +++++++------ megatron/core/ssm/mamba_layer.py | 12 +++--- megatron/core/ssm/mamba_mixer.py | 11 ++---- .../flextron_elasticity_hooks.py | 2 +- .../test_schedule_quantization_context.py | 6 ++- .../test_hybrid_fine_grained_callables.py | 18 +++++++++ 10 files changed, 93 insertions(+), 83 deletions(-) diff --git a/megatron/core/models/hybrid/fine_grained_callables.py b/megatron/core/models/hybrid/fine_grained_callables.py index af2ccf7bb98..e2e3103c5dd 100644 --- a/megatron/core/models/hybrid/fine_grained_callables.py +++ b/megatron/core/models/hybrid/fine_grained_callables.py @@ -20,6 +20,7 @@ StageDispatchBwdGrad, get_comm_stream, ) +from megatron.core.ssm.mamba_mixer import MambaMixer from megatron.core.transformer.transformer_layer import make_viewless_tensor @@ -138,8 +139,7 @@ def get_hybrid_stack_moe_metadata(layer, layer_type: Optional[LayerPatternItem] def _maybe_apply_final_norm(node: ScheduleNode, hidden_states: Tensor): - final_norm = getattr(node.chunk_state.model.decoder, "final_norm", None) - final_norm = final_norm or getattr(node.chunk_state.model.decoder, "final_layernorm", None) + final_norm = node.chunk_state.model.decoder.final_norm if not node.is_mtp and final_norm is not None and node.is_last_layer: hidden_states = final_norm(hidden_states) hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) @@ -177,6 +177,7 @@ def _run_moe_preprocess(layer, node: ScheduleNode, hidden_states: Tensor): local_tokens, probs = layer.mlp.preprocess(pre_mlp_layernorm_output, probs, routing_map) node.layer_state.residual = node.detach(residual) + node.layer_state.shared_expert_output = None if layer.mlp.use_shared_expert and not layer.mlp.shared_expert_overlap: node.layer_state.shared_expert_output = node.detach(shared_expert_output) @@ -215,7 +216,7 @@ def _run_moe_experts(layer, node: ScheduleNode, dispatched_tokens: Tensor): def _run_moe_combine(layer, node: ScheduleNode, output: Tensor): residual = node.layer_state.residual - shared_expert_output = getattr(node.layer_state, 'shared_expert_output', None) + shared_expert_output = node.layer_state.shared_expert_output output = layer.mlp.combine(output) output = layer.mlp.postprocess(output, shared_expert_output) # Inline bda instead of calling ``layer._forward_post_mlp``, which would register @@ -229,7 +230,7 @@ def _run_moe_combine(layer, node: ScheduleNode, output: Tensor): output = layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( mlp_output_with_bias, residual, layer.hidden_dropout ) - mlp_norm_manager = getattr(node.layer_state, "mlp_norm_manager", None) + mlp_norm_manager = node.layer_state.mlp_norm_manager if mlp_norm_manager is not None: output = mlp_norm_manager.group_offload(output, forced_released_tensors=[residual]) node.layer_state.mlp_norm_manager = None @@ -263,7 +264,6 @@ def pre_dispatch_computation(node: ScheduleNode, hidden_states: Tensor): hidden_states = item_layer( hidden_states=hidden_states, attention_mask=node.chunk_state.attention_mask, - inference_context=getattr(node.chunk_state, "inference_context", None), packed_seq_params=node.chunk_state.packed_seq_params, ) elif item_type in ( @@ -380,6 +380,13 @@ def raise_not_implemented(*args): # in turn calls backward_dw on the in_proj / out_proj linears. The # schedule node iterates this list and calls .backward_dw() on each; # registering the layer directly is sufficient. + if item_layer.config.delay_wgrad_compute and not isinstance( + item_layer.mixer, MambaMixer + ): + raise ValueError( + "delay_wgrad_compute with overlap_moe_expert_parallel_comm does not support " + f"Mamba layers with a {type(item_layer.mixer).__name__} mixer." + ) pre_bwd_dw.append(item_layer) if is_moe: # Each slot owns the hooks for exactly the parameters whose delayed wgrad diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index d3be1ef3a42..79bbe5ce10e 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -477,6 +477,8 @@ def __init__( hidden_size=self.config.hidden_size, eps=self.config.layernorm_epsilon, ) + else: + self.final_norm = None if self.config.enable_mhc_connections and self.post_process and not self.is_mtp_layer: hc_mult = self.config.mhc_num_residual_streams @@ -556,10 +558,10 @@ def final_layernorm(self): """Alias for ``final_norm`` matching the attribute name on TransformerBlock. Lets generic decoder consumers (e.g. ``GPTModel.PostProcessNode``) discover the - final norm via the same attribute name they use for non-hybrid decoders, while - keeping ``final_norm`` as the registered submodule for local state-dict compatibility. + final norm via the same attribute name they use for non-hybrid decoders. + ``final_norm`` remains the registered submodule name. """ - return getattr(self, "final_norm", None) + return self.final_norm @staticmethod def _get_layer_cp_layout(layer_config: LayerConfigItem, boundary_layout: CPLayout) -> CPLayout: @@ -668,9 +670,6 @@ def forward( attention_mask: Tensor, inference_context: Optional[BaseInferenceContext] = None, rotary_pos_emb: Optional[Tensor] = None, - rotary_pos_cos: Optional[Tensor] = None, - rotary_pos_sin: Optional[Tensor] = None, - rotary_pos_cos_sin: Optional[Tensor] = None, sequence_len_offset: Optional[Tensor] = None, *, inference_params: Optional[BaseInferenceContext] = None, @@ -679,7 +678,6 @@ def forward( packed_seq_params_by_layout: dict[CPLayout, PackedSeqParams | None] | None = None, cp_layout_plan: THDCPLayoutPlan | None = None, input_ids: Optional[Tensor] = None, - _checkpointed_forward_in_parent: bool = False, ): """ Forward function of the HybridStack class. @@ -697,14 +695,8 @@ def forward( Defaults to None. input_ids (Tensor, optional): Token IDs forwarded to hash-routed TransformerLayer instances. Defaults to None. - rotary_pos_cos / rotary_pos_sin / rotary_pos_cos_sin (Tensor, optional): - precomputed rotary embeddings forwarded to transformer layers (flash-decode / - fused-rope inference paths). Defaults to None. sequence_len_offset (Tensor, optional): precomputed per-sample sequence offsets for static-batching inference. Computed here when None. - _checkpointed_forward_in_parent (bool): set by ``checkpointed_forward`` when the - enclosing stack already checkpoints this nested group stack, so the group - must not checkpoint its layers a second time. Returns: Tensor: the output tensor. """ @@ -822,10 +814,12 @@ def get_inner_quant_context(config, layer_number): ) with outer_fp8_context: + # A bracketed group stack runs inside its enclosing stack's checkpointed + # segment, so only the outer stack applies full recomputation. if ( self.config.recompute_granularity == 'full' and self.training - and not _checkpointed_forward_in_parent + and not self.is_layer_group_stack ): hidden_states = checkpointed_forward( self, @@ -908,9 +902,6 @@ def get_inner_quant_context(config, layer_number): attention_mask=attention_mask, inference_context=inference_context, rotary_pos_emb=rotary_pos_emb, - rotary_pos_cos=rotary_pos_cos, - rotary_pos_sin=rotary_pos_sin, - rotary_pos_cos_sin=rotary_pos_cos_sin, sequence_len_offset=sequence_len_offset, packed_seq_params=layer_packed_seq_params, padding_mask=padding_mask, @@ -928,12 +919,6 @@ def get_inner_quant_context(config, layer_number): packed_seq_params=layer_packed_seq_params, padding_mask=padding_mask, ) - if rotary_pos_cos is not None: - layer_kwargs["rotary_pos_cos"] = rotary_pos_cos - if rotary_pos_sin is not None: - layer_kwargs["rotary_pos_sin"] = rotary_pos_sin - if rotary_pos_cos_sin is not None: - layer_kwargs["rotary_pos_cos_sin"] = rotary_pos_cos_sin if layer_cp_metadata is not None: layer_kwargs["packed_sequence_cp_metadata"] = layer_cp_metadata if residual_stream_recompute_context is not None: diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 123755f60ba..3aa9763f210 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -725,7 +725,6 @@ def _postprocess( runtime_gather_output=None, extra_block_kwargs=None, inference_context=None, - is_spec_decode=None, output_processor=None, output_processor_context=None, compute_mtp_loss=True, @@ -758,13 +757,12 @@ def _postprocess( # Speculative decoding: when active, MTP must run *after* verification so # it conditions on verified tokens rather than stale speculative ones. - if is_spec_decode is None: - is_spec_decode = ( - in_inference_mode - and inference_context is not None - and inference_context.is_dynamic_batching() - and inference_context.num_speculative_tokens > 0 - ) + is_spec_decode = ( + in_inference_mode + and inference_context is not None + and inference_context.is_dynamic_batching() + and inference_context.num_speculative_tokens > 0 + ) # Whether the MTP auxiliary objective is computed in this call. # ``self.mtp_process`` guards against models built without an MTP block. diff --git a/megatron/core/models/hybrid/model_chunk_schedule_plan.py b/megatron/core/models/hybrid/model_chunk_schedule_plan.py index 3aebbc4590b..fe51591e71d 100644 --- a/megatron/core/models/hybrid/model_chunk_schedule_plan.py +++ b/megatron/core/models/hybrid/model_chunk_schedule_plan.py @@ -19,6 +19,14 @@ TransformerLayerSchedulePlan, TransformerModelChunkSchedulePlan, ) +from megatron.core.models.hybrid.fine_grained_callables import ( + HybridStackNode, + build_hybrid_stack_callables, +) +from megatron.core.models.hybrid.hybrid_block import HybridStack +from megatron.core.pipeline_parallel.utils import NoopScheduleNode +from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer +from megatron.core.transformer.transformer_layer import TransformerLayer class HybridStackSchedulePlan(TransformerLayerSchedulePlan): @@ -40,14 +48,6 @@ def _build_callable_nodes(self, event, comp_stream, comm_stream, extra_args): if self.layer_type is None: return super()._build_callable_nodes(event, comp_stream, comm_stream, extra_args) - # Hybrid grouped path. Imports are local because hybrid pulls in TE / SSM - # extensions that we don't want to load when only the GPT path is used. - from megatron.core.models.hybrid.fine_grained_callables import ( - HybridStackNode, - build_hybrid_stack_callables, - ) - from megatron.core.pipeline_parallel.utils import NoopScheduleNode - fwd_callables, bwd_dw_callable_map, is_moe, num_local_experts = ( build_hybrid_stack_callables(self.layer, layer_type=self.layer_type) ) @@ -98,12 +98,16 @@ def create_node(stream, module, name): self.mtp_post_process = NoopScheduleNode() def get_low_precision_context(self): - """Let hybrid callables manage each physical layer's quantization context.""" - # HybridStack and MambaLayer do not expose the TransformerLayer context - # hook. Their callables enter the appropriate context for each inner layer. - if self.layer_type is not None or not hasattr(self.layer, "get_inner_quantization_context"): - return nullcontext() - return super().get_low_precision_context() + """Return the layer-level quantization context for GPT-path layers. + + Hybrid callables enter the quantization context of each physical layer + themselves, so hybrid layer plans use a null context here. + """ + if self.layer_type is None and isinstance( + self.layer, (TransformerLayer, MultiTokenPredictionLayer) + ): + return super().get_low_precision_context() + return nullcontext() class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): @@ -111,8 +115,8 @@ class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): Threads HybridStack's ``layer_type_list[layer_idx]`` symbol into each layer plan's ``extra_args`` so the per-layer plan can dispatch grouped - layers correctly. Ordinary GPT/MTP layers (no ``layer_type_list``) - default to ``layer_type=None`` and follow the GPT path. The pre/post + layers correctly. Layers of other modules (e.g. MTP layers) get + ``layer_type=None`` and follow the GPT path. The pre/post process nodes inherit from the GPT base class — they already dispatch on ``model._preprocess`` / ``model._postprocess`` which a HybridModel implements. @@ -123,6 +127,6 @@ class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): def _extra_args_for_layer(self, module, layer_idx, num_layers): extra_args = super()._extra_args_for_layer(module, layer_idx, num_layers) extra_args["layer_type"] = ( - module.layer_type_list[layer_idx] if hasattr(module, "layer_type_list") else None + module.layer_type_list[layer_idx] if isinstance(module, HybridStack) else None ) return extra_args diff --git a/megatron/core/recompute.py b/megatron/core/recompute.py index a8c73785113..5a7162b50ac 100644 --- a/megatron/core/recompute.py +++ b/megatron/core/recompute.py @@ -62,6 +62,9 @@ def checkpointed_forward( If extract_layer_indices is empty: hidden_states tensor If extract_layer_indices is non-empty: (hidden_states, intermediate_hidden_states) tuple """ + # Imported here because hybrid_block imports this module. + from megatron.core.models.hybrid.hybrid_block import HybridStack + if extract_layer_indices is None: extract_layer_indices = set() intermediate_hidden_states: List[Tensor] = [] @@ -99,7 +102,8 @@ def custom_forward( ) # Keep both residuals in the layer's layout, inside the CP conversions. residual_accumulator = hidden_states - is_hybrid_group = getattr(layer, "is_layer_group_stack", False) + # A HybridStack layer is always a nested bracketed group. + is_hybrid_group = isinstance(layer, HybridStack) # Get appropriate inner quantization context if use_inner_quantization_context and not is_hybrid_group: @@ -145,19 +149,18 @@ def custom_forward( layer_kwargs.pop(k, None) hidden_states, context = layer(**layer_kwargs) elif is_hybrid_group: - # Nested HybridStack group: run its physical layers inside this - # checkpoint segment (it must not checkpoint them again) and let it - # build its own CP layout state from the prebuilt per-layout metadata. + # Nested HybridStack group: its physical layers run inside this + # checkpoint segment (a group stack never checkpoints them again), + # and it builds its own CP layout state from the prebuilt + # per-layout metadata. for k in ("context", "context_mask", "attention_bias"): layer_kwargs.pop(k, None) - if input_ids is not None: - layer_kwargs["input_ids"] = input_ids - if packed_seq_params_by_layout is not None or cp_layout_plan is not None: - layer_kwargs["packed_seq_params_by_layout"] = ( - packed_seq_params_by_layout - ) - layer_kwargs["cp_layout_plan"] = cp_layout_plan - hidden_states = layer(**layer_kwargs, _checkpointed_forward_in_parent=True) + hidden_states = layer( + **layer_kwargs, + input_ids=input_ids, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, + ) context = None else: # MambaLayer (HybridStack `M` slot) for k in ( diff --git a/megatron/core/ssm/mamba_layer.py b/megatron/core/ssm/mamba_layer.py index 61ee94249a5..76845d3d309 100644 --- a/megatron/core/ssm/mamba_layer.py +++ b/megatron/core/ssm/mamba_layer.py @@ -320,15 +320,13 @@ def forward( return hidden_states def backward_dw(self): - """Compute weight gradients for the layer's linear projections. + """Compute the delayed weight gradients of the mixer's linear projections. - Delegates to the mixer; lets the hybrid EP-overlap schedule plan - register a Mamba pre-layer's wgrad alongside attention/GDN pre-layers - so the schedule node iterates a uniform set of callables. No-op when - the linears in the spec do not support delayed wgrad. + Lets the hybrid EP-overlap schedule plan register a Mamba pre-layer's wgrad + alongside attention/GDN pre-layers, so the schedule node iterates a uniform set + of callables. """ - if hasattr(self.mixer, "backward_dw"): - self.mixer.backward_dw() + self.mixer.backward_dw() def sharded_state_dict( self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None diff --git a/megatron/core/ssm/mamba_mixer.py b/megatron/core/ssm/mamba_mixer.py index 729e9e04db1..16604d445c6 100644 --- a/megatron/core/ssm/mamba_mixer.py +++ b/megatron/core/ssm/mamba_mixer.py @@ -1330,15 +1330,10 @@ def backward_dw(self): Mirrors ``GatedDeltaNet.backward_dw``. The selective-scan kernel is a single autograd function whose wgrad runs in the regular backward pass, - so only the input/output projections need delayed wgrad here. Each - ``backward_dw`` call is a no-op unless the underlying linear is built - from a TE primitive that supports delayed wgrad; if the spec uses - non-TE linears, ``backward_dw`` simply does nothing. + so only the input/output projections need delayed wgrad here. """ - if hasattr(self.in_proj, "backward_dw"): - self.in_proj.backward_dw() - if hasattr(self.out_proj, "backward_dw"): - self.out_proj.backward_dw() + self.in_proj.backward_dw() + self.out_proj.backward_dw() def _get_states_from_cache(self, inference_context, batch_size, *, inference_params=None): """Initializes or retrieves the SSM state tensors from the cache. diff --git a/megatron/elastification/flextron_elasticity_hooks.py b/megatron/elastification/flextron_elasticity_hooks.py index bb379916ede..606239b329e 100644 --- a/megatron/elastification/flextron_elasticity_hooks.py +++ b/megatron/elastification/flextron_elasticity_hooks.py @@ -1862,7 +1862,7 @@ def apply_flextron_elasticity_to_model(model, config): managers.append(manager) # Also add hooks to HybridStack if present - if hasattr(model, 'decoder') and hasattr(model.decoder, 'final_norm'): + if hasattr(model, 'decoder') and getattr(model.decoder, 'final_norm', None) is not None: stack_manager = add_flextron_stack_elasticity(model.decoder, config) managers.append(stack_manager) diff --git a/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py b/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py index 19e1032f28f..f8931aab765 100644 --- a/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py +++ b/tests/unit_tests/a2a_overlap/test_schedule_quantization_context.py @@ -63,9 +63,11 @@ def test_hybrid_schedule_runs_without_a_transformer_context_hook(layer_type): ] -def test_hybrid_schedule_preserves_plain_layer_quantization_context(): +@pytest.mark.parametrize("layer_cls", [TransformerLayer, MultiTokenPredictionLayer]) +def test_hybrid_schedule_preserves_plain_layer_quantization_context(layer_cls): expected_context = nullcontext() - layer = SimpleNamespace(get_inner_quantization_context=Mock(return_value=expected_context)) + layer = Mock(spec=layer_cls) + layer.get_inner_quantization_context.return_value = expected_context plan = HybridStackSchedulePlan.__new__(HybridStackSchedulePlan) plan.layer = layer plan.layer_type = None diff --git a/tests/unit_tests/models/test_hybrid_fine_grained_callables.py b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py index a909b765ade..c84c46afeb1 100644 --- a/tests/unit_tests/models/test_hybrid_fine_grained_callables.py +++ b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py @@ -12,6 +12,7 @@ import megatron.core.models.hybrid.fine_grained_callables as hybrid_callables import megatron.core.pipeline_parallel.utils as schedule_utils from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols +from megatron.core.ssm.mamba_mixer import MambaMixer from megatron.core.transformer.moe.moe_layer import MoELayer @@ -176,6 +177,23 @@ def test_attention_half_layer_forward_and_wgrad_are_scheduled(symbol): assert not is_moe +@pytest.mark.parametrize("delay_wgrad_compute", [False, True]) +@pytest.mark.parametrize("mamba_mixer", [False, True]) +def test_mamba_pre_layer_wgrad_requires_a_mixer_that_defers_it(delay_wgrad_compute, mamba_mixer): + """A mixer without delayed wgrad must fail when built, not skip its weight gradients.""" + mixer = MambaMixer.__new__(MambaMixer) if mamba_mixer else SimpleNamespace() + layer = SimpleNamespace(config=_config(delay_wgrad_compute=delay_wgrad_compute), mixer=mixer) + if delay_wgrad_compute and not mamba_mixer: + with pytest.raises(ValueError, match="SimpleNamespace mixer"): + hybrid_callables.build_hybrid_stack_callables(layer, Symbols.MAMBA) + else: + _, backward_dw, is_moe, _ = hybrid_callables.build_hybrid_stack_callables( + layer, Symbols.MAMBA + ) + assert backward_dw["pre_dispatch_computation"] == [layer] + assert not is_moe + + @pytest.mark.parametrize("zero_copy", [False, True]) def test_ncclep_probabilities_do_not_reconnect_schedule_graphs(cpu_streams, zero_copy): """Backward across all three real schedule nodes must match the unsplit computation.""" From d3ca30f981babc0f9995d599218ca6e7d241b183 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Mon, 5 Oct 2026 10:27:08 -0700 Subject: [PATCH 09/11] Pass the batch-first padding mask to hybrid EP-overlap MoE routing MoELayer.route takes the batch-first [b, s] padding mask and transposes it for the router, and MoELayer.preprocess takes the same mask so that dropless HybridEP drops padded rows before dispatch. The hybrid overlap callable transposed the mask before calling route, so the router flattened a [b, s] mask against sequence-first tokens, which misaligns the mask whenever the micro-batch size is greater than 1. It also never passed the mask to preprocess. Pass node.chunk_state.padding_mask to route and preprocess unchanged, as MoELayer.forward and the GPT overlap callables do. Signed-off-by: Yan Xu Co-Authored-By: Claude Opus 5.5 --- .../models/hybrid/fine_grained_callables.py | 17 ++++----- .../test_hybrid_fine_grained_callables.py | 35 ++++++++++++++++++- 2 files changed, 41 insertions(+), 11 deletions(-) diff --git a/megatron/core/models/hybrid/fine_grained_callables.py b/megatron/core/models/hybrid/fine_grained_callables.py index e2e3103c5dd..b1ca8a2f492 100644 --- a/megatron/core/models/hybrid/fine_grained_callables.py +++ b/megatron/core/models/hybrid/fine_grained_callables.py @@ -146,14 +146,6 @@ def _maybe_apply_final_norm(node: ScheduleNode, hidden_states: Tensor): return hidden_states -def _get_moe_padding_mask(node: ScheduleNode): - padding_mask = node.chunk_state.padding_mask - if padding_mask is not None: - # MoELayer.forward receives [batch, seq] and transposes before routing. - padding_mask = padding_mask.transpose(0, 1).bool() - return padding_mask - - def _run_moe_preprocess(layer, node: ScheduleNode, hidden_states: Tensor): pre_mlp_layernorm_output = layer._forward_pre_mlp_layernorm(hidden_states) # A subsequent microbatch can reuse this layer before this node combines. @@ -173,8 +165,13 @@ def _run_moe_preprocess(layer, node: ScheduleNode, hidden_states: Tensor): residual = residual.float() shared_expert_output = layer.mlp.shared_experts_compute(pre_mlp_layernorm_output) - probs, routing_map = layer.mlp.route(pre_mlp_layernorm_output, _get_moe_padding_mask(node)) - local_tokens, probs = layer.mlp.preprocess(pre_mlp_layernorm_output, probs, routing_map) + # Same routing inputs as the eager MoELayer.forward path: route and preprocess take the + # batch-first padding mask (dropless HybridEP excludes padded tokens). + padding_mask = node.chunk_state.padding_mask + probs, routing_map = layer.mlp.route(pre_mlp_layernorm_output, padding_mask) + local_tokens, probs = layer.mlp.preprocess( + pre_mlp_layernorm_output, probs, routing_map, padding_mask + ) node.layer_state.residual = node.detach(residual) node.layer_state.shared_expert_output = None diff --git a/tests/unit_tests/models/test_hybrid_fine_grained_callables.py b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py index c84c46afeb1..e3e38288415 100644 --- a/tests/unit_tests/models/test_hybrid_fine_grained_callables.py +++ b/tests/unit_tests/models/test_hybrid_fine_grained_callables.py @@ -250,6 +250,39 @@ def preprocess(node, hidden_states): assert dispatch_grad_ptrs == [manager._zc_bwd_token_buf.data_ptr()] +def test_moe_routing_receives_the_batch_first_padding_mask(): + """Route and preprocess get the chunk's padding mask unchanged, as in MoELayer.forward.""" + layer = _moe_layer(_config()) + layer._forward_pre_mlp_layernorm = lambda hidden_states: hidden_states + layer.mlp_norm_manager = None + layer.mlp.shared_experts_compute = lambda hidden_states: None + received = {} + + def route(hidden_states, padding_mask): + received["route"] = padding_mask + return hidden_states, None + + def preprocess(hidden_states, probs, routing_map, padding_mask): + received["preprocess"] = padding_mask + return hidden_states, probs + + layer.mlp.route = route + layer.mlp.preprocess = preprocess + # Batch-first [b, s] with b != s, so a transposed mask cannot stand in for it. + padding_mask = torch.tensor([[False, False, True], [False, True, True]]) + node = SimpleNamespace( + layer_state=SimpleNamespace(), + chunk_state=_chunk_state(), + detach=lambda tensor: tensor.detach(), + ) + node.chunk_state.padding_mask = padding_mask + + hybrid_callables._run_moe_preprocess(layer, node, torch.ones(3, 2, 1)) + + assert received["route"] is padding_mask + assert received["preprocess"] is padding_mask + + @pytest.mark.parametrize("recompute", [False, True]) def test_norm_offload_uses_its_microbatch_manager_after_bda(cpu_streams, recompute): """Two in-flight microbatches keep distinct offload managers and single recompute hooks.""" @@ -262,7 +295,7 @@ def test_norm_offload_uses_its_microbatch_manager_after_bda(cpu_streams, recompu layer.mlp_bda = lambda *args: lambda output, residual, dropout: output[0] + residual layer.mlp.shared_experts_compute = lambda hidden: None layer.mlp.route = lambda hidden, mask: (hidden, None) - layer.mlp.preprocess = lambda hidden, probs, routing: (hidden, probs) + layer.mlp.preprocess = lambda hidden, probs, routing, mask: (hidden, probs) layer.mlp.routed_experts_compute = lambda hidden, probs: (hidden * 3, None) layer.mlp.combine = lambda output: output layer.mlp.postprocess = lambda output, shared: output From a38db438e433840ea7aef3b0855761c6f4c08c5a Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 6 Oct 2026 09:15:06 -0700 Subject: [PATCH 10/11] Reject chunkwise linear CP with hybrid EP overlap The EP-overlap schedule plan calls the layer callables directly and bypasses HybridStack.forward, which builds the packed-sequence CP metadata that chunkwise linear CP needs and rejects padding masks for it. With context_parallel_size > 1 and linear_cp_mode='chunkwise', the model built and the first packed step failed inside the Gated Delta Product mixer. Reject the combination at construction for every HybridStack, including nested groups and MTP stacks, like the other features the overlap schedule cannot run. Signed-off-by: Yan Xu Co-Authored-By: Claude Opus 5.5 --- megatron/core/models/hybrid/hybrid_model.py | 17 +++++++++++----- .../a2a_overlap/test_hybrid_schedule_plan.py | 20 +++++++++++++++++++ 2 files changed, 32 insertions(+), 5 deletions(-) diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 3aa9763f210..2e1199d7560 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -499,13 +499,15 @@ def _validate_ep_overlap_support(self) -> None: ) # The schedule plan calls the layer callables directly and bypasses # ``HybridStack.forward``, which is where per-layer context-parallel layout - # conversion happens. A bracketed group presents the boundary layout to its - # enclosing stack even when its own layers need conversion, so check every - # stack, including the nested ones. + # conversion happens and where chunkwise linear CP gets its packed-sequence + # metadata and padding-mask check. A bracketed group presents the boundary + # layout to its enclosing stack even when its own layers need conversion, so + # check every stack, including the nested ones. for module in self.modules(): + if not isinstance(module, HybridStack): + continue if ( - isinstance(module, HybridStack) - and module._cp_layout_manager is not None + module._cp_layout_manager is not None and module._cp_layout_manager.requires_conversion ): raise ValueError( @@ -513,6 +515,11 @@ def _validate_ep_overlap_support(self) -> None: "context-parallel layouts (linear_cp_layout != attention_cp_layout with " "context_parallel_size > 1)." ) + if module._has_linear_layer_with_chunkwise_cp: + raise ValueError( + "overlap_moe_expert_parallel_comm with HybridModel does not support " + "linear_cp_mode='chunkwise' with context_parallel_size > 1." + ) def set_input_tensor(self, input_tensor: Tensor) -> None: """Sets input tensor to the model. diff --git a/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py index 19f04d8766a..823ad1812b9 100644 --- a/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py +++ b/tests/unit_tests/a2a_overlap/test_hybrid_schedule_plan.py @@ -25,6 +25,7 @@ def _hybrid_stack_stub(cp_layout_manager=None): stack = HybridStack.__new__(HybridStack) torch.nn.Module.__init__(stack) stack._cp_layout_manager = cp_layout_manager + stack._has_linear_layer_with_chunkwise_cp = False return stack @@ -80,6 +81,25 @@ def layout_manager(layer_configs): HybridModel._validate_ep_overlap_support(model) +@pytest.mark.parametrize("location", ["decoder", "group", "mtp"]) +def test_hybrid_ep_overlap_rejects_chunkwise_linear_cp(location): + """Chunkwise linear CP needs packed-sequence metadata that only HybridStack.forward builds.""" + model = _overlap_model_stub() + HybridModel._validate_ep_overlap_support(model) + if location == "group": + model.decoder.group = _hybrid_stack_stub() + target = model.decoder.group + elif location == "mtp": + model.mtp = _hybrid_stack_stub() + target = model.mtp + else: + target = model.decoder + target._has_linear_layer_with_chunkwise_cp = True + + with pytest.raises(ValueError, match="linear_cp_mode='chunkwise'"): + HybridModel._validate_ep_overlap_support(model) + + @pytest.mark.parametrize( "config_overrides, message", [ From 089b33b8104aa3ddf10d97e601ab03eaf6836ae2 Mon Sep 17 00:00:00 2001 From: Yan Xu Date: Tue, 6 Oct 2026 21:26:27 -0700 Subject: [PATCH 11/11] Document hybrid layer groups and move their namespace check to layers - Import HybridStack and the hybrid callable builders at module level in common/fine_grained_callables.py; there is no import cycle. - Document bracketed layer groups in docs/user-guide/hybrid-model-migration.md, and point Symbols.GROUP_START, the --hybrid-layer-pattern help and the HybridModel docstring to it. - Add a module docstring to hybrid/fine_grained_callables.py and name GPTModel explicitly in the hybrid schedule-plan and callable docstrings. - Move the layer-symbol to checkpoint-namespace map and the duplicate- namespace check for layer groups into hybrid/layers/utils.py. Signed-off-by: Yan Xu Co-Authored-By: Claude Opus 5.5 --- docs/user-guide/hybrid-model-migration.md | 30 +++++++++++++- .../models/common/fine_grained_callables.py | 11 +++-- .../models/hybrid/fine_grained_callables.py | 22 +++++++++- .../models/hybrid/hybrid_layer_allocation.py | 28 ++----------- megatron/core/models/hybrid/hybrid_model.py | 5 ++- megatron/core/models/hybrid/layers/utils.py | 41 ++++++++++++++++++- .../hybrid/model_chunk_schedule_plan.py | 21 +++++----- megatron/training/arguments.py | 4 +- 8 files changed, 116 insertions(+), 46 deletions(-) diff --git a/docs/user-guide/hybrid-model-migration.md b/docs/user-guide/hybrid-model-migration.md index 8f9d02bf18f..73aefab5889 100644 --- a/docs/user-guide/hybrid-model-migration.md +++ b/docs/user-guide/hybrid-model-migration.md @@ -46,7 +46,8 @@ the first `-` or `E`. Source GPT layer 1 maps to the next pair, and so on. The pattern can also describe execution layout. A `|` marks a pipeline segment boundary, and `/` introduces a repeated Multi-Token Prediction (MTP) pattern. For example, `*-*-|*-*-` places four GPT-equivalent blocks across two pipeline -segments. Separators do not count as layers. +segments. Separators do not count as layers. Square brackets group layers into +one logical layer, as described in [Bracketed layer groups](#bracketed-layer-groups). HybridModel provides the following benefits: @@ -68,6 +69,33 @@ An architecture-preserving `*-` or `*E` migration should be validated for numerical equivalence, and a pattern that adds another layer family should be treated as a new architecture and benchmarked independently. +### Bracketed layer groups + +Square brackets group consecutive layers into one logical layer, for example +`[*E][*E]` or `M[M*E]`. HybridModel builds each group as one nested +`HybridStack`, and the group occupies a single layer index in checkpoint keys: +`[*-]` stores its attention and MLP under one `layers..` prefix, as a GPT +transformer block does. Adding or removing brackets therefore changes the +checkpoint keys of the affected layers. Each symbol inside a group still counts +as one layer for layer numbering, so layer-indexed settings such as FP8 layer +ranges and hash-routed MoE thresholds are unchanged. + +Use groups with MoE expert-parallel communication overlap +(`--overlap-moe-expert-parallel-comm`). The overlap schedule interleaves one +microbatch's MoE all-to-all with another microbatch's compute one logical layer +at a time, so group each MoE layer with the layers that run before it, as in +`[M*E]`. + +Groups have the following constraints: + +- A group cannot be empty or nested, cannot span a `|` pipeline boundary, and + cannot appear in an MTP pattern. +- An MoE layer must be the last layer of its group. +- The layers of a group must use different checkpoint namespaces: at most one + Mamba layer, one attention or GDN layer, and one MLP or MoE layer. +- Groups cannot be combined with mHC connections, wide residual streams, or MoE + shortcut connections. + ## 2. How to Convert a Checkpoint There are two ways to bring `GPTModel` weights into a `HybridModel` run. Both diff --git a/megatron/core/models/common/fine_grained_callables.py b/megatron/core/models/common/fine_grained_callables.py index aa83dcdd6f0..12ae5d48e3a 100644 --- a/megatron/core/models/common/fine_grained_callables.py +++ b/megatron/core/models/common/fine_grained_callables.py @@ -18,6 +18,11 @@ from megatron.core import tensor_parallel from megatron.core.models.gpt.fine_grained_callables import build_transformer_layer_callables +from megatron.core.models.hybrid.fine_grained_callables import ( + build_hybrid_stack_callables, + get_hybrid_stack_moe_metadata, +) +from megatron.core.models.hybrid.hybrid_block import HybridStack from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.multi_token_prediction import ( MultiTokenPredictionLayer, @@ -155,13 +160,10 @@ def rng_context_wrapper(func, *args, **kwargs): def get_layer_moe_metadata(layer): """Return ``(is_moe, num_local_experts)`` for schedule-node construction.""" - from megatron.core.models.hybrid.hybrid_block import HybridStack if isinstance(layer, MultiTokenPredictionLayer): return get_layer_moe_metadata(layer.mtp_model_layer) if isinstance(layer, HybridStack): - from megatron.core.models.hybrid.fine_grained_callables import get_hybrid_stack_moe_metadata - return get_hybrid_stack_moe_metadata(layer) if isinstance(layer, TransformerLayer): is_moe = isinstance(layer.mlp, MoELayer) @@ -176,13 +178,10 @@ def build_layer_callables(layer): Returns ``(forward_funcs, backward_dw)``. """ - from megatron.core.models.hybrid.hybrid_block import HybridStack if isinstance(layer, MultiTokenPredictionLayer): return build_mtp_layer_callables(layer) if isinstance(layer, HybridStack): - from megatron.core.models.hybrid.fine_grained_callables import build_hybrid_stack_callables - forward_funcs, backward_dw, _, _ = build_hybrid_stack_callables(layer) return forward_funcs, backward_dw if isinstance(layer, TransformerLayer): diff --git a/megatron/core/models/hybrid/fine_grained_callables.py b/megatron/core/models/hybrid/fine_grained_callables.py index b1ca8a2f492..0d4b8517c61 100644 --- a/megatron/core/models/hybrid/fine_grained_callables.py +++ b/megatron/core/models/hybrid/fine_grained_callables.py @@ -1,5 +1,23 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +"""Layer callables for HybridModel's EP-overlap (combined 1F1B) schedule plan. + +``build_hybrid_stack_callables`` splits one logical HybridStack layer, either a bracketed +group such as ``[M*E]`` or a single layer symbol, into the schedule plan's slots: + +* pre-dispatch compute: the Mamba, attention, MLA or GDN layers that precede the + terminal MLP/MoE layer, then the MoE layer's pre-MLP norm, shared experts, router + and dispatch preprocessing; +* token dispatch; +* the dense MLP or the routed experts; +* token combine and the MoE residual add, plus the final norm on the last layer. + +``HybridStackNode`` is the schedule node that runs these callables, and +``_MoEBackwardDWWrapper`` gives each slot the delayed MoE weight-gradient work it +owns. The equivalent callables for GPTModel's ``TransformerLayer`` are in +``megatron/core/models/gpt/fine_grained_callables.py``. +""" + from contextlib import nullcontext from functools import partial from typing import Optional @@ -80,7 +98,7 @@ class HybridStackNode(TransformerLayerNode): def _resolve_free_input(name, is_moe, config, num_local_experts): """Hybrid free-input policy. - Same as the GPT default: dense layers always retain their + Same as the GPTModel default: dense layers always retain their input for backward; MoE-only "moe_dispatch", "mlp", and "moe_combine" slots can free, subject to the dispatcher / cuda-graph constraints encoded in ``should_free_input``. Hybrid groups have a @@ -221,7 +239,7 @@ def _run_moe_combine(layer, node: ScheduleNode, output: Tensor): # already registered on ``expert_output`` in ``_run_moe_experts``. A second one # would run the layernorm recompute before attention's backward in the same # pre-dispatch slot of a bracketed group (``[*E]``), which corrupts the attention - # gradients. GPT's ``submodule_combine_forward`` inlines bda for the same reason. + # gradients. GPTModel's ``submodule_combine_forward`` inlines bda for the same reason. mlp_output_with_bias = (output, None) with layer.bias_dropout_add_exec_handler(): output = layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index b49b1066db5..a31bb551cc6 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -72,36 +72,16 @@ def layer_type_list_to_str(layer_type_list: Sequence[LayerPatternItem]) -> str: def validate_layer_group(layer_types: Sequence[str]) -> None: - """Require group members to have distinct sharded checkpoint namespaces. + """Validate the layer symbols of one bracketed group. - A group shares one logical checkpoint layer index, so it can contain at most - one mixer, one attention module (including GDN), and one MLP or MoE module. + A group must not be empty, an MoE layer must be its last member, and its members + must use distinct sharded checkpoint namespaces. """ if not layer_types: raise ValueError("Layer groups cannot be empty.") if Symbols.MOE in layer_types[:-1]: raise ValueError(f"MoE layer '{Symbols.MOE}' must be the last symbol inside a layer group.") - namespaces = { - Symbols.MAMBA: "mixer", - Symbols.GDN: "self_attention", - Symbols.ATTENTION: "self_attention", - Symbols.DS_ATTENTION: "self_attention", - Symbols.MLA: "self_attention", - Symbols.CSA: "self_attention", - Symbols.HCA: "self_attention", - Symbols.WINDOW: "self_attention", - Symbols.MLP: "mlp", - Symbols.MOE: "mlp", - } - seen = set() - for layer_type in layer_types: - namespace = namespaces[layer_type] - if namespace in seen: - raise ValueError( - f"Layer group '{layer_type_list_to_str([tuple(layer_types)])}' contains " - f"multiple layers in checkpoint namespace '{namespace}'." - ) - seen.add(namespace) + layer_utils.validate_layer_group_checkpoint_namespaces(layer_types) @dataclass diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 2e1199d7560..10033b32882 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -126,11 +126,14 @@ class HybridModel(LanguageModule, GraphableMegatronModule): hybrid_layer_pattern (str): Unified hybrid layer pattern with optional MTP and pipeline stage boundaries. Format: "///..." - The main pattern may contain "|" to define pipeline stage boundaries. + The main pattern may contain "|" to define pipeline stage boundaries and + "[...]" to group layers into one logical layer (see "Bracketed layer groups" + in docs/user-guide/hybrid-model-migration.md). Examples: - "M*M*" -> main decoder only, no MTP - "M*M*/MM/MM" -> main="M*M*", mtp="MM", 2 depths - "M-M-|M-M*-|M-M-|M-M*-" -> 4 pipeline segments + - "[M*E][M*E]" -> 2 logical layers, each grouping Mamba, attention and MoE hybrid_attention_ratio (float, optional): Deprecated. Use hybrid_layer_pattern instead. If set to a value > 0.0 and hybrid_layer_pattern is None, a pattern will be generated from the ratio with a deprecation warning. diff --git a/megatron/core/models/hybrid/layers/utils.py b/megatron/core/models/hybrid/layers/utils.py index 5d065fb03a7..bbb79e77551 100644 --- a/megatron/core/models/hybrid/layers/utils.py +++ b/megatron/core/models/hybrid/layers/utils.py @@ -1,5 +1,7 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +from typing import Sequence + from megatron.core.ssm.gdn_layer_config import GDNLayerConfig from megatron.core.ssm.mamba_layer_config import MambaLayerConfig from megatron.core.ssm.mlp_layer_config import MLPLayerConfig @@ -28,7 +30,9 @@ class Symbols: MOE = 'E' PIPE = '|' MTP_SEPARATOR = "/" - # Bracketed groups (e.g. ``[M*E]``) build one nested HybridStack logical layer. + # Brackets group layers into one logical layer that HybridStack builds as a nested + # HybridStack, e.g. ``[M*E]``. See "Bracketed layer groups" in + # docs/user-guide/hybrid-model-migration.md. GROUP_START = "[" GROUP_END = "]" LAYER_CONFIG_MAP = { @@ -44,6 +48,20 @@ class Symbols: MOE: MoELayerConfig, } DSV4_COMPRESS_RATIO_MAP = {CSA: 4, HCA: 128, WINDOW: 0} + # Submodule under which each layer type stores its sharded checkpoint state. The layers + # of a bracketed group share one checkpoint layer index, so they need distinct namespaces. + CHECKPOINT_NAMESPACE_MAP = { + MAMBA: "mixer", + GDN: "self_attention", + ATTENTION: "self_attention", + DS_ATTENTION: "self_attention", + CSA: "self_attention", + HCA: "self_attention", + MLA: "self_attention", + WINDOW: "self_attention", + MLP: "mlp", + MOE: "mlp", + } MLA_ATTENTION = {MLA, DS_ATTENTION, CSA, HCA, WINDOW} ATTENTION_LAYER_CONFIGS = {AttentionLayerConfig, DSALayerConfig, CSALayerConfig, MLALayerConfig} @@ -58,6 +76,27 @@ def name_sorted_valid_layer_symbols(cls) -> list[str]: return [value for (_, value) in valid_layer_attrs] +def validate_layer_group_checkpoint_namespaces(layer_symbols: Sequence[str]) -> None: + """Require the layers of one bracketed group to use distinct checkpoint namespaces. + + A group shares one checkpoint layer index, so it can contain at most one mixer, one + attention module (including GDN), and one MLP or MoE module. + + Args: + layer_symbols: Symbols of the group's layers, in pattern order. + """ + seen = set() + for layer_symbol in layer_symbols: + namespace = Symbols.CHECKPOINT_NAMESPACE_MAP[layer_symbol] + if namespace in seen: + group = f"{Symbols.GROUP_START}{''.join(layer_symbols)}{Symbols.GROUP_END}" + raise ValueError( + f"Layer group '{group}' contains multiple layers in checkpoint namespace " + f"'{namespace}'." + ) + seen.add(namespace) + + def is_valid_symbol(layer_symbol: str, allow_pipe: bool = False) -> bool: """Return whether ``layer_symbol`` identifies a supported layer or allowed pipe. diff --git a/megatron/core/models/hybrid/model_chunk_schedule_plan.py b/megatron/core/models/hybrid/model_chunk_schedule_plan.py index fe51591e71d..95b6862e2ec 100644 --- a/megatron/core/models/hybrid/model_chunk_schedule_plan.py +++ b/megatron/core/models/hybrid/model_chunk_schedule_plan.py @@ -2,13 +2,14 @@ """Schedule-plan classes for HybridStack-based decoders. -These extend the GPT-side ``TransformerLayerSchedulePlan`` / -``TransformerModelChunkSchedulePlan`` with the per-layer ``layer_type`` symbol +These extend GPTModel's schedule plans, ``TransformerLayerSchedulePlan`` and +``TransformerModelChunkSchedulePlan``, with the per-layer ``layer_type`` symbol that HybridStack assigns to each entry of its ``layer_type_list`` (including -bracketed groups like ``[*-]``). The base classes remain GPT-only; this module -adds the hybrid-specific dispatch into ``build_hybrid_stack_callables`` and -uses ``HybridStackNode`` so the schedule node's free-input policy can diverge -from the GPT default. The pre/post-process nodes from +bracketed groups like ``[*-]``). The base classes build callables from the layer +module alone; this module adds the hybrid-specific dispatch into +``build_hybrid_stack_callables``, which also needs the layer's symbol, and uses +``HybridStackNode`` so the schedule node's free-input policy can diverge from the +GPTModel default. The pre/post-process nodes from ``core.models.common.utils`` are reused as-is — they already call ``model._preprocess`` / ``model._postprocess`` which work on a HybridModel. """ @@ -34,7 +35,7 @@ class HybridStackSchedulePlan(TransformerLayerSchedulePlan): Adds the ``layer_type`` extra-arg propagation; routes through ``build_hybrid_stack_callables`` when ``layer_type`` is set (i.e. the layer - is a HybridStack entry, possibly a bracketed group); falls back to the GPT + is a HybridStack entry, possibly a bracketed group); falls back to the GPTModel path for plain TransformerLayer / MTP layers when ``layer_type`` is None. """ @@ -98,7 +99,7 @@ def create_node(stream, module, name): self.mtp_post_process = NoopScheduleNode() def get_low_precision_context(self): - """Return the layer-level quantization context for GPT-path layers. + """Return the layer-level quantization context for GPTModel-path layers. Hybrid callables enter the quantization context of each physical layer themselves, so hybrid layer plans use a null context here. @@ -116,8 +117,8 @@ class HybridStackModelChunkSchedulePlan(TransformerModelChunkSchedulePlan): Threads HybridStack's ``layer_type_list[layer_idx]`` symbol into each layer plan's ``extra_args`` so the per-layer plan can dispatch grouped layers correctly. Layers of other modules (e.g. MTP layers) get - ``layer_type=None`` and follow the GPT path. The pre/post - process nodes inherit from the GPT base class — they already dispatch + ``layer_type=None`` and follow the GPTModel path. The pre/post + process nodes inherit from the GPTModel base class — they already dispatch on ``model._preprocess`` / ``model._postprocess`` which a HybridModel implements. """ diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index d0cbb54407f..1ea60f8b260 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -3600,7 +3600,9 @@ def _add_experimental_args(parser): help='Specify a hybrid layer pattern using M (mamba), G (gdn), ' '* (attention), D (dsa), - (mlp), E (moe). Use | to define pipeline ' 'stage boundaries for flexible virtual pipeline parallel (fVPP). ' - 'Use / to separate MTP patterns. ' + 'Use / to separate MTP patterns. Use [...] to group layers into one ' + 'logical layer, e.g. "[M*E][M*E]" (see the "Bracketed layer groups" ' + 'section of docs/user-guide/hybrid-model-migration.md). ' 'Example: "M-M-|M-M*-|M-M-|M-M*-" or "M-M-|M-M*-/MM/MM". ' 'When this flag is used, it is the sole indicator that a hybrid model ' 'is being run.')