diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index fe91f0146fb..eb8e19b4156 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -8,7 +8,7 @@ import copy from contextlib import nullcontext from dataclasses import dataclass -from typing import Optional, Tuple, Union +from typing import Optional, Sequence, Tuple, Union import torch from torch import Tensor, nn @@ -21,14 +21,25 @@ from megatron.core.fp8_utils import get_fp8_context from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode -from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols +from megatron.core.models.hybrid.hybrid_layer_allocation import ( + get_layer_type_list_from_layer_config_list, + validate_segment_layers, +) +from megatron.core.models.hybrid.layer_utils import normalize_tp_comm_overlap from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.recompute import checkpointed_forward +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 from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.attention_layer_config import AttentionLayerConfig from megatron.core.transformer.cuda_graphs import annotate_first_last_layer +from megatron.core.transformer.experimental_attention_variant.dsa_layer_config import DSALayerConfig from megatron.core.transformer.identity_op import IdentityOp +from megatron.core.transformer.mla_layer_config import MLALayerConfig from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.moe.moe_layer_config import MoELayerConfig from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_layer import TransformerLayer @@ -61,9 +72,13 @@ 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, optional): pre-computed list of layer type symbols for - this pipeline segment. When provided (by HybridModel), pipeline stage - selection has already been done via '|' separators in the pattern. + layer_type_list (list[str], optional): backward-compatible list of layer + type symbols for this pipeline segment. 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 segment. Defaults to 0. post_layer_norm (bool, optional): whether to include a final layer norm. @@ -82,7 +97,7 @@ def __init__( config: TransformerConfig, submodules: HybridStackSubmodules, pre_process: bool = True, - layer_type_list: Optional[list[str]] = None, + layer_type_list: list[str] | None = None, pp_layer_offset: int = 0, post_layer_norm: bool = True, post_process: bool = True, @@ -91,11 +106,24 @@ def __init__( pg_collection: ProcessGroupCollection = None, is_mtp_layer: bool = False, name: str | None = None, + layer_config_list: Sequence[TransformerConfig] | None = None, ) -> None: """ Args: 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) + normalize_tp_comm_overlap(config, segment) + layer_config_list = validate_segment_layers(segment, config) + super().__init__(config=config) self.pre_process = pre_process self.post_layer_norm = post_layer_norm @@ -111,39 +139,40 @@ def __init__( self.input_tensor = None self.pg_collection = pg_collection - assert layer_type_list is not None, ( - "layer_type_list must be provided. It should be pre-computed from " - "--hybrid-layer-pattern by HybridModel." - ) - self.layer_type_list = layer_type_list + assert layer_config_list is not None + self.layer_config_list = layer_config_list if getattr(self.config, "mla_down_proj_fusion", False): submodules = self._fuse_mla_down_proj(submodules) # Build layers from the pre-selected segment self.layers = nn.ModuleList() - for i, layer_type in enumerate(self.layer_type_list): + for i, layer_config in enumerate(self.layer_config_list): layer_number = i + 1 + pp_layer_offset - if self.config.fp8: - quant_init_context = get_fp8_context(self.config, i + pp_layer_offset, is_init=True) - elif self.config.fp4: - quant_init_context = get_fp4_context(self.config, i + pp_layer_offset, is_init=True) + if layer_config.fp8: + quant_init_context = get_fp8_context( + layer_config, i + pp_layer_offset, is_init=True + ) + elif layer_config.fp4: + quant_init_context = get_fp4_context( + layer_config, i + pp_layer_offset, is_init=True + ) else: quant_init_context = nullcontext() with quant_init_context: - if layer_type == LayerSymbols.MAMBA: + if isinstance(layer_config, MambaLayerConfig): layer = build_module( submodules.mamba_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pp_layer_offset=pp_layer_offset, pg_collection=pg_collection, name=(name + f".layers.{i}") if name is not None else None, ) - elif layer_type == LayerSymbols.ATTENTION: + elif isinstance(layer_config, AttentionLayerConfig): layer = build_module( submodules.attention_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, is_mtp_layer=is_mtp_layer, @@ -151,10 +180,10 @@ def __init__( pp_layer_offset=pp_layer_offset, name=(name + f".layers.{i}") if name is not None else None, ) - elif layer_type == LayerSymbols.DS_ATTENTION: + elif isinstance(layer_config, DSALayerConfig): layer = build_module( submodules.dsa_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, is_mtp_layer=is_mtp_layer, @@ -162,38 +191,38 @@ def __init__( pp_layer_offset=pp_layer_offset, name=(name + f".layers.{i}") if name is not None else None, ) - elif layer_type == LayerSymbols.MLA: + elif isinstance(layer_config, MLALayerConfig): layer = build_module( submodules.mla_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, is_mtp_layer=is_mtp_layer, add_layer_offset=False, pp_layer_offset=pp_layer_offset, ) - elif layer_type == LayerSymbols.MLP: + elif isinstance(layer_config, MLPLayerConfig): layer = build_module( submodules.mlp_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, add_layer_offset=False, name=(name + f".layers.{i}") if name is not None else None, ) - elif layer_type == LayerSymbols.MOE: + elif isinstance(layer_config, MoELayerConfig): layer = build_module( submodules.moe_layer, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, is_mtp_layer=is_mtp_layer, add_layer_offset=False, name=(name + f".layers.{i}") if name is not None else None, ) - elif layer_type == LayerSymbols.GDN: + elif isinstance(layer_config, GDNLayerConfig): gdn_layer_spec = submodules.gdn_layer - if self.config.experimental_attention_variant == "gdn2": + if layer_config.experimental_attention_variant == "gdn2": # 'G' layers build the GDN2 variant when the gdn2 experimental # attention variant is selected. from megatron.core.ssm.gated_delta_net import GatedDeltaNet2 @@ -202,7 +231,7 @@ def __init__( gdn_layer_spec.submodules.self_attention.module = GatedDeltaNet2 layer = build_module( gdn_layer_spec, - config=self.config, + config=layer_config, layer_number=layer_number, pg_collection=pg_collection, # Set to False as we do not want to change offset. @@ -211,7 +240,10 @@ def __init__( name=(name + f".layers.{i}") if name is not None else None, ) else: - raise ValueError("unexpected layer_type") + raise ValueError( + f"Unexpected hybrid layer config type: {type(layer_config).__name__}" + ) + self.layers.append(layer) if self.config.cuda_graph_impl == "local": @@ -228,6 +260,11 @@ def __init__( eps=self.config.layernorm_epsilon, ) + @property + def layer_type_list(self) -> list[str]: + """Return layer symbols derived from the per-layer configs for compatibility.""" + return get_layer_type_list_from_layer_config_list(self.layer_config_list) + def _fuse_mla_down_proj(self, submodules: HybridStackSubmodules) -> HybridStackSubmodules: # Avoid modifying the original object so users don't get surprised about their `submodules` # being modified underneath them. @@ -263,10 +300,10 @@ def mamba_state_shapes_per_request(self) -> Optional[Tuple[Tuple[int], Tuple[int Returns the recurrent mixer's conv and SSM state shapes per input sequence if this block contains Mamba or GDN layers (this may not be the case with PP > 1). """ - for layer_type, layer in zip(self.layer_type_list, self.layers): - if layer_type == LayerSymbols.MAMBA: + for layer_config, layer in zip(self.layer_config_list, self.layers, strict=True): + if isinstance(layer_config, MambaLayerConfig): return layer.mamba_state_shapes_per_request() - if layer_type == LayerSymbols.GDN: + if isinstance(layer_config, GDNLayerConfig): return layer.self_attention.mamba_state_shapes_per_request() return None @@ -371,10 +408,10 @@ def get_inner_quant_context(config, layer_number): use_inner_quantization_context=(use_inner_fp8_context or use_fp4_context), ) else: - for layer in self.layers: + for layer_config, layer in zip(self.layer_config_list, self.layers, strict=True): # Layers have 1-indexed layer numbers attribute. inner_quant_context = get_inner_quant_context( - self.config, layer.layer_number - 1 + layer_config, layer.layer_number - 1 ) with inner_quant_context: if isinstance(layer, TransformerLayer): diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index 83a6163b88d..e3af6c9f099 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -2,42 +2,24 @@ import logging from dataclasses import dataclass -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Sequence, Tuple import torch +from megatron.core.models.hybrid.layer_utils import Symbols, normalize_tp_comm_overlap +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 +from megatron.core.transformer.attention_layer_config import AttentionLayerConfig +from megatron.core.transformer.experimental_attention_variant.dsa_layer_config import DSALayerConfig +from megatron.core.transformer.mla_layer_config import MLALayerConfig +from megatron.core.transformer.moe.moe_layer_config import MoELayerConfig +from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import log_on_each_pipeline_stage, log_single_rank logger = logging.getLogger(__name__) -class Symbols: - """Symbols for different layer types and pattern separators.""" - - MAMBA = "M" - GDN = 'G' - ATTENTION = "*" - DS_ATTENTION = "D" - MLA = "+" - MLP = "-" - MOE = 'E' - PIPE = '|' - MTP_SEPARATOR = "/" - VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, MLA, MLP, MOE} - - @classmethod - def name_sorted_valid_layer_symbols(cls) -> list[str]: - """Return the valid layer symbols sorted lexicographically by their public attribute - name. - """ - valid_layer_attrs = [] - for name, value in vars(cls).items(): - if not name.startswith('_') and value in cls.VALID_LAYERS: - valid_layer_attrs.append((name, value)) - valid_layer_attrs.sort() - return [value for (_, value) in valid_layer_attrs] - - @dataclass class ParsedHybridPattern: """Result of parsing a unified hybrid pattern string. @@ -170,14 +152,14 @@ def get_hybrid_layer_counts(pattern: str) -> Dict[str, int]: Returns: Dictionary mapping layer symbol to count. Keys are all valid layer symbols - (Symbols.VALID_LAYERS). + (``Symbols.VALID_LAYERS``). Examples: >>> get_hybrid_layer_counts("M*M*") - {'*': 2, 'G': 0, 'D': 0, 'M': 2, '-': 0, 'E': 0} + {'*': 2, 'D': 0, 'G': 0, 'M': 2, '+': 0, '-': 0, 'E': 0} >>> get_hybrid_layer_counts("M-M-|M-M*-/MM/MM") - {'*': 1, 'G': 0, 'D': 0, 'M': 8, '-': 4, 'E': 0} + {'*': 1, 'D': 0, 'G': 0, 'M': 8, '+': 0, '-': 4, 'E': 0} """ parsed = parse_hybrid_pattern(pattern) counts = {symbol: 0 for symbol in Symbols.name_sorted_valid_layer_symbols()} @@ -298,45 +280,78 @@ def _validate_pattern(pattern: str, pattern_name: str, allow_pipe: bool = False) raise ValueError("Not supported to have both Attention and MLA/DSA in one model") -def validate_segment_layers(segment: str) -> List[str]: - """Validate and convert a single pipeline segment pattern to a layer type list. +def _validate_segment_layer_symbols(segment: str) -> None: + """Validate the layer symbols in a single pipeline segment.""" + for layer_symbol in segment: + if layer_symbol not in Symbols.VALID_LAYERS: + raise ValueError( + f"In hybrid layer pattern segment, '{layer_symbol}' is not " + f"one of {Symbols.VALID_LAYERS}" + ) + + # Disallow Attention + MLA/DSA hybridity. + if Symbols.ATTENTION in segment and (Symbols.DS_ATTENTION in segment or Symbols.MLA in segment): + raise ValueError("Not supported to have both Attention and MLA/DSA in one model") + + +def _create_layer_config(config: TransformerConfig, layer_symbol: str) -> TransformerConfig: + """Create a layer-specific config from a normalized stack-level config.""" + if layer_symbol == Symbols.MAMBA: + return MambaLayerConfig.from_config(config) + if layer_symbol == Symbols.GDN: + return GDNLayerConfig.from_config(config) + if layer_symbol == Symbols.ATTENTION: + return AttentionLayerConfig.from_config(config) + if layer_symbol == Symbols.DS_ATTENTION: + return DSALayerConfig.from_config(config) + if layer_symbol == Symbols.MLA: + return MLALayerConfig.from_config(config) + if layer_symbol == Symbols.MLP: + return MLPLayerConfig.from_config(config) + if layer_symbol == Symbols.MOE: + return MoELayerConfig.from_config(config) + raise ValueError(f"Unexpected hybrid layer symbol: {layer_symbol}") + + +def validate_segment_layers(segment: str, config: TransformerConfig) -> List[TransformerConfig]: + """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 '|'). + The source config is expected to be normalized before this function is called. + Each layer config is created from that state without running ``__post_init__`` + a second time. + Args: segment: A single pipeline segment pattern string (e.g., "M-M*-") + config: Normalized stack-level config to copy for each layer. Returns: - List of layer type characters. + List of independent per-layer configs. Raises: ValueError: If segment contains invalid layer symbols. """ - layer_type_list = list(segment) - for layer_char in layer_type_list: - if layer_char not in Symbols.VALID_LAYERS: - raise ValueError( - f"In hybrid layer pattern segment, '{layer_char}' is not " - f"one of {Symbols.VALID_LAYERS}" - ) + _validate_segment_layer_symbols(segment) - # Disallow Attention + MLA/DSA hybridity. - if Symbols.ATTENTION in segment and (Symbols.DS_ATTENTION in segment or Symbols.MLA in segment): - raise ValueError("Not supported to have both Attention and MLA/DSA in one model") + layer_configs: list[TransformerConfig] = [] + for layer_symbol in segment: + layer_configs.append(_create_layer_config(config, layer_symbol)) - return layer_type_list + return layer_configs def select_pipeline_segment( 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[str], int]: +) -> Tuple[List[TransformerConfig], 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 @@ -349,6 +364,7 @@ def select_pipeline_segment( Args: main_pattern: Main decoder pattern (may contain '|' separators). Empty string is allowed (produces one empty segment). + config: Normalized stack-level config to copy for each selected layer. pp_group: Pipeline parallel process group, or None if not using PP. vp_stage: Virtual pipeline stage, or None if not using VPP. first_stage_layers: Number of layers on the first pipeline stage for @@ -359,8 +375,8 @@ def select_pipeline_segment( dp_cp_group: Optional data/context-parallel process group used for per-stage logging. Returns: - Tuple of (layer_type_list, layer_offset) where layer_type_list is - the list of layer type characters for this segment, and layer_offset + 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. Raises: @@ -399,8 +415,8 @@ def select_pipeline_segment( "Example: 'M*M*M*M*' with pp_size=2 should become 'M*M*|M*M*'.", ) full_pattern = segments[0] - layer_type_list = validate_segment_layers(full_pattern) - num_layers = len(layer_type_list) + _validate_segment_layer_symbols(full_pattern) + num_layers = len(full_pattern) if first_stage_layers is not None or last_stage_layers is not None: first = first_stage_layers or 0 @@ -443,12 +459,14 @@ def select_pipeline_segment( offset = pp_rank * layers_per_rank count = layers_per_rank - selected = layer_type_list[offset : offset + count] + selected_pattern = full_pattern[offset : offset + count] + normalize_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='{''.join(selected)}' ({len(selected)} layers), " + f"layers='{selected_pattern}' ({len(selected)} layers), " f"layer_offset={offset} (auto-split)", tp_group=tp_group, dp_cp_group=dp_cp_group, @@ -477,20 +495,21 @@ def select_pipeline_segment( layer_offset = sum(len(segments[i]) for i in range(segment_index)) my_segment = segments[segment_index] - layer_type_list = validate_segment_layers(my_segment) + normalize_tp_comm_overlap(config, my_segment) + layer_config_list = validate_segment_layers(my_segment, config) log_on_each_pipeline_stage( logger, 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_type_list)} layers), " + f"layers='{my_segment}' ({len(layer_config_list)} layers), " f"layer_offset={layer_offset}", tp_group=tp_group, dp_cp_group=dp_cp_group, ) - return layer_type_list, layer_offset + return layer_config_list, layer_offset def get_layer_maps_from_layer_type_list(layer_type_list: list[str]) -> dict[str, dict[int, int]]: @@ -505,3 +524,47 @@ def get_layer_maps_from_layer_type_list(layer_type_list: list[str]) -> dict[str, local_layer_idx = len(layer_map) layer_map[global_layer_idx] = local_layer_idx return layer_maps + + +def _get_layer_symbol_from_config(layer_config: TransformerConfig) -> str: + """Return the canonical symbol for a layer config, including subclasses.""" + matching_symbols = [] + if isinstance(layer_config, MambaLayerConfig): + matching_symbols.append(Symbols.MAMBA) + if isinstance(layer_config, GDNLayerConfig): + matching_symbols.append(Symbols.GDN) + if isinstance(layer_config, AttentionLayerConfig): + matching_symbols.append(Symbols.ATTENTION) + if isinstance(layer_config, DSALayerConfig): + matching_symbols.append(Symbols.DS_ATTENTION) + if isinstance(layer_config, MLALayerConfig): + matching_symbols.append(Symbols.MLA) + if isinstance(layer_config, MLPLayerConfig): + matching_symbols.append(Symbols.MLP) + if isinstance(layer_config, MoELayerConfig): + matching_symbols.append(Symbols.MOE) + if not matching_symbols: + raise ValueError(f"Unexpected hybrid layer config type: {type(layer_config).__name__}") + if len(matching_symbols) > 1: + raise ValueError( + f"Ambiguous hybrid layer config type: {type(layer_config).__name__} " + f"matches symbols {matching_symbols}" + ) + return matching_symbols[0] + + +def get_layer_type_list_from_layer_config_list( + layer_config_list: Sequence[TransformerConfig], +) -> list[str]: + """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``. + + Args: + layer_config_list: Per-layer configs in layer order. + + Returns: + The canonical layer symbol for each config. + """ + return [_get_layer_symbol_from_config(layer_config) for layer_config in layer_config_list] diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 84b00b9b50e..2ef96edcce0 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -14,6 +14,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.layer_utils import normalize_tp_comm_overlap from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, @@ -205,17 +206,6 @@ def __init__( self.mtp_pattern = parsed.mtp_pattern self.mtp_num_depths = parsed.mtp_num_depths - logging_pg_kwargs = _hybrid_logging_pg_kwargs(self.pg_collection) - - layer_type_list, layer_offset = select_pipeline_segment( - parsed.main_pattern or '', - 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, - ) - # Determine if MTP is needed (based on pattern parsing) self.mtp_process = ( self.mtp_pattern is not None @@ -231,6 +221,19 @@ def __init__( vp_stage=self.vp_stage, ) ) + normalize_tp_comm_overlap(self.config, '', has_mtp=self.mtp_process) + + 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, + ) # megatron core pipelining currently depends on model type # TODO: remove this dependency ? @@ -282,7 +285,7 @@ def __init__( hybrid_stack_spec, self.config, pre_process=self.pre_process, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=layer_offset, post_process=self.post_process, dtype=config.params_dtype, diff --git a/megatron/core/models/hybrid/layer_utils.py b/megatron/core/models/hybrid/layer_utils.py new file mode 100644 index 00000000000..d8e8ba857fd --- /dev/null +++ b/megatron/core/models/hybrid/layer_utils.py @@ -0,0 +1,62 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import warnings + +from megatron.core.transformer.transformer_config import TransformerConfig + + +class Symbols: + """Symbols for different layer types and pattern separators.""" + + MAMBA = "M" + GDN = 'G' + ATTENTION = "*" + DS_ATTENTION = "D" + MLA = "+" + MLP = "-" + MOE = 'E' + PIPE = '|' + MTP_SEPARATOR = "/" + VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, MLA, MLP, MOE} + + @classmethod + def name_sorted_valid_layer_symbols(cls) -> list[str]: + """Return valid layer symbols sorted by their public attribute names.""" + valid_layer_attrs = [] + for name, value in vars(cls).items(): + if not name.startswith('_') and value in cls.VALID_LAYERS: + valid_layer_attrs.append((name, value)) + valid_layer_attrs.sort() + return [value for (_, value) in valid_layer_attrs] + + +def normalize_tp_comm_overlap( + config: TransformerConfig, segment: str, has_mtp: bool = False +) -> None: + """Disable TP communication overlap unsupported by built-in hybrid layers. + + This must run before ``validate_segment_layers`` copies the stack-level config so + every generated layer config receives the normalized value. + + Args: + config: Stack-level config that will be copied for each layer. + segment: Selected pipeline segment, containing only layer symbols. + has_mtp: Whether this model instance will build an MTP block. + """ + unsupported_features: list[str] = [] + if Symbols.MLA in segment: + unsupported_features.append("MLA") + if Symbols.DS_ATTENTION in segment: + unsupported_features.append("DSA") + if has_mtp: + unsupported_features.append("MTP") + + if not config.tp_comm_overlap or not unsupported_features: + return + + config.tp_comm_overlap = False + warnings.warn( + "TP communication overlap is not supported with hybrid " + f"{'/'.join(unsupported_features)} layers. Disabling tp_comm_overlap.", + stacklevel=2, + ) diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 11f5d9dd462..eeff6a06afa 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -1203,7 +1203,7 @@ def __init__( self.mtp_model_layer = HybridStack( config=self.config, submodules=hybrid_submodules, - layer_type_list=validate_segment_layers(mtp_layer_pattern), + layer_config_list=validate_segment_layers(mtp_layer_pattern, self.config), pp_layer_offset=0, pre_process=True, # Always receives input from eh_proj post_layer_norm=False, # MTP has its own final_layernorm diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 4815c876a90..0148d772bd3 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -3,8 +3,9 @@ import logging import math import warnings +from copy import deepcopy from dataclasses import dataclass, field -from typing import Callable, List, Literal, Optional, Tuple, Union +from typing import Callable, List, Literal, Optional, Self, Tuple, Union import torch import torch.nn.functional as F @@ -1390,6 +1391,24 @@ class TransformerConfig(ModelParallelConfig): insert these joins. This feature is particularly useful when using with full-iteration CUDA graphs""" + @classmethod + def from_config(cls, config: "TransformerConfig") -> Self: + """Create this config type from an existing normalized transformer config. + + The source config's complete instance state is deep-copied without invoking + the target class's initializer or ``__post_init__``. This preserves normalized + values and dynamically added attributes while producing an independent config. + + Args: + config: The transformer config to copy. + + Returns: + An independent copy of ``config`` whose type is ``cls``. + """ + new_config = cls.__new__(cls) + new_config.__dict__ = deepcopy(config.__dict__, {id(config): new_config}) + return new_config + def __post_init__(self): """Python dataclass method that is used to modify attributes after initialization. See https://docs.python.org/3/library/dataclasses.html#post-init-processing for more diff --git a/megatron/core/transformer/utils.py b/megatron/core/transformer/utils.py index 9983f2f6dc0..af7d3fda164 100644 --- a/megatron/core/transformer/utils.py +++ b/megatron/core/transformer/utils.py @@ -400,7 +400,7 @@ def set_model_to_sequence_parallel(model, set_to=False, exclude_modules=None): if _sequence_parallel_attr_cache is None or model_id not in _sequence_parallel_attr_cache: _init_sequence_parallel_cache(model, exclude_modules) - model.config.sequence_parallel = set_to + set_model_config_attribute(model, "sequence_parallel", set_to) # Set all cached attributes to desired value for attr, modules in _sequence_parallel_attr_cache[model_id].items(): @@ -484,7 +484,7 @@ def toggle_cuda_graphs(model, set_to="none"): init_cuda_graph_cache(model) assert set_to in ["none", "local"], f"Invalid CUDA graph implementation: {set_to}" - model.config.cuda_graph_impl = set_to + set_model_config_attribute(model, "cuda_graph_impl", set_to) # Collect all modules that have any of the CUDA graph attributes for attribute, modules in cuda_graph_attr_cache[model_id].items(): diff --git a/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py b/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py index 51568243d0d..e1cacc7a2c0 100644 --- a/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py +++ b/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py @@ -196,9 +196,10 @@ def _build_mamba_model( post_process: bool = True, ) -> HybridModel: """Build a HybridModel with the given hybrid layer pattern.""" - layer_type_list = validate_segment_layers(layer_pattern) mamba_config = copy.deepcopy(config) - mamba_config.num_layers = len(layer_type_list) + mamba_config.num_layers = len(layer_pattern) + layer_config_list = validate_segment_layers(layer_pattern, mamba_config) + assert len(layer_config_list) == mamba_config.num_layers assert mamba_config.num_layers == _NUM_GPT_LAYERS * 2 model = HybridModel( config=mamba_config, diff --git a/tests/unit_tests/models/test_hybrid_model.py b/tests/unit_tests/models/test_hybrid_model.py index 95bcaa2d7d0..04ab99c13e6 100644 --- a/tests/unit_tests/models/test_hybrid_model.py +++ b/tests/unit_tests/models/test_hybrid_model.py @@ -21,11 +21,15 @@ from megatron.core.inference.sampling_params import SamplingParams from megatron.core.inference.utils import InferenceMode from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding +from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec from megatron.core.models.hybrid.hybrid_model import HybridModel, _hybrid_logging_pg_kwargs from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.ssm.mamba_layer_config import MambaLayerConfig +from megatron.core.ssm.mlp_layer_config import MLPLayerConfig 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.module import Float16Module from megatron.core.utils import divide, is_fa_min_version, is_torch_min_version @@ -234,9 +238,55 @@ def test_constructor(self): assert self.model.max_sequence_length == 4 + decoder = self.model.decoder + assert "layer_type_list" not in decoder.__dict__ + assert decoder.layer_type_list == [Symbols.MAMBA, Symbols.ATTENTION, Symbols.MLP] + assert [type(config) for config in decoder.layer_config_list] == [ + MambaLayerConfig, + AttentionLayerConfig, + MLPLayerConfig, + ] + assert len({id(config) for config in decoder.layer_config_list}) == 3 + assert all(config is not self.model.config for config in decoder.layer_config_list) + assert all( + layer.config is layer_config + for layer, layer_config in zip(decoder.layers, decoder.layer_config_list, strict=True) + ) + num_weights = sum([p.numel() for p in self.model.parameters()]) assert num_weights == 1774872 + def test_mtp_tp_overlap_is_normalized_before_decoder_configs(self, monkeypatch): + class CheckingMTPBlock(torch.nn.Module): + + def __init__(self, config, **kwargs): + super().__init__() + assert config.tp_comm_overlap is False + + monkeypatch.setattr( + "megatron.core.models.hybrid.hybrid_model.MultiTokenPredictionBlock", CheckingMTPBlock + ) + + model_config = TransformerConfig( + num_layers=1, + hidden_size=256, + num_attention_heads=4, + use_cpu_initialization=True, + mtp_num_layers=1, + tp_comm_overlap=True, + ) + with pytest.warns(UserWarning, match="Disabling tp_comm_overlap"): + model = HybridModel( + config=model_config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern="-/M", + ) + + assert model.config.tp_comm_overlap is False + assert all(config.tp_comm_overlap is False for config in model.decoder.layer_config_list) + def test_set_input_tensor(self): config: TransformerConfig = self.model.config sequence_length = self.model.max_sequence_length diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index 322f1d78911..e8c3593ba18 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -1,8 +1,12 @@ # Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. +from types import SimpleNamespace + import pytest import torch +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.hybrid_block import HybridStack from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols, validate_segment_layers @@ -13,13 +17,17 @@ from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.ssm.gated_delta_net import GatedDeltaNet from megatron.core.ssm.mamba_layer import MambaLayer +from megatron.core.ssm.mamba_layer_config import MambaLayerConfig +from megatron.core.ssm.mlp_layer_config import MLPLayerConfig from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from megatron.core.transformer.attention import SelfAttention +from megatron.core.transformer.attention_layer_config import AttentionLayerConfig from megatron.core.transformer.experimental_attention_variant.absorbed_mla import ( AbsorbedMLASelfAttention, ) from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention +from megatron.core.transformer.mla_layer_config import MLALayerConfig from megatron.core.transformer.mlp import MLP from megatron.core.transformer.multi_latent_attention import MLASelfAttention from megatron.core.transformer.transformer_config import MLATransformerConfig @@ -27,6 +35,288 @@ from tests.unit_tests.test_utilities import Utils +@pytest.mark.parametrize( + ("layer_pattern", "expected_spec_names"), + [ + ( + Symbols.MAMBA + Symbols.GDN + Symbols.ATTENTION + Symbols.MLP + Symbols.MOE, + ["mamba_layer", "gdn_layer", "attention_layer", "mlp_layer", "moe_layer"], + ), + (Symbols.DS_ATTENTION + Symbols.MLA, ["dsa_layer", "mla_layer"]), + ], +) +def test_all_layer_configs_route_to_matching_specs(monkeypatch, layer_pattern, expected_spec_names): + """Each config marker selects its matching layer spec and config instance.""" + + 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((module_spec, kwargs)) + return BuiltLayer(kwargs["config"], kwargs["layer_number"]) + + monkeypatch.setattr(hybrid_block_module, "build_module", fake_build_module) + + config = MLATransformerConfig( + num_layers=len(layer_pattern), hidden_size=64, num_attention_heads=4 + ) + layer_config_list = validate_segment_layers(layer_pattern, config) + submodules = hybrid_stack_spec.submodules + expected_specs = [getattr(submodules, spec_name) for spec_name in expected_spec_names] + + block = HybridStack( + config=config, + submodules=submodules, + layer_config_list=layer_config_list, + pre_process=False, + pp_layer_offset=5, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + name="decoder", + ) + + assert "layer_type_list" not in block.__dict__ + assert block.layer_type_list == list(layer_pattern) + assert [module_spec for module_spec, _ in build_calls] == expected_specs + assert all( + kwargs["config"] is layer_config + for (_, kwargs), layer_config in zip(build_calls, layer_config_list) + ) + expected_layer_numbers = list(range(6, 6 + len(layer_pattern))) + assert [kwargs["layer_number"] for _, kwargs in build_calls] == expected_layer_numbers + assert [layer.layer_number for layer in block.layers] == expected_layer_numbers + + +def test_hybrid_stack_accepts_layer_config_subclasses(monkeypatch): + """Layer config subclasses retain their parent layer's routing behavior.""" + + class CustomMambaLayerConfig(MambaLayerConfig): + pass + + 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(module_spec) + return BuiltLayer(kwargs["config"], kwargs["layer_number"]) + + monkeypatch.setattr(hybrid_block_module, "build_module", fake_build_module) + + root_config = TransformerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + layer_config = CustomMambaLayerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + block = HybridStack( + config=root_config, + submodules=hybrid_stack_spec.submodules, + layer_config_list=[layer_config], + pre_process=False, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + + assert build_calls == [hybrid_stack_spec.submodules.mamba_layer] + assert block.layers[0].config is layer_config + + +def test_layer_type_list_normalizes_tp_overlap_before_copying_configs(monkeypatch): + """The positional layer-type API normalizes the root config before conversion.""" + + class BuiltLayer(torch.nn.Module): + + def __init__(self, config, layer_number): + super().__init__() + self.config = config + self.layer_number = layer_number + + submodules = hybrid_stack_spec.submodules + + def fake_build_module(module_spec, **kwargs): + return BuiltLayer(kwargs["config"], kwargs["layer_number"]) + + monkeypatch.setattr(hybrid_block_module, "build_module", fake_build_module) + + config = MLATransformerConfig( + num_layers=3, hidden_size=64, num_attention_heads=4, tp_comm_overlap=True + ) + with pytest.warns(UserWarning, match="Disabling tp_comm_overlap"): + block = HybridStack( + config, + submodules, + False, + [Symbols.MAMBA, Symbols.MLA, Symbols.MLP], + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + layer_config_list = block.layer_config_list + + assert "layer_type_list" not in block.__dict__ + assert block.layer_type_list == [Symbols.MAMBA, Symbols.MLA, Symbols.MLP] + assert type(layer_config_list) is list + assert [type(layer_config) for layer_config in layer_config_list] == [ + MambaLayerConfig, + MLALayerConfig, + MLPLayerConfig, + ] + assert len({id(layer_config) for layer_config in layer_config_list}) == len(layer_config_list) + assert all(layer_config is not config for layer_config in layer_config_list) + assert all( + layer.config is layer_config for layer, layer_config in zip(block.layers, layer_config_list) + ) + assert config.tp_comm_overlap is False + assert all(layer_config.tp_comm_overlap is False for layer_config in layer_config_list) + + block.position_embedding_type = "rope" + config.sequence_parallel = True + for layer_config in layer_config_list: + layer_config.sequence_parallel = True + + monkeypatch.setattr(transformer_utils, "_sequence_parallel_attr_cache", None) + transformer_utils.set_model_to_sequence_parallel(block, set_to=False) + + assert config.sequence_parallel is False + assert all(layer_config.sequence_parallel is False for layer_config in layer_config_list) + + +def test_explicit_layer_config_mutations_are_isolated(monkeypatch): + """Mutating one explicitly supplied layer config does not affect the others.""" + + class BuiltLayer(torch.nn.Module): + + def __init__(self, config, layer_number): + super().__init__() + self.config = config + self.layer_number = layer_number + + submodules = hybrid_stack_spec.submodules + + def fake_build_module(module_spec, **kwargs): + if module_spec is submodules.mla_layer: + kwargs["config"].tp_comm_overlap = False + return BuiltLayer(kwargs["config"], kwargs["layer_number"]) + + monkeypatch.setattr(hybrid_block_module, "build_module", fake_build_module) + + root_config = MLATransformerConfig( + num_layers=2, hidden_size=64, num_attention_heads=4, tp_comm_overlap=True + ) + layer_configs = validate_segment_layers(Symbols.MLA + Symbols.MLP, root_config) + HybridStack( + config=root_config, + submodules=submodules, + layer_config_list=layer_configs, + pre_process=False, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + + assert type(layer_configs) is list + assert root_config.tp_comm_overlap is True + assert [layer_config.tp_comm_overlap for layer_config in layer_configs] == [False, True] + + +@pytest.mark.parametrize( + ("provide_layer_type_list", "provide_layer_config_list"), + [(False, False), (True, True)], + ids=["neither", "both"], +) +def test_hybrid_stack_requires_exactly_one_layer_list( + provide_layer_type_list, provide_layer_config_list +): + """HybridStack requires exactly one legacy symbol list or per-layer config list.""" + config = TransformerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + layer_type_list = [Symbols.MAMBA] if provide_layer_type_list else None + layer_config_list = ( + validate_segment_layers(Symbols.MAMBA, config) if provide_layer_config_list else None + ) + + with pytest.raises( + ValueError, match="Exactly one of layer_type_list or layer_config_list must be provided" + ): + HybridStack( + config=config, + submodules=hybrid_stack_spec.submodules, + layer_type_list=layer_type_list, + layer_config_list=layer_config_list, + pre_process=False, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + + +def test_hybrid_stack_rejects_multi_character_layer_type(): + """The legacy list treats each entry as one layer symbol.""" + config = TransformerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + + with pytest.raises(ValueError, match="Each entry in layer_type_list must be a single"): + HybridStack( + config=config, + submodules=hybrid_stack_spec.submodules, + layer_type_list=[Symbols.MAMBA + Symbols.ATTENTION], + pre_process=False, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + + +def test_mamba_state_shapes_are_selected_by_layer_config_type(): + """Mamba state shape lookup does not depend on layer symbols or module methods alone.""" + + class CustomMambaLayerConfig(MambaLayerConfig): + pass + + attention_config = object.__new__(AttentionLayerConfig) + mamba_config = object.__new__(CustomMambaLayerConfig) + attention_shapes = ((1,), (2,)) + mamba_shapes = ((3,), (4,)) + block = SimpleNamespace( + layer_config_list=[attention_config, mamba_config], + layers=[ + SimpleNamespace(mamba_state_shapes_per_request=lambda: attention_shapes), + SimpleNamespace(mamba_state_shapes_per_request=lambda: mamba_shapes), + ], + ) + + assert HybridStack.mamba_state_shapes_per_request(block) == mamba_shapes + + block.layer_config_list = [attention_config] + block.layers = block.layers[:1] + assert HybridStack.mamba_state_shapes_per_request(block) is None + + +def test_hybrid_stack_rejects_same_named_config_type(): + root_config = TransformerConfig(num_layers=1, hidden_size=64, num_attention_heads=4) + same_named_config_class = type("MambaLayerConfig", (TransformerConfig,), {}) + layer_config = same_named_config_class(num_layers=1, hidden_size=64, num_attention_heads=4) + + with pytest.raises(ValueError, match="Unexpected hybrid layer config type: MambaLayerConfig"): + HybridStack( + config=root_config, + submodules=hybrid_stack_spec.submodules, + layer_config_list=[layer_config], + pre_process=False, + post_layer_norm=False, + post_process=False, + pg_collection=SimpleNamespace(pp=None, tp=None), + ) + + @pytest.mark.internal class TestHybridBlock: @@ -38,32 +328,31 @@ def get_pg_collection(self): return ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'pp', 'cp']) def get_hybrid_block(self, layer_pattern, **config_kwargs): - layer_type_list = validate_segment_layers(layer_pattern) transformer_config = TransformerConfig( 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_type_list), + num_layers=len(layer_pattern), num_attention_heads=4, use_cpu_initialization=True, **config_kwargs, ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) modules = hybrid_stack_spec.submodules return HybridStack( transformer_config, modules, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=0, pg_collection=self.get_pg_collection(), ) def get_dsa_hybrid_block(self, layer_pattern): - layer_type_list = validate_segment_layers(layer_pattern) transformer_config = MLATransformerConfig( 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_type_list), + num_layers=len(layer_pattern), num_attention_heads=16, use_cpu_initialization=True, bf16=True, @@ -81,22 +370,22 @@ def get_dsa_hybrid_block(self, layer_pattern): dsa_indexer_topk=32, add_bias_linear=False, ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) modules = hybrid_stack_spec.submodules return HybridStack( transformer_config, modules, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=0, pg_collection=self.get_pg_collection(), ) def get_mla_hybrid_block(self, layer_pattern): - layer_type_list = validate_segment_layers(layer_pattern) transformer_config = MLATransformerConfig( 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_type_list), + num_layers=len(layer_pattern), num_attention_heads=16, use_cpu_initialization=True, bf16=True, @@ -110,11 +399,12 @@ def get_mla_hybrid_block(self, layer_pattern): rotary_base=10000, rotary_percent=1.0, ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) modules = hybrid_stack_spec.submodules return HybridStack( transformer_config, modules, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=0, pg_collection=self.get_pg_collection(), ) @@ -245,11 +535,16 @@ def test_layer_types(self): assert isinstance(layers[1].self_attention, SelfAttention) assert isinstance(layers[2], TransformerLayer) assert isinstance(layers[2].mlp, MLP) + assert len({id(config) for config in block.layer_config_list}) == len(layer_pattern) + assert all( + layer.config is layer_config + for layer, layer_config in zip(block.layers, block.layer_config_list) + ) def test_invalid_layer_types_cause_failure(self): - invalid_symbol = 'X' - assert invalid_symbol not in Symbols.VALID_LAYERS # sanity check. - layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP + invalid_symbol + invalid_pattern_char = 'X' + assert invalid_pattern_char not in Symbols.VALID_LAYERS # sanity check. + layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP + invalid_pattern_char # validate_segment_layers() in hybrid_layer_allocation.py throws a ValueError. with pytest.raises(ValueError): block = self.get_hybrid_block(layer_pattern) @@ -277,19 +572,19 @@ def test_gdn_inference_spec(self): def test_gdn_gpu_forward(self): """Test GPU forward pass with GDN, attention, and Mamba layers.""" layer_pattern = Symbols.GDN + Symbols.ATTENTION + Symbols.MAMBA - layer_type_list = validate_segment_layers(layer_pattern) transformer_config = TransformerConfig( hidden_size=256, - num_layers=len(layer_type_list), + num_layers=len(layer_pattern), num_attention_heads=4, use_cpu_initialization=True, activation_func=torch.nn.functional.silu, ) + layer_config_list = validate_segment_layers(layer_pattern, transformer_config) modules = hybrid_stack_spec.submodules block = HybridStack( transformer_config, modules, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=0, pg_collection=self.get_pg_collection(), ) diff --git a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py index 8b4c181ee30..be34bec3710 100644 --- a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py +++ b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py @@ -1,5 +1,6 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +import functools import operator from unittest.mock import patch @@ -17,6 +18,48 @@ select_pipeline_segment, validate_segment_layers, ) +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 +from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.attention_layer_config import AttentionLayerConfig +from megatron.core.transformer.experimental_attention_variant.dsa_layer_config import DSALayerConfig +from megatron.core.transformer.mla_layer_config import MLALayerConfig +from megatron.core.transformer.moe.moe_layer_config import MoELayerConfig +from megatron.core.transformer.transformer_config import MLATransformerConfig + +_EXPECTED_LAYER_CONFIG_CLASSES = { + Symbols.MAMBA: MambaLayerConfig, + Symbols.GDN: GDNLayerConfig, + Symbols.ATTENTION: AttentionLayerConfig, + Symbols.DS_ATTENTION: DSALayerConfig, + Symbols.MLA: MLALayerConfig, + Symbols.MLP: MLPLayerConfig, + Symbols.MOE: MoELayerConfig, +} + + +def _make_transformer_config() -> TransformerConfig: + return TransformerConfig(num_layers=7, hidden_size=64, num_attention_heads=4) + + +def _assert_layer_config_types(layer_config_list, pattern: str) -> None: + assert [type(config) for config in layer_config_list] == [ + _EXPECTED_LAYER_CONFIG_CLASSES[layer_symbol] for layer_symbol in pattern + ] + + +def _assert_config_contents_equal(actual, expected) -> None: + assert vars(actual).keys() == vars(expected).keys() + for field_name, expected_value in vars(expected).items(): + actual_value = getattr(actual, field_name) + if isinstance(expected_value, functools.partial): + assert isinstance(actual_value, functools.partial) + assert actual_value.func is expected_value.func + assert actual_value.args == expected_value.args + assert actual_value.keywords == expected_value.keywords + else: + assert actual_value == expected_value @pytest.mark.internal @@ -67,46 +110,105 @@ def test_returns_string(self): @pytest.mark.internal class TestValidateSegmentLayers: - def test_valid_patterns(self): - """Test that valid segment patterns produce the correct layer type lists.""" - test_cases = [ - ("M*-M*-M*-", ['M', '*', '-', 'M', '*', '-', 'M', '*', '-']), - ("MMMMMMMMM", ['M'] * 9), - ("MM*-MM*-", ['M', 'M', '*', '-', 'M', 'M', '*', '-']), - ("E", ['E']), - ("", []), - ("GGG*GGG*", ['G', 'G', 'G', '*', 'G', 'G', 'G', '*']), - ("GEGEGE*E", ['G', 'E', 'G', 'E', 'G', 'E', '*', 'E']), - ("MDMD", ['M', 'D', 'M', 'D']), - ("M+M+", ['M', '+', 'M', '+']), - ] - for pattern, expected in test_cases: - result = validate_segment_layers(pattern) - assert result == expected, f"Failed for pattern: {pattern}" + def setup_method(self): + self.config = _make_transformer_config() - def test_all_valid_symbols(self): + def test_valid_patterns(self): + """Test that valid segment patterns produce configs in the correct order.""" + for pattern in [ + "M*-M*-M*-", + "MMMMMMMMM", + "MM*-MM*-", + "E", + "", + "GGG*GGG*", + "GEGEGE*E", + "MDMD", + "M+M+", + ]: + result = validate_segment_layers(pattern, self.config) + _assert_layer_config_types(result, pattern) + + def test_all_valid_pattern_characters(self): """Make sure all returned layers are valid.""" for pattern in ["M*-M*-M*-", "MMMMMMMMM", "MM*-", "MEME"]: - layer_types = validate_segment_layers(pattern) - for layer_type in layer_types: - assert layer_type in Symbols.VALID_LAYERS + layer_config_list = validate_segment_layers(pattern, self.config) + for layer_config in layer_config_list: + assert type(layer_config) in _EXPECTED_LAYER_CONFIG_CLASSES.values() + + @pytest.mark.parametrize( + ("layer_symbol", "config_class"), list(_EXPECTED_LAYER_CONFIG_CLASSES.items()) + ) + def test_all_symbols_map_to_layer_configs(self, layer_symbol, config_class): + layer_config_list = validate_segment_layers(layer_symbol, self.config) + + assert len(layer_config_list) == 1 + assert type(layer_config_list[0]) is config_class + assert isinstance(layer_config_list[0], TransformerConfig) + assert layer_config_list[0] is not self.config + _assert_config_contents_equal(layer_config_list[0], self.config) + assert layer_config_list[0].hidden_size == self.config.hidden_size + _assert_layer_config_types(layer_config_list, layer_symbol) + + def test_all_layer_symbols_have_an_expected_config_class(self): + assert set(_EXPECTED_LAYER_CONFIG_CLASSES) == Symbols.VALID_LAYERS + assert Symbols.PIPE not in _EXPECTED_LAYER_CONFIG_CLASSES + assert Symbols.MTP_SEPARATOR not in _EXPECTED_LAYER_CONFIG_CLASSES + + def test_repeated_layers_receive_independent_config_copies(self): + self.config.test_mutable_value = {"items": []} + + layer_config_list = validate_segment_layers("MMM", self.config) + + assert type(layer_config_list) is list + assert len({id(config) for config in layer_config_list}) == 3 + assert all(config is not self.config for config in layer_config_list) + assert all(config.test_mutable_value == {"items": []} for config in layer_config_list) + assert len({id(config.test_mutable_value) for config in layer_config_list}) == 3 + + layer_config_list[0].test_mutable_value["items"].append("changed") + assert layer_config_list[1].test_mutable_value == {"items": []} + assert self.config.test_mutable_value == {"items": []} + + def test_mla_configs_preserve_specialized_fields(self): + config = MLATransformerConfig( + num_layers=2, + hidden_size=128, + num_attention_heads=8, + q_lora_rank=32, + kv_lora_rank=16, + qk_head_dim=32, + qk_pos_emb_head_dim=16, + v_head_dim=32, + rope_type="rope", + ) + + dsa_config, mla_config = validate_segment_layers("D+", config) + + assert type(dsa_config) is DSALayerConfig + assert type(mla_config) is MLALayerConfig + assert isinstance(dsa_config, MLATransformerConfig) + assert isinstance(mla_config, MLATransformerConfig) + assert dsa_config.q_lora_rank == config.q_lora_rank + assert mla_config.kv_lora_rank == config.kv_lora_rank + assert dsa_config is not mla_config def test_invalid_symbols_cause_failure(self): """Test that invalid symbols raise ValueError.""" with pytest.raises(ValueError): - validate_segment_layers("M*X") + validate_segment_layers("M*X", self.config) with pytest.raises(ValueError): - validate_segment_layers("M|M") # pipe not valid in a segment + validate_segment_layers("M|M", self.config) # pipe not valid in a segment with pytest.raises(ValueError): - validate_segment_layers("M/M") # MTP separator not valid in a segment + validate_segment_layers("M/M", self.config) # MTP separator not valid in a segment with pytest.raises(ValueError): # Not allowed to have both standard Attention and MLA/DSA - validate_segment_layers("MDM*-") + validate_segment_layers("MDM*-", self.config) with pytest.raises(ValueError): # Not allowed to have both standard Attention and MLA (same reason # as DSA: * uses the model-level rotary_pos_emb while + uses MLA's # own decoupled RoPE). - validate_segment_layers("M+M*-") + validate_segment_layers("M+M*-", self.config) @pytest.mark.internal @@ -489,87 +591,134 @@ class TestSelectPipelineSegment: is simply the vp_stage value. """ + def setup_method(self): + self.config = _make_transformer_config() + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_single_segment_no_vp(self, mock_log): """Single segment, no VPP.""" - layer_types, offset = select_pipeline_segment("M*M*", pp_group=None, vp_stage=None) - assert layer_types == ['M', '*', 'M', '*'] + layer_configs, offset = select_pipeline_segment( + "M*M*", self.config, pp_group=None, vp_stage=None + ) + _assert_layer_config_types(layer_configs, "M*M*") assert offset == 0 + @pytest.mark.parametrize("layer_symbol", [Symbols.MLA, Symbols.DS_ATTENTION]) + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') + def test_normalizes_tp_overlap_before_copying_selected_segment(self, mock_log, layer_symbol): + self.config.tp_comm_overlap = True + + with pytest.warns(UserWarning, match="Disabling tp_comm_overlap"): + layer_configs, _ = select_pipeline_segment( + layer_symbol, self.config, pp_group=None, vp_stage=None + ) + + assert self.config.tp_comm_overlap is False + assert all(layer_config.tp_comm_overlap is False for layer_config in layer_configs) + + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') + def test_does_not_normalize_tp_overlap_for_unselected_mla_segment(self, mock_log): + self.config.tp_comm_overlap = True + + layer_configs, _ = select_pipeline_segment( + f"{Symbols.MLA}|{Symbols.MAMBA}", self.config, pp_group=None, vp_stage=1 + ) + + assert self.config.tp_comm_overlap is True + assert all(layer_config.tp_comm_overlap is True for layer_config in layer_configs) + @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_two_segments_vp0(self, mock_log): """Two segments, select first (vp_stage=0).""" - layer_types, offset = select_pipeline_segment("M-M-|M-M*-", pp_group=None, vp_stage=0) - assert layer_types == ['M', '-', 'M', '-'] + layer_configs, offset = select_pipeline_segment( + "M-M-|M-M*-", self.config, pp_group=None, vp_stage=0 + ) + _assert_layer_config_types(layer_configs, "M-M-") assert offset == 0 @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_two_segments_vp1(self, mock_log): """Two segments, select second (vp_stage=1).""" - layer_types, offset = select_pipeline_segment("M-M-|M-M*-", pp_group=None, vp_stage=1) - assert layer_types == ['M', '-', 'M', '*', '-'] + layer_configs, offset = select_pipeline_segment( + "M-M-|M-M*-", self.config, pp_group=None, vp_stage=1 + ) + _assert_layer_config_types(layer_configs, "M-M*-") assert offset == 4 @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_four_segments(self, mock_log): """Four segments, verify each vp_stage selects correctly.""" pattern = "MM|M*|M-|ME" - expected = [(['M', 'M'], 0), (['M', '*'], 2), (['M', '-'], 4), (['M', 'E'], 6)] - for vp_stage, (expected_layers, expected_offset) in enumerate(expected): - layer_types, offset = select_pipeline_segment(pattern, pp_group=None, vp_stage=vp_stage) - assert layer_types == expected_layers, f"Failed for vp_stage={vp_stage}" + expected = [("MM", 0), ("M*", 2), ("M-", 4), ("ME", 6)] + for vp_stage, (expected_pattern, expected_offset) in enumerate(expected): + layer_configs, offset = select_pipeline_segment( + pattern, self.config, pp_group=None, vp_stage=vp_stage + ) + _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_empty_segment(self, mock_log): """Empty segments are allowed for pipeline balancing.""" - layer_types, offset = select_pipeline_segment("||M*", pp_group=None, vp_stage=0) - assert layer_types == [] + layer_configs, offset = select_pipeline_segment( + "||M*", self.config, pp_group=None, vp_stage=0 + ) + assert layer_configs == [] assert offset == 0 - layer_types, offset = select_pipeline_segment("||M*", pp_group=None, vp_stage=2) - assert layer_types == ['M', '*'] + layer_configs, offset = select_pipeline_segment( + "||M*", self.config, pp_group=None, vp_stage=2 + ) + _assert_layer_config_types(layer_configs, "M*") assert offset == 0 @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_uneven_segments(self, mock_log): """Segments of different lengths.""" pattern = "MMM|M|MMMMM" - layer_types, offset = select_pipeline_segment(pattern, pp_group=None, vp_stage=0) - assert len(layer_types) == 3 + layer_configs, offset = select_pipeline_segment( + pattern, self.config, pp_group=None, vp_stage=0 + ) + assert len(layer_configs) == 3 assert offset == 0 - layer_types, offset = select_pipeline_segment(pattern, pp_group=None, vp_stage=1) - assert len(layer_types) == 1 + layer_configs, offset = select_pipeline_segment( + pattern, self.config, pp_group=None, vp_stage=1 + ) + assert len(layer_configs) == 1 assert offset == 3 - layer_types, offset = select_pipeline_segment(pattern, pp_group=None, vp_stage=2) - assert len(layer_types) == 5 + layer_configs, offset = select_pipeline_segment( + pattern, self.config, pp_group=None, vp_stage=2 + ) + assert len(layer_configs) == 5 assert offset == 4 @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_empty_main_pattern(self, mock_log): """Empty main pattern produces one empty segment.""" - layer_types, offset = select_pipeline_segment("", pp_group=None, vp_stage=None) - assert layer_types == [] + layer_configs, offset = select_pipeline_segment( + "", self.config, pp_group=None, vp_stage=None + ) + assert layer_configs == [] assert offset == 0 @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_invalid_segment_raises(self, mock_log): """Invalid layer symbols in a segment should raise ValueError.""" with pytest.raises(ValueError): - select_pipeline_segment("MX|M*", pp_group=None, vp_stage=0) + select_pipeline_segment("MX|M*", self.config, pp_group=None, vp_stage=0) @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_out_of_range_segment_raises(self, mock_log): """Segment index out of range should raise ValueError.""" with pytest.raises(ValueError, match="out of range"): - select_pipeline_segment("M*|M*", pp_group=None, vp_stage=5) + select_pipeline_segment("M*|M*", self.config, pp_group=None, vp_stage=5) @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_logging_is_called(self, mock_log): """Verify that log_on_each_pipeline_stage is called.""" - select_pipeline_segment("M*M*", pp_group=None, vp_stage=None) + select_pipeline_segment("M*M*", self.config, pp_group=None, vp_stage=None) mock_log.assert_called_once() @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') @@ -577,7 +726,12 @@ def test_logging_receives_explicit_groups(self, mock_log): tp_group = object() dp_cp_group = object() select_pipeline_segment( - "M*M*", pp_group=None, vp_stage=None, tp_group=tp_group, dp_cp_group=dp_cp_group + "M*M*", + self.config, + pp_group=None, + vp_stage=None, + tp_group=tp_group, + dp_cp_group=dp_cp_group, ) assert mock_log.call_args.kwargs["tp_group"] is tp_group assert mock_log.call_args.kwargs["dp_cp_group"] is dp_cp_group @@ -586,13 +740,17 @@ def test_logging_receives_explicit_groups(self, mock_log): def test_mutual_exclusivity_pipes_with_first_stage(self, mock_log): """Pipe separators + first_stage_layers should raise ValueError.""" with pytest.raises(ValueError, match="Cannot specify"): - select_pipeline_segment("M*|M*", pp_group=None, vp_stage=0, first_stage_layers=1) + select_pipeline_segment( + "M*|M*", self.config, pp_group=None, vp_stage=0, first_stage_layers=1 + ) @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_mutual_exclusivity_pipes_with_last_stage(self, mock_log): """Pipe separators + last_stage_layers should raise ValueError.""" with pytest.raises(ValueError, match="Cannot specify"): - select_pipeline_segment("M*|M*", pp_group=None, vp_stage=0, last_stage_layers=1) + select_pipeline_segment( + "M*|M*", self.config, pp_group=None, vp_stage=0, last_stage_layers=1 + ) @patch('megatron.core.models.hybrid.hybrid_layer_allocation.log_on_each_pipeline_stage') def test_segment_count_not_divisible_by_pp_size(self, mock_log): @@ -603,7 +761,7 @@ def test_segment_count_not_divisible_by_pp_size(self, mock_log): patch('torch.distributed.get_world_size', return_value=2), ): with pytest.raises(ValueError, match="evenly divisible"): - select_pipeline_segment("M|M|M", pp_group=mock_group, vp_stage=None) + select_pipeline_segment("M|M|M", self.config, pp_group=mock_group, vp_stage=None) @pytest.mark.internal @@ -614,6 +772,9 @@ class TestSelectPipelineSegmentLegacyFallback: activates when the pattern has no pipe separators but pp_size > 1. """ + def setup_method(self): + self.config = _make_transformer_config() + def _call_for_rank( self, pattern, @@ -633,6 +794,7 @@ def _call_for_rank( ): return select_pipeline_segment( pattern, + self.config, pp_group=mock_group, vp_stage=vp_stage, first_stage_layers=first_stage_layers, @@ -642,11 +804,11 @@ def _call_for_rank( def test_even_split_2_ranks(self): """4 layers across 2 ranks -> 2 each.""" layers0, off0 = self._call_for_rank("M*M-", pp_rank=0, pp_size=2) - assert layers0 == ['M', '*'] + _assert_layer_config_types(layers0, "M*") assert off0 == 0 layers1, off1 = self._call_for_rank("M*M-", pp_rank=1, pp_size=2) - assert layers1 == ['M', '-'] + _assert_layer_config_types(layers1, "M-") assert off1 == 2 def test_even_split_4_ranks(self): @@ -736,7 +898,7 @@ def test_deprecation_warning_logged(self): 'megatron.core.models.hybrid.hybrid_layer_allocation.log_single_rank' ) as mock_warn, ): - select_pipeline_segment("M*M*", pp_group=mock_group, vp_stage=None) + select_pipeline_segment("M*M*", self.config, pp_group=mock_group, vp_stage=None) mock_warn.assert_called_once() call_args = mock_warn.call_args assert "DEPRECATION" in call_args[0][2] @@ -750,7 +912,7 @@ def test_all_ranks_cover_full_pattern(self): layers, offset = self._call_for_rank(pattern, pp_rank=rank, pp_size=pp_size) assert offset == len(all_layers) all_layers.extend(layers) - assert all_layers == ['M', '*', 'M', '*', 'M', '*'] + _assert_layer_config_types(all_layers, "M*M*M*") @pytest.mark.internal @@ -759,9 +921,11 @@ class TestGetLayerMapsFromLayerTypeList: def test_standard_layer_types(self): """Standard symbols each produce a single-entry map at local index 0.""" - maps = get_layer_maps_from_layer_type_list(["*", "M", "-", "E"]) - # We always get all symbols returned, not only those contained in the pattern. - assert len(maps) == 7 + maps = get_layer_maps_from_layer_type_list( + [Symbols.ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE] + ) + # We always get all symbols, not only those contained in the pattern. + assert len(maps) == len(Symbols.VALID_LAYERS) attention_map, mamba_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE )(maps) @@ -771,8 +935,10 @@ def test_standard_layer_types(self): assert moe_map == {3: 0} def test_dsa(self): - """D (DSA) layers are treated as separate layers for KV cache mapping.""" - maps = get_layer_maps_from_layer_type_list(["D", "M", "D", "M"]) + """DSA layers have their own local cache indices.""" + maps = get_layer_maps_from_layer_type_list( + [Symbols.DS_ATTENTION, Symbols.MAMBA, Symbols.DS_ATTENTION, Symbols.MAMBA] + ) attention_map, dsa_map, mamba_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.DS_ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE )(maps) @@ -783,8 +949,10 @@ def test_dsa(self): assert moe_map == {} def test_mixed_attention_and_dsa(self): - """Both * and D contribute to the different maps with non-consecutive local indices.""" - maps = get_layer_maps_from_layer_type_list(["*", "D", "M", "-"]) + """Attention and DSA layers maintain separate local indices.""" + maps = get_layer_maps_from_layer_type_list( + [Symbols.ATTENTION, Symbols.DS_ATTENTION, Symbols.MAMBA, Symbols.MLP] + ) attention_map, dsa_map, mamba_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.DS_ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE )(maps) @@ -795,8 +963,8 @@ def test_mixed_attention_and_dsa(self): assert moe_map == {} def test_all_mamba(self): - """All-mamba pattern leaves attention, mlp, and moe maps empty.""" - maps = get_layer_maps_from_layer_type_list(["M", "M", "M"]) + """All-Mamba patterns leave the other maps empty.""" + maps = get_layer_maps_from_layer_type_list([Symbols.MAMBA] * 3) attention_map, mamba_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE )(maps) @@ -806,8 +974,10 @@ def test_all_mamba(self): assert moe_map == {} def test_mla(self): - """+ (MLA) layers are mapped independently of other attention types.""" - maps = get_layer_maps_from_layer_type_list(["+", "M", "+", "M"]) + """MLA layers have their own local cache indices.""" + maps = get_layer_maps_from_layer_type_list( + [Symbols.MLA, Symbols.MAMBA, Symbols.MLA, Symbols.MAMBA] + ) attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.DS_ATTENTION, @@ -824,8 +994,10 @@ def test_mla(self): assert moe_map == {} def test_mixed_dsa_and_mla(self): - """D and + can coexist (both are MLA-based and use decoupled RoPE).""" - maps = get_layer_maps_from_layer_type_list(["D", "+", "M", "-"]) + """DSA and MLA layers maintain separate local indices.""" + maps = get_layer_maps_from_layer_type_list( + [Symbols.DS_ATTENTION, Symbols.MLA, Symbols.MAMBA, Symbols.MLP] + ) attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.DS_ATTENTION, diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index 2eb2ca49d27..e9db6a7a81d 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -1054,21 +1054,21 @@ def get_pg_collection(): return ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'pp', 'cp']) def get_mamba_block(hybrid_layer_pattern): - layer_type_list = validate_segment_layers(hybrid_layer_pattern) transformer_config = TransformerConfig( 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_type_list), + num_layers=len(hybrid_layer_pattern), num_attention_heads=4, use_cpu_initialization=True, cuda_graph_impl="local", ) + layer_config_list = validate_segment_layers(hybrid_layer_pattern, transformer_config) modules = hybrid_stack_spec.submodules return HybridStack( transformer_config, modules, - layer_type_list=layer_type_list, + layer_config_list=layer_config_list, pp_layer_offset=0, pg_collection=get_pg_collection(), ) diff --git a/tests/unit_tests/transformer/test_transformer_config.py b/tests/unit_tests/transformer/test_transformer_config.py index e13b1b78986..27938dc7ab4 100644 --- a/tests/unit_tests/transformer/test_transformer_config.py +++ b/tests/unit_tests/transformer/test_transformer_config.py @@ -59,6 +59,32 @@ def test_gdp_num_householder_accepts_positive_values(): assert config.gdp_num_householder == 5 +def test_from_config_creates_independent_target_config_without_reinitializing(): + class LayerConfig(TransformerConfig): + + def __post_init__(self): + raise AssertionError("from_config must not reinitialize the target config") + + config = TransformerConfig(num_layers=1, hidden_size=128, num_attention_heads=4) + config.dynamic_value = {"items": []} + config.dynamic_alias = config.dynamic_value + config.self_reference = config + config.state_reference = config.__dict__ + + layer_config = LayerConfig.from_config(config) + + assert type(layer_config) is LayerConfig + assert vars(layer_config).keys() == vars(config).keys() + assert layer_config.dynamic_value == config.dynamic_value + assert layer_config.dynamic_value is not config.dynamic_value + assert layer_config.dynamic_alias is layer_config.dynamic_value + assert layer_config.self_reference is layer_config + assert layer_config.state_reference is layer_config.__dict__ + + layer_config.dynamic_value["items"].append("changed") + assert config.dynamic_value == {"items": []} + + @pytest.mark.parametrize("num_householder", [0, -1]) def test_gdp_num_householder_rejects_non_positive_values(num_householder: int): with pytest.raises(ValueError, match="gdp_num_householder must be positive"):