diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index a270161ddd6..e29265c04f4 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -4,7 +4,7 @@ import warnings from abc import abstractmethod from functools import partial -from typing import Optional, Protocol, cast +from typing import Literal, Optional, Protocol, cast from megatron.core.extensions.transformer_engine import ( TEColumnParallelGroupedLinear, @@ -200,3 +200,19 @@ def grouped_mlp_modules(self, moe_use_grouped_gemm: bool) -> ExpertsBuilder: activation_func=self.activation_func(), ), ) + + +def get_backend( + transformer_impl: Literal["local", "transformer_engine", "inference_optimized"] +) -> BackendSpecProvider: + """Return the backend that's selected with the given `transformer_impl`.""" + if transformer_impl == "transformer_engine": + from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider + + return TESpecProvider() + elif transformer_impl == "inference_optimized": + return InferenceSpecProvider() + elif transformer_impl == "local": + return LocalSpecProvider() + else: + raise ValueError(f"unknown transformer_impl='{transformer_impl}'") diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index 22322e1b346..0042cbea010 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -5,6 +5,7 @@ # This source code is licensed under the Apache license found in the # LICENSE file in the root directory of this source tree. +import copy from contextlib import nullcontext from dataclasses import dataclass from typing import Optional, Tuple, Union @@ -15,7 +16,7 @@ from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe -from megatron.core.extensions.transformer_engine import TENorm +from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear, TENorm from megatron.core.fp4_utils import get_fp4_context from megatron.core.fp8_utils import get_fp8_context from megatron.core.inference.contexts import BaseInferenceContext @@ -28,6 +29,7 @@ from megatron.core.transformer.cuda_graphs import annotate_first_last_layer from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.module import MegatronModule +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 from megatron.core.transformer.utils import sharded_state_dict_default @@ -44,6 +46,7 @@ class HybridStackSubmodules: gdn_layer: Union[ModuleSpec, type] = IdentityOp attention_layer: Union[ModuleSpec, type] = IdentityOp dsa_layer: Union[ModuleSpec, type] = IdentityOp + mla_layer: Union[ModuleSpec, type] = IdentityOp mlp_layer: Union[ModuleSpec, type] = IdentityOp moe_layer: Union[ModuleSpec, type] = IdentityOp mtp_block_spec: Optional[ModuleSpec] = None @@ -114,6 +117,9 @@ def __init__( ) self.layer_type_list = layer_type_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): @@ -156,6 +162,16 @@ 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: + layer = build_module( + submodules.mla_layer, + config=self.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: layer = build_module( submodules.mlp_layer, @@ -202,6 +218,26 @@ def __init__( eps=self.config.layernorm_epsilon, ) + 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. + submodules = copy.deepcopy(submodules) + mla_spec = submodules.mla_layer + # We always fuse the input layernorm because Hybrid always uses TransformerEngine. + mla_spec.submodules.input_layernorm = IdentityOp + mla_spec.submodules.self_attention.module = FusedMLASelfAttention + mla_spec.submodules.self_attention.submodules.linear_qkv_down_proj = ( + TELayerNormColumnParallelLinear + ) + mla_spec.submodules.self_attention.submodules.linear_q_down_proj = None + mla_spec.submodules.self_attention.submodules.linear_kv_down_proj = None + mla_spec.submodules.sharded_state_dict_keys_map = { + "self_attention.linear_q_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_kv_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_qkv_down_proj.layer_norm_": "input_layernorm.", + } + return submodules + def set_input_tensor(self, input_tensor: Tensor): """Set input tensor to be used instead of forward()'s input. diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index 67103fe67f1..83a6163b88d 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -18,11 +18,12 @@ class Symbols: GDN = 'G' ATTENTION = "*" DS_ATTENTION = "D" + MLA = "+" MLP = "-" MOE = 'E' PIPE = '|' MTP_SEPARATOR = "/" - VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, MLP, MOE} + VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, MLA, MLP, MOE} @classmethod def name_sorted_valid_layer_symbols(cls) -> list[str]: @@ -293,7 +294,7 @@ def _validate_pattern(pattern: str, pattern_name: str, allow_pipe: bool = False) ) # Disallow Attention + MLA/DSA hybridity. - if Symbols.ATTENTION in pattern and Symbols.DS_ATTENTION in pattern: + if Symbols.ATTENTION in pattern and (Symbols.DS_ATTENTION in pattern or Symbols.MLA in pattern): raise ValueError("Not supported to have both Attention and MLA/DSA in one model") @@ -321,7 +322,7 @@ def validate_segment_layers(segment: str) -> List[str]: ) # Disallow Attention + MLA/DSA hybridity. - if Symbols.ATTENTION in segment and Symbols.DS_ATTENTION in segment: + 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") return layer_type_list diff --git a/megatron/core/models/hybrid/hybrid_layer_specs.py b/megatron/core/models/hybrid/hybrid_layer_specs.py index e1624293b5a..03fef58159f 100755 --- a/megatron/core/models/hybrid/hybrid_layer_specs.py +++ b/megatron/core/models/hybrid/hybrid_layer_specs.py @@ -169,6 +169,28 @@ self_attn_bda=get_bias_dropout_add, ), ), + mla_layer=ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=TENorm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=TEColumnParallelLinear, + linear_q_down_proj=TELinear, + linear_q_up_proj=TEColumnParallelLinear, + linear_kv_down_proj=TELinear, + linear_kv_up_proj=TEColumnParallelLinear, + core_attention=TEDotProductAttention, + linear_proj=TERowParallelLinear, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + ), + ), # Started with spec from gpt_layer_specs.py # Using the TE spec because we had problems getting the non-TE spec # working @@ -264,6 +286,28 @@ self_attn_bda=get_bias_dropout_add, ), ), + mla_layer=ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=TENorm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=TEColumnParallelLinear, + linear_q_down_proj=TELinear, + linear_q_up_proj=TEColumnParallelLinear, + linear_kv_down_proj=TELinear, + linear_kv_up_proj=TEColumnParallelLinear, + core_attention=TEDotProductAttention, + linear_proj=InferenceRowParallelLinear, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + ), + ), # Started with spec from gpt_layer_specs.py # Using the TE spec because we had problems getting the non-TE spec # working diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index fccf674d785..e0b6af7aa7f 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -35,6 +35,7 @@ ) from megatron.core.transformer.attention import Attention from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.mla_qk_norm_config import QKNormConfigResolver from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.core.utils import deprecate_inference_params, get_pg_size, is_te_min_version @@ -162,6 +163,10 @@ def __init__( name=name, ) + # Resolve which classes to use for Q and KV linear up projections and norms, based on + # QK-norm selection. + layer_classes = QKNormConfigResolver(self.config, submodules).resolve() + assert not config.add_bias_linear, "add_bias_linear is not supported for AbsorbedMLA" assert not ( config.tensor_model_parallel_size > 1 and not config.sequence_parallel @@ -260,7 +265,7 @@ def __init__( if self.config.q_lora_rank is None: # Not projecting query self.linear_q_proj = build_module( - submodules.linear_q_proj, + layer_classes["linear_q_proj"], self.config.hidden_size, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -306,7 +311,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -353,7 +358,7 @@ def __init__( ) self.linear_kv_up_proj = build_module( - submodules.linear_kv_up_proj, + layer_classes["linear_kv_up_proj"], self.config.kv_lora_rank, self.config.num_attention_heads * (self.config.qk_head_dim + self.config.v_head_dim), config=self.config, @@ -369,14 +374,14 @@ def __init__( if self.config.q_lora_rank is not None: self.q_layernorm = build_module( - submodules.q_layernorm, + layer_classes["q_layernorm"], hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) self.kv_layernorm = build_module( - submodules.kv_layernorm, + layer_classes["kv_layernorm"], hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, diff --git a/megatron/core/transformer/mla_qk_norm_config.py b/megatron/core/transformer/mla_qk_norm_config.py new file mode 100644 index 00000000000..e14066a105d --- /dev/null +++ b/megatron/core/transformer/mla_qk_norm_config.py @@ -0,0 +1,293 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +""" +Resolve MLA and DSA Q/KV norm configuration from a layer specification. +""" + +from typing import NoReturn + +from megatron.core.models.backends import get_backend +from megatron.core.transformer.identity_op import IdentityOp +from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.torch_norm import LayerNormBuilder +from megatron.core.transformer.transformer_config import MLATransformerConfig + +__all__ = [] + +_QKNormResolvedConfig = dict[str, ModuleSpec | type | LayerNormBuilder] + + +class QKNormConfigResolver: + """Validate and resolve Q/KV norm placement for MLA and DSA. + + Q/KV norm can be represented either by a standalone norm module or by + a fused norm+linear projection. MLA can use the fused form; DSA cannot + because it needs the normalized Q/KV values outside the projection. + + Constraints: + - `qk_l2_norm` is unsupported for MLA/DSA. + - A standalone Q norm is only usable when `q_lora_rank` is set. + - Explicit norm modules cannot be paired with fused norm+linear projections. + - Disabled QK norm rejects both explicit norms and fused norm+linear projections. + - DSA with QK norm requires non-fused projections and standalone Q/KV norms. + """ + + def __init__(self, config: MLATransformerConfig, submodules) -> None: + """Capture the configuration, requested modules, and backend implementations.""" + self.config = config + self.submodules = submodules + self.has_q_lora = config.q_lora_rank is not None + self.is_dsa = config.experimental_attention_variant == "dsa" + self.variant_str = "DSA" if self.is_dsa else "MLA" + + backend = get_backend(config.transformer_impl) + self.qk_norm_impl = backend.layer_norm( + rms_norm=config.normalization == "RMSNorm", for_qk=True + ) + self.linear_impl = backend.column_parallel_linear() + self.fused_norm_linear_impl = backend.column_parallel_layer_norm_linear() + + def resolve(self) -> _QKNormResolvedConfig: + """Validate the specification and return the modules to instantiate. + + Returns: + The Q/KV norms and projections after applying the MLA or DSA constraints. + + Raises: + ValueError: If the requested norm placement is unsupported or conflicting. + """ + if self.config.qk_l2_norm: + raise ValueError(f"qk_l2_norm is not supported with {self.variant_str}.") + + self._reject_common_spec_conflicts() + if not self.config.qk_layernorm: + return self._resolve_disabled_qk_layernorm() + if self.is_dsa: + return self._resolve_dsa_qk_layernorm() + return self._resolve_mla_qk_layernorm() + + def _resolve_disabled_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve projections when Q/KV normalization is disabled. + + Explicit norm modules and fused norm-linear projections are rejected because + they would still introduce Q/KV normalization. + """ + linear_q_proj_cls = IdentityOp + linear_q_up_proj_cls = IdentityOp + + if self.has_q_lora: + self._reject_disabled_norm( + self.submodules.linear_q_up_proj, + self.submodules.q_layernorm, + "linear_q_up_proj", + "q_layernorm", + ) + linear_q_up_proj_cls = self.submodules.linear_q_up_proj or self.linear_impl + else: + if self._is_fused_norm_linear(self.submodules.linear_q_proj): + raise ValueError( + f"spec sets linear_q_proj={self.submodules.linear_q_proj}, but " + "qk_layernorm/qk_l2_norm are supposed to be disabled" + ) + linear_q_proj_cls = self.submodules.linear_q_proj or self.linear_impl + + self._reject_disabled_norm( + self.submodules.linear_kv_up_proj, + self.submodules.kv_layernorm, + "linear_kv_up_proj", + "kv_layernorm", + ) + return self._result( + linear_q_proj=linear_q_proj_cls, + linear_q_up_proj=linear_q_up_proj_cls, + linear_kv_up_proj=self.submodules.linear_kv_up_proj or self.linear_impl, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ) + + def _resolve_dsa_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve DSA's standalone Q/KV norms and non-fused projections. + + DSA consumes the normalized Q/KV values outside the projection, so it cannot + use fused norm-linear projections. + """ + if not self.has_q_lora: + raise ValueError( + "`qk_layernorm=True` with `q_lora_rank is None` is not supported for DSA " + "because DSA cannot fuse Q norm into `linear_q_proj`." + ) + + return self._result( + linear_q_proj=IdentityOp, + linear_q_up_proj=self._dsa_linear_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ), + linear_kv_up_proj=self._dsa_linear_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ), + q_layernorm=self._default_if_trivial(self.submodules.q_layernorm, self.qk_norm_impl), + kv_layernorm=self._default_if_trivial(self.submodules.kv_layernorm, self.qk_norm_impl), + ) + + def _resolve_mla_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve MLA norms, fusing them into projections when no norm is explicit.""" + q_norm_cls = self.submodules.q_layernorm or IdentityOp + linear_q_proj_cls = IdentityOp + linear_q_up_proj_cls = IdentityOp + + if self.has_q_lora: + if self._is_trivial(q_norm_cls): + linear_q_up_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ) + else: + linear_q_up_proj_cls = self._non_fused_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ) + else: + linear_q_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_q_proj, "linear_q_proj" + ) + + kv_norm_cls = self.submodules.kv_layernorm or IdentityOp + if self._is_trivial(kv_norm_cls): + linear_kv_up_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ) + else: + linear_kv_up_proj_cls = self._non_fused_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ) + + return self._result( + linear_q_proj=linear_q_proj_cls, + linear_q_up_proj=linear_q_up_proj_cls, + linear_kv_up_proj=linear_kv_up_proj_cls, + q_layernorm=q_norm_cls, + kv_layernorm=kv_norm_cls, + ) + + def _reject_common_spec_conflicts(self) -> None: + """Reject conflicts that apply regardless of the selected attention variant.""" + if not self.has_q_lora and not self._is_trivial(self.submodules.q_layernorm): + self._raise_unused_q_norm() + if self.has_q_lora: + self._reject_explicit_norm_with_fused_linear( + self.submodules.linear_q_up_proj, + self.submodules.q_layernorm, + "linear_q_up_proj", + "q_layernorm", + ) + self._reject_explicit_norm_with_fused_linear( + self.submodules.linear_kv_up_proj, + self.submodules.kv_layernorm, + "linear_kv_up_proj", + "kv_layernorm", + ) + + def _reject_disabled_norm(self, module_spec, norm_spec, module_name, norm_name) -> None: + """Reject a norm module or fused projection when Q/KV norm is disabled.""" + if self._is_fused_norm_linear(module_spec) or not self._is_trivial(norm_spec): + raise ValueError( + f"spec sets {module_name}={module_spec} and " + f"{norm_name}={norm_spec}, but " + "qk_layernorm/qk_l2_norm are supposed to be disabled" + ) + + def _reject_explicit_norm_with_fused_linear( + self, module_spec, norm_spec, module_name, norm_name + ) -> None: + """Reject specifying the same norm both explicitly and inside a projection.""" + if not self._is_trivial(norm_spec) and self._is_fused_norm_linear(module_spec): + raise ValueError( + f"`{norm_name}={norm_spec}` is non-trivial " + f"and `{module_name}={module_spec}` is a " + f"fused norm+linear; either unset `{norm_name}` or use a " + f"linear layer without norm fusion for `{module_name}`" + ) + + def _non_fused_or_default(self, module_spec, module_name): + """Return a linear implementation, requiring it not to fuse normalization.""" + linear_cls = module_spec or self.linear_impl + self._require_linear(linear_cls, module_name) + if self._is_fused_norm_linear(linear_cls): + raise ValueError( + f"`{module_name}={module_spec}` is fused norm+linear, but a non-fused linear " + f"is required" + ) + return linear_cls + + def _dsa_linear_or_default(self, module_spec, module_name): + """Return DSA's non-fused projection implementation. + + This uses a DSA-specific diagnostic so the rejected constraint is clear. + """ + linear_cls = module_spec or self.linear_impl + self._require_linear(linear_cls, module_name) + if self._is_fused_norm_linear(linear_cls): + raise ValueError( + f"`{module_name}={module_spec}` is fused norm+linear, " + f"which is not supported for DSA." + ) + return linear_cls + + def _mla_fused_linear_or_default(self, module_spec, module_name): + """Return a fused MLA projection, using the backend default when available.""" + if self._is_fused_norm_linear(module_spec): + return module_spec + return self._require_linear(self.fused_norm_linear_impl, module_name) + + def _require_linear(self, module_spec, module_name): + """Return a configured projection or report that no viable implementation exists.""" + if module_spec is None: + raise RuntimeError( + "qk_layernorm requires TransformerEngine or " + "q_layernorm/kv_layernorm to be set in the spec " + f"to build `{module_name}`." + ) + return module_spec + + def _raise_unused_q_norm(self) -> NoReturn: + """Report an explicit Q norm that has no Q-LoRA projection to consume it.""" + help_msg = "" + if not self._is_fused_norm_linear(self.submodules.linear_q_proj): + help_msg = ( + f"Please use a fused norm+linear for " + f"`linear_q_proj={self.submodules.linear_q_proj}` if " + f"you intend to have a Q-norm." + ) + raise ValueError( + f"`q_layernorm={self.submodules.q_layernorm}` is non-trivial, " + f"but `q_lora_rank is None`, meaning it will not be used." + f"{help_msg}" + ) + + def _is_fused_norm_linear(self, module_spec) -> bool: + """Return whether a module specification selects the backend fused projection.""" + module_cls = module_spec.module if isinstance(module_spec, ModuleSpec) else module_spec + return self.fused_norm_linear_impl is not None and module_cls is self.fused_norm_linear_impl + + @staticmethod + def _is_trivial(module_spec) -> bool: + """Return whether a norm slot is unset or explicitly an identity operation.""" + return module_spec in (None, IdentityOp) + + @classmethod + def _default_if_trivial(cls, module_spec, default): + """Replace an unset or identity specification with the supplied default.""" + if cls._is_trivial(module_spec): + return default + return module_spec + + @staticmethod + def _result( + *, linear_q_proj, linear_q_up_proj, linear_kv_up_proj, q_layernorm, kv_layernorm + ) -> _QKNormResolvedConfig: + """Package the resolved Q/KV norms and projections in the caller's schema.""" + return dict( + linear_q_proj=linear_q_proj, + linear_q_up_proj=linear_q_up_proj, + linear_kv_up_proj=linear_kv_up_proj, + q_layernorm=q_layernorm, + kv_layernorm=kv_layernorm, + ) diff --git a/megatron/core/transformer/multi_latent_attention.py b/megatron/core/transformer/multi_latent_attention.py index 202034986db..50e11151dcd 100644 --- a/megatron/core/transformer/multi_latent_attention.py +++ b/megatron/core/transformer/multi_latent_attention.py @@ -36,6 +36,7 @@ ) from megatron.core.transformer.attention import Attention, LinearProjBuilder from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.mla_qk_norm_config import QKNormConfigResolver from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import MLATransformerConfig @@ -499,10 +500,14 @@ def __init__( name=name, ) + # Resolve which classes to use for Q and KV linear up projections and norms, based on + # QK-norm selection. + layer_classes = self._resolve_qk_norm_config(submodules) + if self.config.q_lora_rank is None: # Not projecting query self.linear_q_proj = build_module( - submodules.linear_q_proj, + layer_classes["linear_q_proj"], self.config.hidden_size, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -549,7 +554,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -596,7 +601,7 @@ def __init__( ) self.linear_kv_up_proj = build_module( - submodules.linear_kv_up_proj, + layer_classes["linear_kv_up_proj"], self.config.kv_lora_rank, self.config.num_attention_heads * (self.config.qk_head_dim + self.config.v_head_dim), config=self.config, @@ -611,18 +616,24 @@ def __init__( ) if self.config.q_lora_rank is not None: - self.q_layernorm = submodules.q_layernorm( + self.q_layernorm = layer_classes["q_layernorm"]( hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) - self.kv_layernorm = submodules.kv_layernorm( + self.kv_layernorm = layer_classes["kv_layernorm"]( hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) + def _resolve_qk_norm_config( + self, submodules + ) -> dict[str, ModuleSpec | type | LayerNormBuilder]: + """Resolve which Q/KV norm and up-projection implementations to build.""" + return QKNormConfigResolver(self.config, submodules).resolve() + def _qkv_down_projection(self, hidden_states): """Unfused q/kv down projection path.""" if self.config.q_lora_rank is not None: @@ -1250,6 +1261,9 @@ def __init__( "FusedMLASelfAttention requires q_lora_rank to be set; " "fallback to MLASelfAttention for q_lora_rank=None." ) + # Resolve which linear class to use for Q and KV up projections, + # based on QK-norm selection. + layer_classes = self._resolve_qk_norm_config(submodules) qkv_down_proj_kwargs = {} if submodules.linear_qkv_down_proj in [TELinear]: @@ -1285,7 +1299,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -1300,7 +1314,7 @@ def __init__( ) self.linear_kv_up_proj = build_module( - submodules.linear_kv_up_proj, + layer_classes["linear_kv_up_proj"], self.config.kv_lora_rank, self.config.num_attention_heads * (self.config.qk_head_dim + self.config.v_head_dim), config=self.config, @@ -1314,12 +1328,12 @@ def __init__( name=(name + ".linear_kv_up_proj") if name is not None else None, ) - self.q_layernorm = submodules.q_layernorm( + self.q_layernorm = layer_classes["q_layernorm"]( hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) - self.kv_layernorm = submodules.kv_layernorm( + self.kv_layernorm = layer_classes["kv_layernorm"]( hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 8174fbd5eb3..591875213b2 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -861,7 +861,10 @@ def validate_args(args, defaults={}): ) # Infer use of MLA from unified pattern - if args.hybrid_layer_pattern and Symbols.DS_ATTENTION in args.hybrid_layer_pattern: + if args.hybrid_layer_pattern and ( + Symbols.MLA in args.hybrid_layer_pattern + or Symbols.DS_ATTENTION in args.hybrid_layer_pattern + ): args.multi_latent_attention = True # === End of hybrid layer pattern: deprecation handling and validation === diff --git a/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py b/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py index 3151b42d22d..01bda68b4ca 100644 --- a/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py +++ b/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py @@ -453,7 +453,12 @@ def test_mtp_layer_overlap(self, dispatcher_type, flex_backend, fp8_flag): Verifies all-to-all overlap optimization in MTP layer produces the same results as the reference implementation. """ - extra_kwargs = {"mtp_num_layers": 1, "mtp_loss_scaling_factor": 1.1} + qk_layernorm = True + extra_kwargs = { + "mtp_num_layers": 1, + "mtp_loss_scaling_factor": 1.1, + "qk_layernorm": qk_layernorm, + } apply_flex_backend_kwargs(extra_kwargs, dispatcher_type, flex_backend) if fp8_flag is not None: extra_kwargs["fp8_recipe"] = fp8_flag[1] @@ -466,7 +471,7 @@ def test_mtp_layer_overlap(self, dispatcher_type, flex_backend, fp8_flag): transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( num_experts=16, moe_grouped_gemm=True, - qk_layernorm=True, + qk_layernorm=qk_layernorm, multi_latent_attention=True, ) mtp_block_spec = get_gpt_mtp_block_spec(config, transformer_layer_spec, True) diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index 80d143a89a3..372f8d0d293 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -220,6 +220,7 @@ def test_update_router_qb_beta_skips_eval(self): class TestAllReduceLNGrads: def init_model(self, share_embeddings_and_output_weights: bool = False): + qk_layernorm = True self.transformer_config = TransformerConfig( num_layers=2, hidden_size=12, @@ -227,13 +228,15 @@ def init_model(self, share_embeddings_and_output_weights: bool = False): use_cpu_initialization=True, tensor_model_parallel_size=self.tp_size, pipeline_model_parallel_size=self.pp_size, - qk_layernorm=True, + qk_layernorm=qk_layernorm, pipeline_dtype=torch.float32, ) self.model = GPTModel( config=self.transformer_config, - transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(qk_layernorm=True), + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec( + qk_layernorm=qk_layernorm + ), vocab_size=100, max_sequence_length=4, share_embeddings_and_output_weights=share_embeddings_and_output_weights, diff --git a/tests/unit_tests/models/test_hybrid_model.py b/tests/unit_tests/models/test_hybrid_model.py index ffc9fe41e99..95bcaa2d7d0 100644 --- a/tests/unit_tests/models/test_hybrid_model.py +++ b/tests/unit_tests/models/test_hybrid_model.py @@ -1,9 +1,12 @@ # Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. +import dataclasses +import functools import os from datetime import timedelta from itertools import accumulate from types import SimpleNamespace +from unittest.mock import patch import pytest import torch @@ -22,12 +25,75 @@ from megatron.core.models.hybrid.hybrid_model import HybridModel, _hybrid_logging_pg_kwargs from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed -from megatron.core.transformer import TransformerConfig +from megatron.core.transformer import MLATransformerConfig, TransformerConfig 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 from tests.unit_tests.test_utilities import Utils +try: + from fast_hadamard_transform import hadamard_transform as _hadamard_transform + + _HAVE_HADAMARD = True +except ImportError: + _HAVE_HADAMARD = False + _hadamard_transform = None + + +def _mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor: + """Identity-with-scale stand-in for `fast_hadamard_transform.hadamard_transform`. + + Mirrors the helper in `tests/unit_tests/transformer/experimental_attention_variant/ + test_attention_variant_dsa.py` so that DSA forward tests run in containers that + don't ship the upstream library. + """ + return x * scale + + +def _is_dataclass_instance(value): + return dataclasses.is_dataclass(value) and not isinstance(value, type) + + +def _assert_equal_with_partial_contents(left, right, path="root"): + """Assert recursive equality while comparing `partial` objects structurally.""" + if isinstance(left, functools.partial) or isinstance(right, functools.partial): + assert isinstance(left, functools.partial), f"{path}: left is not `partial`" + assert isinstance(right, functools.partial), f"{path}: right is not `partial`" + _assert_equal_with_partial_contents(left.func, right.func, f"{path}.func") + _assert_equal_with_partial_contents(left.args, right.args, f"{path}.args") + _assert_equal_with_partial_contents( + left.keywords or {}, right.keywords or {}, f"{path}.keywords" + ) + return + + if _is_dataclass_instance(left) or _is_dataclass_instance(right): + assert _is_dataclass_instance(left), f"{path}: left is not a dataclass" + assert _is_dataclass_instance(right), f"{path}: right is not a dataclass" + assert type(left) is type(right), f"{path}: dataclass types differ" + for field in dataclasses.fields(left): + if field.compare: + _assert_equal_with_partial_contents( + getattr(left, field.name), getattr(right, field.name), f"{path}.{field.name}" + ) + return + + if isinstance(left, dict) or isinstance(right, dict): + assert isinstance(left, dict), f"{path}: left is not a dict" + assert isinstance(right, dict), f"{path}: right is not a dict" + assert left.keys() == right.keys(), f"{path}: dict keys differ" + for key in left: + _assert_equal_with_partial_contents(left[key], right[key], f"{path}[{key!r}]") + return + + if isinstance(left, (list, tuple)) or isinstance(right, (list, tuple)): + assert type(left) is type(right), f"{path}: sequence types differ" + assert len(left) == len(right), f"{path}: sequence lengths differ" + for index, (left_item, right_item) in enumerate(zip(left, right)): + _assert_equal_with_partial_contents(left_item, right_item, f"{path}[{index}]") + return + + assert left == right, f"{path}: values differ" + def test_hybrid_logging_process_groups_are_paired(): tp_group = object() @@ -329,6 +395,12 @@ def test_layer_numbers(self): class TestHybridQKLayernorm: + # Subclasses override these to retarget the same tests at MLA's + # `mla_layer.kv_layernorm` or DSA's `dsa_layer.kv_layernorm`. The base class + # exercises the SelfAttention path with `attention_layer.k_layernorm`. + _attention_layer_attr = 'attention_layer' + _k_norm_attr = 'k_layernorm' + def setup_method(self, method): Utils.initialize_model_parallel(1, 1) model_parallel_cuda_manual_seed(123) @@ -336,7 +408,9 @@ def setup_method(self, method): def teardown_method(self, method): Utils.destroy_model_parallel() - def _build_model(self, **config_overrides): + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec config = TransformerConfig( num_layers=3, hidden_size=256, @@ -346,26 +420,32 @@ def _build_model(self, **config_overrides): ) return HybridModel( config=config, - hybrid_stack_spec=hybrid_stack_spec, + hybrid_stack_spec=spec, vocab_size=100, max_sequence_length=4, hybrid_layer_pattern="M*-", ) def _get_attention_layer(self, model): - """Return the SelfAttention submodule from the attention layer.""" + """Return the self-attention submodule that owns a `q_layernorm`.""" for layer in model.decoder.layers: if hasattr(layer, 'self_attention') and hasattr(layer.self_attention, 'q_layernorm'): return layer.self_attention return None - def test_no_qk_norm_by_default(self): - """Without qk_layernorm, attention has no q/k layernorm.""" + def _get_k_norm(self, attn): + return getattr(attn, self._k_norm_attr) + + def test_trivial_qk_norm_by_default(self): + """Without qk_layernorm, attention has trivial q/k layernorm.""" + from megatron.core.transformer.identity_op import IdentityOp + model = self._build_model() attn = self._get_attention_layer(model) assert attn is not None - assert attn.q_layernorm is None - assert attn.k_layernorm is None + assert attn.q_layernorm is None or isinstance(attn.q_layernorm, IdentityOp) + k_norm = self._get_k_norm(attn) + assert k_norm is None or isinstance(k_norm, IdentityOp) def test_qk_layernorm_from_config(self): """config.qk_layernorm=True creates q/k layernorm even with static spec.""" @@ -375,7 +455,7 @@ def test_qk_layernorm_from_config(self): # TENorm is a factory (__new__ returns a TE LayerNorm/RMSNorm), so we # verify the norm was created rather than checking for a specific type. assert attn.q_layernorm is not None - assert attn.k_layernorm is not None + assert self._get_k_norm(attn) is not None def test_qk_l2_norm_from_config(self): """config.qk_l2_norm=True creates L2Norm q/k layernorm.""" @@ -385,57 +465,649 @@ def test_qk_l2_norm_from_config(self): attn = self._get_attention_layer(model) assert attn is not None assert isinstance(attn.q_layernorm, L2Norm) - assert isinstance(attn.k_layernorm, L2Norm) + assert isinstance(self._get_k_norm(attn), L2Norm) def test_spec_provided_norm_not_overwritten(self): """When the spec already provides q/k layernorm, config doesn't override it.""" import copy - from megatron.core.extensions.transformer_engine import ( - TEDotProductAttention, - TELayerNormColumnParallelLinear, - TERowParallelLinear, - ) - from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules - from megatron.core.transformer.enums import AttnMaskType from megatron.core.transformer.identity_op import IdentityOp - from megatron.core.transformer.spec_utils import ModuleSpec - from megatron.core.transformer.transformer_layer import ( - TransformerLayer, - TransformerLayerSubmodules, - ) - # Build a spec that explicitly sets q/k layernorm to IdentityOp + # Build a spec that explicitly sets q/k layernorm to IdentityOp on the + # attention layer that this subclass exercises. spec = copy.deepcopy(hybrid_stack_spec) - spec.submodules.attention_layer.submodules.self_attention.submodules.q_layernorm = ( - IdentityOp + attn_submodules = getattr( + spec.submodules, self._attention_layer_attr + ).submodules.self_attention.submodules + attn_submodules.q_layernorm = IdentityOp + setattr(attn_submodules, self._k_norm_attr, IdentityOp) + + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_attention_layer(model) + assert attn is not None + assert isinstance(attn.q_layernorm, IdentityOp) + assert isinstance(self._get_k_norm(attn), IdentityOp) + + def test_forward_with_qk_layernorm(self): + """HybridModel forward pass works with qk_layernorm enabled.""" + model = self._build_model(qk_layernorm=True) + model.cuda() + + sequence_length = 4 + micro_batch_size = 2 + data = list(range(sequence_length)) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + attention_mask = torch.ones( + (micro_batch_size, 1, sequence_length, sequence_length), dtype=bool + ).cuda() + + logits = model.forward( + input_ids=input_ids, position_ids=position_ids, attention_mask=attention_mask + ) + + assert logits.shape[0] == micro_batch_size + assert logits.shape[1] == sequence_length + assert logits.shape[2] == 100 + + +class TestHybridMLAQKLayernorm(TestHybridQKLayernorm): + """Tests QK norm configuration of HybridModel with MLA.""" + + _attention_layer_attr = 'mla_layer' + _k_norm_attr = 'kv_layernorm' + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config = MLATransformerConfig( + num_layers=3, + hidden_size=256, + num_attention_heads=4, + use_cpu_initialization=True, + **config_overrides, ) - spec.submodules.attention_layer.submodules.self_attention.submodules.k_layernorm = ( - IdentityOp + return HybridModel( + config=config, + hybrid_stack_spec=spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern="M+-", ) - config = TransformerConfig( + def test_qk_l2_norm_from_config(self): + with pytest.raises(ValueError, match="qk_l2_norm is not supported"): + super().test_qk_l2_norm_from_config() + + +class TestHybridDSAQKLayernorm(TestHybridQKLayernorm): + """Tests QK norm configuration of HybridModel with DSA.""" + + _attention_layer_attr = 'dsa_layer' + _k_norm_attr = 'kv_layernorm' + + @pytest.fixture(autouse=True) + def _patch_hadamard_if_needed(self): + if not _HAVE_HADAMARD: + with patch( + 'megatron.core.transformer.experimental_attention_variant.dsa.hadamard_transform', + _mock_hadamard_transform, + ): + yield + else: + yield + + def test_spec_provided_norm_not_overwritten(self): + # DSA cannot fuse the QK norm into the up-projection, so a trivial + # `IdentityOp` spec is auto-promoted to `TENorm` when `qk_layernorm=True`. + # Finer-grained spec-respect behavior is covered by TestDSAQKNormResolution. + pytest.skip("DSA auto-promotes IdentityOp to TENorm; covered by TestDSAQKNormResolution.") + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config_kwargs = dict( num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True, - qk_layernorm=True, + add_bias_linear=False, + # AbsorbedMLASelfAttention forwards `x` and `qr` to the DSA core attention; without + # this, the DSA core attention's forward fails on missing positional arguments. + experimental_attention_variant="dsa", + # DSA-specific settings; defaults are None and DSAIndexer requires them. + dsa_indexer_n_heads=8, + dsa_indexer_head_dim=64, + dsa_indexer_topk=32, + # The indexer-loss path runs in training mode and multiplies by this coefficient; + # leaving it at the default `None` raises `TypeError: ... 'Tensor' and 'NoneType'`. + dsa_indexer_loss_coeff=1.0, + # DSA's `rotate_activation` (Hadamard rotation) only supports bf16 input. + bf16=True, + params_dtype=torch.bfloat16, ) - model = HybridModel( + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( config=config, hybrid_stack_spec=spec, vocab_size=100, max_sequence_length=4, - hybrid_layer_pattern="M*-", + hybrid_layer_pattern="MD-", ) - attn = self._get_attention_layer(model) + + def test_qk_l2_norm_from_config(self): + with pytest.raises(ValueError, match="qk_l2_norm is not supported"): + super().test_qk_l2_norm_from_config() + + +class _MLAQKNormTestBase: + """Common machinery for MLA/DSA QK-norm spec tests. + + Subclasses override `experimental_attention_variant` and + `hybrid_layer_pattern` to target the MLA vs. DSA code path. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def setup_method(self, method): + Utils.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def _make_spec(self, **submodule_overrides): + """Return a copy of `hybrid_stack_spec` with MLA/DSA submodule overrides.""" + import copy + + spec = copy.deepcopy(hybrid_stack_spec) + mla_submodules = getattr( + spec.submodules, self.mla_layer_attr + ).submodules.self_attention.submodules + for key, value in submodule_overrides.items(): + setattr(mla_submodules, key, value) + return spec + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config_kwargs = dict( + num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + if self.experimental_attention_variant is not None: + config_kwargs["experimental_attention_variant"] = self.experimental_attention_variant + if self.experimental_attention_variant == "dsa": + # Must not be True for DSA. + config_kwargs.setdefault("add_bias_linear", False) + # DSAIndexer requires these; their config defaults are None. + config_kwargs.setdefault("dsa_indexer_n_heads", 8) + config_kwargs.setdefault("dsa_indexer_head_dim", 64) + config_kwargs.setdefault("dsa_indexer_topk", 32) + + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( + config=config, + hybrid_stack_spec=spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern=self.hybrid_layer_pattern, + ) + + def _get_mla_attention(self, model): + """Return the attention submodule for the selected MLA variant, or None.""" + if self.experimental_attention_variant == "dsa": + from megatron.core.transformer.experimental_attention_variant.absorbed_mla import ( + AbsorbedMLASelfAttention, + ) + + attention_cls = AbsorbedMLASelfAttention + else: + from megatron.core.transformer.multi_latent_attention import MLASelfAttention + + attention_cls = MLASelfAttention + + for layer in model.decoder.layers: + if hasattr(layer, 'self_attention') and isinstance(layer.self_attention, attention_cls): + return layer.self_attention + return None + + +class TestMLAQKNormSpecValidation(_MLAQKNormTestBase): + """Tests QK norm spec validation in `MLASelfAttention`. + + These errors guard against silently ignoring a configured norm or + double-applying one through a fused norm+linear. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def test_q_norm_without_q_lora_rank_raises(self): + """When `q_lora_rank is None`, a non-trivial `q_layernorm` would + never be reached and must error out. + """ + from megatron.core.extensions.transformer_engine import TENorm + + spec = self._make_spec(q_layernorm=TENorm) + with pytest.raises(ValueError, match=r"q_lora_rank is None"): + self._build_model(spec=spec, q_lora_rank=None) + + def test_q_norm_without_q_lora_rank_hint_for_non_fused_linear(self): + """Error message hints at fused linear when `linear_q_proj` is non-fused.""" + from megatron.core.extensions.transformer_engine import TENorm + + spec = self._make_spec(q_layernorm=TENorm) + with pytest.raises(ValueError, match=r"fused norm\+linear for"): + self._build_model(spec=spec, q_lora_rank=None) + + def test_fused_linear_q_up_with_q_norm_raises(self): + """Non-trivial `q_layernorm` combined with a fused `linear_q_up_proj` + would apply the norm twice. + """ + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(q_layernorm=TENorm, linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"fused norm\+linear"): + self._build_model(spec=spec) + + def test_fused_linear_kv_up_with_kv_norm_raises(self): + """Non-trivial `kv_layernorm` combined with a fused `linear_kv_up_proj` + would apply the norm twice. + """ + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec( + kv_layernorm=TENorm, linear_kv_up_proj=TELayerNormColumnParallelLinear + ) + with pytest.raises(ValueError, match=r"fused norm\+linear"): + self._build_model(spec=spec) + + +class TestMLAQKNormResolution(_MLAQKNormTestBase): + """Tests `_resolve_qk_norm_config` for MLA. + + Covers fusion auto-selection, spec overrides, and the "disabled"-path + guards that reject fused/explicit norms when `qk_layernorm` is off. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def test_qk_layernorm_fuses_kv_up_by_default(self): + """With default (trivial) `kv_layernorm`, enabling `qk_layernorm` + auto-selects the fused `TELayerNormColumnParallelLinear` for KV up. + """ + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + from megatron.core.transformer.identity_op import IdentityOp + + model = self._build_model(qk_layernorm=True) + attn = self._get_mla_attention(model) assert attn is not None - assert isinstance(attn.q_layernorm, IdentityOp) - assert isinstance(attn.k_layernorm, IdentityOp) + assert isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + assert isinstance(attn.kv_layernorm, IdentityOp) + + def test_spec_q_norm_disables_q_up_fusion(self): + """A non-trivial `q_layernorm` from the spec must force a non-fused + `linear_q_up_proj` so the norm isn't applied on top of a fused one. + """ + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(q_layernorm=TENorm) + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_q_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_q_up_proj, TELayerNormColumnParallelLinear) + # The spec's norm is actually used; it's not reset to IdentityOp. + assert attn.q_layernorm is not None + from megatron.core.transformer.identity_op import IdentityOp + + assert not isinstance(attn.q_layernorm, IdentityOp) + + def test_spec_kv_norm_disables_kv_up_fusion(self): + """Mirror of `test_spec_q_norm_disables_q_up_fusion` for KV.""" + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(kv_layernorm=TENorm) + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_kv_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + from megatron.core.transformer.identity_op import IdentityOp + + assert not isinstance(attn.kv_layernorm, IdentityOp) + + def test_disabled_qk_layernorm_rejects_fused_linear_q_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_q_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_fused_linear_kv_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_kv_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_kv_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_spec_norms(self): + """When `qk_layernorm` is off, spec must not carry explicit q/kv layernorms.""" + from megatron.core.extensions.transformer_engine import TENorm + + for overrides in ( + {"q_layernorm": TENorm}, + {"kv_layernorm": TENorm}, + {"q_layernorm": TENorm, "kv_layernorm": TENorm}, + ): + spec = self._make_spec(**overrides) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + +class TestDSAQKNormResolution(_MLAQKNormTestBase): + """Tests `_resolve_qk_norm_config` for DSA. + + DSA requires non-fused Q/KV up projections and explicit norms; + the fused optimization valid for MLA must be rejected here. + """ + + experimental_attention_variant = "dsa" + hybrid_layer_pattern = "MD-" + mla_layer_attr = "dsa_layer" + + def test_qk_layernorm_uses_unfused_linear_and_te_norm(self): + """With default spec, DSA + `qk_layernorm=True` uses non-fused + `TEColumnParallelLinear` and `TENorm` for Q/KV. + """ + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + ) + from megatron.core.transformer.identity_op import IdentityOp - def test_forward_with_qk_layernorm(self): - """HybridModel forward pass works with qk_layernorm enabled.""" model = self._build_model(qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_q_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_q_up_proj, TELayerNormColumnParallelLinear) + assert isinstance(attn.linear_kv_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + assert not isinstance(attn.q_layernorm, IdentityOp) + assert not isinstance(attn.kv_layernorm, IdentityOp) + + def test_qk_layernorm_without_q_lora_rank_raises(self): + """DSA cannot apply Q norm when `q_lora_rank is None`.""" + with pytest.raises(ValueError, match=r"q_lora_rank is None.*not supported for DSA"): + self._build_model(qk_layernorm=True, q_lora_rank=None) + + def test_qk_layernorm_rejects_fused_linear_q_up(self): + """DSA does not support the fused norm+linear optimization.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"not supported for DSA"): + self._build_model(spec=spec, qk_layernorm=True) + + def test_qk_layernorm_without_q_lora_rejects_fused_linear_q(self): + """DSA does not support fused `linear_q_proj` when `q_lora_rank=None`.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"not supported for DSA"): + self._build_model(spec=spec, qk_layernorm=True, q_lora_rank=None) + + def test_disabled_qk_layernorm_rejects_fused_linear_kv_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_kv_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_kv_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_spec_norms(self): + """When `qk_layernorm` is off, spec must not carry explicit q/kv layernorms.""" + from megatron.core.extensions.transformer_engine import TENorm + + for overrides in ( + {"q_layernorm": TENorm}, + {"kv_layernorm": TENorm}, + {"q_layernorm": TENorm, "kv_layernorm": TENorm}, + ): + spec = self._make_spec(**overrides) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + +class TestMLADownProjFusion: + """Tests `HybridStack._fuse_mla_down_proj`. + + The method rewrites the MLA `ModuleSpec` in place on a deep-copied + `HybridStackSubmodules` when `config.mla_down_proj_fusion=True`, swapping + the self-attention module to `FusedMLASelfAttention` and collapsing the + separate q/kv down projections into a single fused `linear_qkv_down_proj` + that also absorbs the input layernorm. + """ + + def setup_method(self, method): + Utils.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def _fresh_submodules(self): + """Return a deep copy of `hybrid_stack_spec.submodules` so tests don't + share state through `hybrid_stack_spec`. + """ + import copy + + return copy.deepcopy(hybrid_stack_spec.submodules) + + def _call_fuse(self, submodules, *, mla_down_proj_fusion): + """Invoke `_fuse_mla_down_proj` as an unbound method with a minimal + stub for `self`. The method only reads `self.config`, so we can avoid + constructing a full `HybridStack`. + """ + from megatron.core.models.hybrid.hybrid_block import HybridStack + + stub = SimpleNamespace(config=SimpleNamespace(mla_down_proj_fusion=mla_down_proj_fusion)) + # Mimic the call-site check in `HybridStack.__init__`. + if getattr(stub.config, "mla_down_proj_fusion", False): + submodules = HybridStack._fuse_mla_down_proj(stub, submodules) + return submodules + + def _build_model(self, pattern="M+-", **config_overrides): + config_kwargs = dict( + num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( + config=config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern=pattern, + ) + + def _get_layer_with_mla(self, model): + """Return the layer whose self-attention is an `MLASelfAttention` + (which includes its `FusedMLASelfAttention` subclass). + """ + from megatron.core.transformer.multi_latent_attention import MLASelfAttention + + for layer in model.decoder.layers: + if hasattr(layer, 'self_attention') and isinstance( + layer.self_attention, MLASelfAttention + ): + return layer + return None + + def test_disabled_returns_spec_unchanged(self): + """Flag off: method returns the same object, no copying or rewriting.""" + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=False) + assert result is submodules + + def test_enabled_rewrites_mla_spec(self): + """Flag on: MLA spec is swapped to the fused module and fused linear.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + from megatron.core.transformer.identity_op import IdentityOp + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + mla_spec = result.mla_layer + assert mla_spec.submodules.input_layernorm is IdentityOp + assert mla_spec.submodules.self_attention.module is FusedMLASelfAttention + + attn_submodules = mla_spec.submodules.self_attention.submodules + assert attn_submodules.linear_qkv_down_proj is TELayerNormColumnParallelLinear + assert attn_submodules.linear_q_down_proj is None + assert attn_submodules.linear_kv_down_proj is None + + def test_enabled_sets_sharded_state_dict_keys_map(self): + """The keys map is written on the MLA layer submodules for checkpoint + compatibility with pre-fusion checkpoints. + """ + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + keys_map = result.mla_layer.submodules.sharded_state_dict_keys_map + assert keys_map == { + "self_attention.linear_q_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_kv_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_qkv_down_proj.layer_norm_": "input_layernorm.", + } + + def test_enabled_deep_copies_input_submodules(self): + """The caller's submodules object must not be mutated – the method + deep-copies before rewriting, so callers can safely reuse their spec. + """ + from megatron.core.transformer.multi_latent_attention import ( + FusedMLASelfAttention, + MLASelfAttention, + ) + + submodules = self._fresh_submodules() + original_mla_module = submodules.mla_layer.submodules.self_attention.module + original_q_down_proj = ( + submodules.mla_layer.submodules.self_attention.submodules.linear_q_down_proj + ) + assert original_mla_module is MLASelfAttention # sanity check of baseline + + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + # Original is unchanged. + assert submodules.mla_layer.submodules.self_attention.module is original_mla_module + assert ( + submodules.mla_layer.submodules.self_attention.submodules.linear_q_down_proj + is original_q_down_proj + ) + # And result is a different object than the input. + assert result is not submodules + assert result.mla_layer is not submodules.mla_layer + # Plus the fused module only shows up on the returned copy. + assert result.mla_layer.submodules.self_attention.module is FusedMLASelfAttention + + def test_enabled_leaves_dsa_layer_alone(self): + """MLA fusion must not rewrite the absorbed DSA attention specification.""" + from megatron.core.transformer.experimental_attention_variant.absorbed_mla import ( + AbsorbedMLASelfAttention, + ) + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + assert result.dsa_layer.submodules.self_attention.module is AbsorbedMLASelfAttention + assert result.dsa_layer.submodules.self_attention.module is not FusedMLASelfAttention + # DSA's down projections must remain non-`None` (they're still used + # via the unfused path). + assert result.dsa_layer.submodules.self_attention.submodules.linear_q_down_proj is not None + assert result.dsa_layer.submodules.self_attention.submodules.linear_kv_down_proj is not None + + def test_enabled_leaves_non_mla_layers_alone(self): + """Unrelated layer specs (mamba, attention, mlp) must survive unchanged.""" + submodules = self._fresh_submodules() + original_mamba = submodules.mamba_layer + original_attention = submodules.attention_layer + original_mlp = submodules.mlp_layer + + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + _assert_equal_with_partial_contents(result.mamba_layer, original_mamba) + _assert_equal_with_partial_contents(result.attention_layer, original_attention) + _assert_equal_with_partial_contents(result.mlp_layer, original_mlp) + + def test_model_uses_fused_mla_when_enabled(self): + """Integration: a full HybridModel built with the flag uses + `FusedMLASelfAttention`. + """ + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + model = self._build_model(mla_down_proj_fusion=True) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.self_attention, FusedMLASelfAttention) + # And the fused down projection is present on the attention module. + assert hasattr(layer.self_attention, "linear_qkv_down_proj") + + def test_model_uses_unfused_mla_when_disabled(self): + """Integration: with the flag off, MLA layers use the standard + `MLASelfAttention` (never the fused subclass). + """ + from megatron.core.transformer.multi_latent_attention import ( + FusedMLASelfAttention, + MLASelfAttention, + ) + + model = self._build_model(mla_down_proj_fusion=False) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.self_attention, MLASelfAttention) + assert not isinstance(layer.self_attention, FusedMLASelfAttention) + + def test_enabled_replaces_input_layernorm_with_identity(self): + """Integration: because the fused down-proj absorbs the input + layernorm, the transformer layer's own `input_layernorm` must be + `IdentityOp`. + """ + from megatron.core.transformer.identity_op import IdentityOp + + model = self._build_model(mla_down_proj_fusion=True) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.input_layernorm, IdentityOp) + + def test_forward_with_fused_mla(self): + """Integration: forward pass works with `mla_down_proj_fusion=True`.""" + model = self._build_model(mla_down_proj_fusion=True) model.cuda() sequence_length = 4 diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index f59a424d5c5..5d3c33264f4 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -3,6 +3,7 @@ import pytest import torch +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 from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec @@ -17,6 +18,7 @@ ) from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention from megatron.core.transformer.mlp import MLP +from megatron.core.transformer.multi_latent_attention import MLASelfAttention from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayer from tests.unit_tests.test_utilities import Utils @@ -52,7 +54,7 @@ def get_hybrid_block(self, layer_pattern, **config_kwargs): pg_collection=self.get_pg_collection(), ) - def get_dsa_mamba_block(self, layer_pattern): + 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 @@ -85,6 +87,35 @@ def get_dsa_mamba_block(self, layer_pattern): 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_attention_heads=16, + use_cpu_initialization=True, + bf16=True, + params_dtype=torch.bfloat16, + q_lora_rank=64, + kv_lora_rank=64, + qk_head_dim=64, + qk_pos_emb_head_dim=32, + v_head_dim=64, + rope_type='rope', + rotary_base=10000, + rotary_percent=1.0, + ) + modules = hybrid_stack_spec.submodules + return HybridStack( + transformer_config, + modules, + layer_type_list=layer_type_list, + pp_layer_offset=0, + pg_collection=self.get_pg_collection(), + ) + def teardown_method(self, method): Utils.destroy_model_parallel() @@ -213,7 +244,7 @@ def test_layer_types(self): assert isinstance(layers[2].mlp, MLP) def test_invalid_layer_types_cause_failure(self): - invalid_symbol = '+' + invalid_symbol = 'X' assert invalid_symbol not in Symbols.VALID_LAYERS # sanity check. layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP + invalid_symbol # validate_segment_layers() in hybrid_layer_allocation.py throws a ValueError. @@ -271,7 +302,7 @@ def test_gdn_gpu_forward(self): def test_dsa_layer_types(self): """D symbol creates a TransformerLayer with absorbed MLA and DSA core attention.""" layer_pattern = Symbols.MAMBA + Symbols.DS_ATTENTION + Symbols.MAMBA - block = self.get_dsa_mamba_block(layer_pattern) + block = self.get_dsa_hybrid_block(layer_pattern) layers = block.layers assert isinstance(layers[0], MambaLayer) assert isinstance(layers[1], TransformerLayer) @@ -283,4 +314,22 @@ def test_mixed_attention_and_dsa_layer_types(self): """* and D in the same block fail.""" layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.DS_ATTENTION + Symbols.MAMBA with pytest.raises(ValueError): - block = self.get_dsa_mamba_block(layer_pattern) + block = self.get_dsa_hybrid_block(layer_pattern) + + def test_mla_layer_types(self): + """+ symbol creates a TransformerLayer with MLASelfAttention but + standard (non-DSA) core attention.""" + layer_pattern = Symbols.MAMBA + Symbols.MLA + Symbols.MAMBA + block = self.get_mla_hybrid_block(layer_pattern) + layers = block.layers + assert isinstance(layers[0], MambaLayer) + assert isinstance(layers[1], TransformerLayer) + assert isinstance(layers[1].self_attention, MLASelfAttention) + assert isinstance(layers[1].self_attention.core_attention, TEDotProductAttention) + assert isinstance(layers[2], MambaLayer) + + def test_mixed_attention_and_mla_layer_types(self): + """* and + in the same block fail (same reason as * and D).""" + layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLA + Symbols.MAMBA + with pytest.raises(ValueError): + block = self.get_mla_hybrid_block(layer_pattern) diff --git a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py index faa553216da..8b4c181ee30 100644 --- a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py +++ b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py @@ -78,6 +78,7 @@ def test_valid_patterns(self): ("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) @@ -101,6 +102,11 @@ def test_invalid_symbols_cause_failure(self): with pytest.raises(ValueError): # Not allowed to have both standard Attention and MLA/DSA validate_segment_layers("MDM*-") + 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*-") @pytest.mark.internal @@ -163,6 +169,8 @@ def test_main_pattern_only(self): ("GEGEGE*E", "GEGEGE*E"), ("MDMD", "MDMD"), ("DM", "DM"), + ("M+M+", "M+M+"), + ("+M", "+M"), ] for pattern, expected_main in test_cases: result = parse_hybrid_pattern(pattern) @@ -287,6 +295,8 @@ def test_complex_patterns(self): ("GEGEGE*E/GG/GG", "GEGEGE*E", "GG", 2), # DSA in main pattern with MTP ("MDMD/MD/MD", "MDMD", "MD", 2), + # MLA in main pattern with MTP + ("M+M+/M+/M+", "M+M+", "M+", 2), ] for pattern, expected_main, expected_mtp, expected_depths in test_cases: result = parse_hybrid_pattern(pattern) @@ -305,21 +315,63 @@ def test_dataclass_equality(self): class TestGetHybridLayerCounts: def test_simple_pattern(self): - assert get_hybrid_layer_counts("M*M*") == {'*': 2, 'D': 0, 'G': 0, 'M': 2, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("M*M*") == { + '*': 2, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 0, + '-': 0, + 'E': 0, + } def test_all_layer_types(self): # Not allowed to have both standard Attention and MLA/DSA, so we do separate asserts. - assert get_hybrid_layer_counts("MG*-E") == {'*': 1, 'D': 0, 'G': 1, 'M': 1, '-': 1, 'E': 1} - assert get_hybrid_layer_counts("MGD-E") == {'*': 0, 'D': 1, 'G': 1, 'M': 1, '-': 1, 'E': 1} + assert get_hybrid_layer_counts("MG*-E") == { + '*': 1, + 'D': 0, + 'G': 1, + 'M': 1, + '+': 0, + '-': 1, + 'E': 1, + } + assert get_hybrid_layer_counts("MGD-E") == { + '*': 0, + 'D': 1, + 'G': 1, + 'M': 1, + '+': 0, + '-': 1, + 'E': 1, + } + assert get_hybrid_layer_counts("MG+-E") == { + '*': 0, + 'D': 0, + 'G': 1, + 'M': 1, + '+': 1, + '-': 1, + 'E': 1, + } def test_with_pipes(self): # Pipes should be skipped in counting - assert get_hybrid_layer_counts("M*|M*") == {'*': 2, 'D': 0, 'G': 0, 'M': 2, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("M*|M*") == { + '*': 2, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 0, + '-': 0, + 'E': 0, + } assert get_hybrid_layer_counts("M-M-|M-M*-") == { '*': 1, 'D': 0, 'G': 0, 'M': 4, + '+': 0, '-': 4, 'E': 0, } @@ -331,6 +383,7 @@ def test_with_mtp(self): 'D': 0, 'G': 0, 'M': 6, + '+': 0, '-': 0, 'E': 0, } @@ -343,12 +396,21 @@ def test_with_pipes_and_mtp(self): 'D': 0, 'G': 0, 'M': 8, + '+': 0, '-': 4, 'E': 0, } def test_moe_pattern(self): - assert get_hybrid_layer_counts("MEME") == {'*': 0, 'D': 0, 'G': 0, 'M': 2, '-': 0, 'E': 2} + assert get_hybrid_layer_counts("MEME") == { + '*': 0, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 0, + '-': 0, + 'E': 2, + } def test_mtp_with_attention(self): # MTP pattern "*M" repeated 3 depths -> 3 attn + 3 mamba from MTP @@ -357,22 +419,66 @@ def test_mtp_with_attention(self): 'D': 0, 'G': 0, 'M': 7, + '+': 0, '-': 0, 'E': 0, } def test_gdn_pattern(self): - assert get_hybrid_layer_counts("GMGM") == {'*': 0, 'D': 0, 'G': 2, 'M': 2, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("GMGM") == { + '*': 0, + 'D': 0, + 'G': 2, + 'M': 2, + '+': 0, + '-': 0, + 'E': 0, + } def test_gdn_hybrid_pattern(self): # GDN + Mamba + Attention - assert get_hybrid_layer_counts("G*GM*") == {'*': 2, 'D': 0, 'G': 2, 'M': 1, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("G*GM*") == { + '*': 2, + 'D': 0, + 'G': 2, + 'M': 1, + '+': 0, + '-': 0, + 'E': 0, + } def test_dsa_pattern(self): - assert get_hybrid_layer_counts("DMDM") == {'*': 0, 'D': 2, 'G': 0, 'M': 2, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("DMDM") == { + '*': 0, + 'D': 2, + 'G': 0, + 'M': 2, + '+': 0, + '-': 0, + 'E': 0, + } + + def test_mla_pattern(self): + assert get_hybrid_layer_counts("+M+M") == { + '*': 0, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 2, + '-': 0, + 'E': 0, + } def test_empty_pattern(self): - assert get_hybrid_layer_counts("") == {'*': 0, 'D': 0, 'G': 0, 'M': 0, '-': 0, 'E': 0} + assert get_hybrid_layer_counts("") == { + '*': 0, + 'D': 0, + 'G': 0, + 'M': 0, + '+': 0, + '-': 0, + 'E': 0, + } @pytest.mark.internal @@ -655,7 +761,7 @@ 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) == 6 + assert len(maps) == 7 attention_map, mamba_map, mlp_map, moe_map = operator.itemgetter( Symbols.ATTENTION, Symbols.MAMBA, Symbols.MLP, Symbols.MOE )(maps) @@ -698,3 +804,39 @@ def test_all_mamba(self): assert mamba_map == {0: 0, 1: 1, 2: 2} assert mlp_map == {} 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"]) + attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( + Symbols.ATTENTION, + Symbols.DS_ATTENTION, + Symbols.MAMBA, + Symbols.MLA, + Symbols.MLP, + Symbols.MOE, + )(maps) + assert attention_map == {} + assert dsa_map == {} + assert mla_map == {0: 0, 2: 1} + assert mamba_map == {1: 0, 3: 1} + assert mlp_map == {} + 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", "-"]) + attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( + Symbols.ATTENTION, + Symbols.DS_ATTENTION, + Symbols.MAMBA, + Symbols.MLA, + Symbols.MLP, + Symbols.MOE, + )(maps) + assert attention_map == {} + assert dsa_map == {0: 0} + assert mla_map == {1: 0} + assert mamba_map == {2: 0} + assert mlp_map == {3: 0} + assert moe_map == {} diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py index 1b81fe73399..fc1778f649f 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py @@ -125,7 +125,7 @@ def _forward_thd(self, q, k, v, packed_seq_params): def get_mock_mla_config( - tensor_model_parallel_size: int, context_parallel_size: int + tensor_model_parallel_size: int, context_parallel_size: int, qk_layernorm: bool ) -> MLATransformerConfig: """Create test config with all attributes used in MLA.""" return MLATransformerConfig( @@ -142,6 +142,7 @@ def get_mock_mla_config( params_dtype=torch.bfloat16, layernorm_epsilon=1e-5, normalization="RMSNorm", + qk_layernorm=qk_layernorm, layernorm_zero_centered_gamma=False, expert_model_parallel_size=1, tensor_model_parallel_size=tensor_model_parallel_size, @@ -399,15 +400,18 @@ def test_functionality(tp_cp: List[int], qkv_format: str, down_proj_use_column_p model_parallel_cuda_manual_seed(123) # Create model - config = get_mock_mla_config(tensor_model_parallel_size=tp_size, context_parallel_size=cp_size) + qk_layernorm = True + config = get_mock_mla_config( + tensor_model_parallel_size=tp_size, context_parallel_size=cp_size, qk_layernorm=qk_layernorm + ) absorbed_submodules = get_absorbed_mla_submodules( down_proj_use_column_parallel=down_proj_use_column_parallel, - qk_layernorm=True, + qk_layernorm=qk_layernorm, rms_norm=True, ) standard_submodules = get_mla_submodules( down_proj_use_column_parallel=down_proj_use_column_parallel, - qk_layernorm=True, + qk_layernorm=qk_layernorm, rms_norm=True, ) absorbed_mla = AbsorbedMLASelfAttention( diff --git a/tests/unit_tests/transformer/test_submodule_callables.py b/tests/unit_tests/transformer/test_submodule_callables.py index 42ba73bc92e..3b111db1548 100644 --- a/tests/unit_tests/transformer/test_submodule_callables.py +++ b/tests/unit_tests/transformer/test_submodule_callables.py @@ -199,9 +199,11 @@ def test_1f1b_overlap(self, dispatcher_type, grouped_gemm, permute_fusion): expert_model_parallel_size=2, virtual_pipeline_model_parallel_size=2, ) + qk_layernorm = True extra_kwargs = { "moe_token_dispatcher_type": dispatcher_type, "moe_permute_fusion": permute_fusion, + "qk_layernorm": qk_layernorm, } if dispatcher_type == "flex": extra_kwargs["moe_flex_dispatcher_backend"] = get_valid_flex_dispatcher_backend() @@ -211,7 +213,7 @@ def test_1f1b_overlap(self, dispatcher_type, grouped_gemm, permute_fusion): transformer_layer_submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=8, moe_grouped_gemm=grouped_gemm, - qk_layernorm=True, + qk_layernorm=qk_layernorm, multi_latent_attention=True, ) model = TransformerLayer(config, transformer_layer_submodules)