Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
39 commits
Select commit Hold shift + click to select a range
9386fd8
Port MLA to HybridModel
janEbert Apr 22, 2026
9568486
Add tests for HybridModel MLA
janEbert Apr 22, 2026
59fbe0b
Rename "Mamba block" to "Hybrid block"
janEbert Apr 23, 2026
5242ca4
Raise an error for twice-configured Q/KV norms
janEbert Apr 24, 2026
1749d95
Support setting QK norm from config for MLA
janEbert Apr 24, 2026
984368d
Do not apply norm+linear fusion automatically
janEbert Apr 24, 2026
6fb703f
Make existing tests more flexible
janEbert Apr 24, 2026
615a6df
Add initial Hybrid MLA/DSA QK config tests
janEbert Apr 24, 2026
af3fd7b
Add additional tests
janEbert Apr 24, 2026
ccaefce
Support `mla_down_proj_fusion=True`
janEbert Apr 24, 2026
e0f58fd
Add tests for Hybrid + `mla_down_proj_fusion`
janEbert Apr 24, 2026
070dae9
Fix type error
janEbert Apr 24, 2026
7d3eda1
Simplify norm resolution
janEbert Apr 24, 2026
e9d776a
Fix norm config resolution
janEbert Apr 28, 2026
4c6aaa5
Fix tests
janEbert Apr 28, 2026
fee07c3
Move conditional imports to unconditionally global
janEbert Apr 28, 2026
e17fcfc
Add backend utility
janEbert Apr 28, 2026
4b25b98
Support all backends
janEbert Apr 28, 2026
f9e2c30
Fix return type
janEbert Apr 28, 2026
6777113
Fix layer type checks
janEbert May 6, 2026
a63c7a8
Check for fusion layers explicitly
janEbert May 6, 2026
f47a2d3
Fix quotes
janEbert May 6, 2026
3f9c155
Be more lenient with specs
janEbert May 6, 2026
662d7ad
Fix automatic MLA inference
janEbert May 6, 2026
50d8018
Fix config not matching specs
janEbert May 7, 2026
d3b5240
Do not repeat oneself
janEbert May 7, 2026
8858492
Improve norm spec test generality
janEbert May 7, 2026
084df1f
Align error types
janEbert May 7, 2026
aea1087
Fix test regexps
janEbert May 7, 2026
4dc68df
Copy less data
janEbert Jun 1, 2026
6586e77
Revert "Copy less data"
janEbert Jun 9, 2026
80d2a26
Refactor condition check
janEbert Jun 9, 2026
a45492a
Fix equality checks
janEbert Jun 9, 2026
e414280
Simplify QK norm validation and resolution
janEbert Jun 25, 2026
9d8d778
Refactor QK norm validation/resolution to new file
janEbert Jun 29, 2026
f5e9978
Fix docstring and variable name
janEbert Jun 29, 2026
bc9b93f
Fix up AbsorbedMLA
janEbert Jun 29, 2026
e4aa27d
Fix tests
janEbert Jun 29, 2026
19af620
Document deepcopy
janEbert Jul 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion megatron/core/models/backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -200,3 +200,19 @@ def grouped_mlp_modules(self, moe_use_grouped_gemm: bool) -> ExpertsBuilder:
activation_func=self.activation_func(),
),
)


def get_backend(
Comment thread
santhnm2 marked this conversation as resolved.
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}'")
38 changes: 37 additions & 1 deletion megatron/core/models/hybrid/hybrid_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Comment thread
janEbert marked this conversation as resolved.
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.

Expand Down
7 changes: 4 additions & 3 deletions megatron/core/models/hybrid/hybrid_layer_allocation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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")


Expand Down Expand Up @@ -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
Expand Down
44 changes: 44 additions & 0 deletions megatron/core/models/hybrid/hybrid_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
Loading
Loading