Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
94 changes: 93 additions & 1 deletion megatron/core/post_training/modelopt/hybrid/model_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,14 +6,32 @@
from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec
from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules
from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec
from megatron.core.post_training.modelopt.layers import Norm
from megatron.core.post_training.modelopt.layers import Linear, Norm
from megatron.core.ssm.gated_delta_net import GatedDeltaNet, GatedDeltaNetSubmodules
from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules
from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules
from megatron.core.transformer.dot_product_attention import DotProductAttention
from megatron.core.transformer.enums import AttnMaskType
from megatron.core.transformer.experimental_attention_variant.dsa import (
DSAIndexer,
DSAIndexerSubmodules,
DSAttention,
DSAttentionSubmodules,
)
from megatron.core.transformer.identity_op import IdentityOp
from megatron.core.transformer.mlp import MLP, MLPSubmodules
from megatron.core.transformer.multi_latent_attention import (
MLASelfAttention,
MLASelfAttentionSubmodules,
)
from megatron.core.transformer.multi_token_prediction import (
MultiTokenPredictionBlock,
MultiTokenPredictionBlockSubmodules,
MultiTokenPredictionLayer,
MultiTokenPredictionLayerSubmodules,
)
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.transformer_layer import (
MoETransformerLayer,
Expand Down Expand Up @@ -98,12 +116,14 @@ def _get_hybrid_stack_local_spec(
"""
mamba_state_dict_keys_map = {}
transformer_state_dict_keys_map = {}
gdn_state_dict_keys_map = {}
if remap_te_layernorm:
mamba_state_dict_keys_map = {'norm.': 'mixer.in_proj.layer_norm_'}
transformer_state_dict_keys_map = {
'input_layernorm.': 'self_attention.linear_qkv.layer_norm_',
'pre_mlp_layernorm.': 'mlp.linear_fc1.layer_norm_',
}
gdn_state_dict_keys_map = {'input_layernorm.': 'self_attention.in_proj.layer_norm_'}

mamba_layer = ModuleSpec(
module=MambaLayer,
Expand All @@ -120,6 +140,21 @@ def _get_hybrid_stack_local_spec(
),
)

gdn_layer = ModuleSpec(
module=TransformerLayer,
submodules=TransformerLayerSubmodules(
input_layernorm=Norm,
self_attention=ModuleSpec(
module=GatedDeltaNet,
submodules=GatedDeltaNetSubmodules(
in_proj=ColumnParallelLinear, out_norm=Norm, out_proj=RowParallelLinear
),
),
self_attn_bda=get_bias_dropout_add,
sharded_state_dict_keys_map=gdn_state_dict_keys_map,
),
)

attn_mask_type = AttnMaskType.causal
core_attention = DotProductAttention if local_core_attention else TEDotProductAttention
attention_layer = ModuleSpec(
Expand All @@ -140,6 +175,42 @@ def _get_hybrid_stack_local_spec(
),
)

dsa_layer = ModuleSpec(
module=TransformerLayer,
submodules=TransformerLayerSubmodules(
input_layernorm=Norm,
self_attention=ModuleSpec(
module=MLASelfAttention,
params={"attn_mask_type": attn_mask_type},
submodules=MLASelfAttentionSubmodules(
linear_q_proj=ColumnParallelLinear,
linear_q_down_proj=Linear,
linear_q_up_proj=ColumnParallelLinear,
linear_kv_down_proj=Linear,
linear_kv_up_proj=ColumnParallelLinear,
core_attention=ModuleSpec(
module=DSAttention,
submodules=DSAttentionSubmodules(
indexer=ModuleSpec(
module=DSAIndexer,
submodules=DSAIndexerSubmodules(
linear_wq_b=Linear,
linear_wk=Linear,
k_norm=Norm,
linear_weights_proj=Linear,
),
)
),
),
linear_proj=RowParallelLinear,
q_layernorm=IdentityOp,
kv_layernorm=IdentityOp,
),
),
self_attn_bda=get_bias_dropout_add,
),
)

mlp_layer = ModuleSpec(
module=TransformerLayer,
submodules=TransformerLayerSubmodules(
Expand All @@ -166,12 +237,33 @@ def _get_hybrid_stack_local_spec(
),
)

mtp_block_spec = ModuleSpec(
module=MultiTokenPredictionBlock,
submodules=MultiTokenPredictionBlockSubmodules(
layer_specs=[
ModuleSpec(
module=MultiTokenPredictionLayer,
submodules=MultiTokenPredictionLayerSubmodules(
enorm=Norm,
hnorm=Norm,
eh_proj=ColumnParallelLinear,
mtp_model_layer=None,
layer_norm=Norm,
),
)
]
),
)

return ModuleSpec(
module=HybridStack,
submodules=HybridStackSubmodules(
mamba_layer=mamba_layer,
gdn_layer=gdn_layer,
attention_layer=attention_layer,
dsa_layer=dsa_layer,
mlp_layer=mlp_layer,
moe_layer=moe_layer,
mtp_block_spec=mtp_block_spec,
),
)
13 changes: 12 additions & 1 deletion megatron/core/post_training/modelopt/layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,11 +123,20 @@ def __init__(
is_expert: bool = False,
tp_comm_buffer_name: str = None, # Not used
disable_grad_reduce: bool = False,
parallel_mode: Optional[str] = None,
tp_group: Optional[torch.distributed.ProcessGroup] = None,
name: str | None = None, # Not used
):
if parallel_mode not in (None, "duplicated"):
raise ValueError(
f"{type(self).__name__} only supports parallel_mode='duplicated' or None"
)
if parallel_mode == "duplicated" and tp_group is not None:
raise ValueError("duplicated Linear should not have tp_group set")

self.config = config
self.tp_group = tp_group
self.parallel_mode = parallel_mode
self.tp_group = None if parallel_mode == "duplicated" else tp_group

self._return_bias = skip_bias_add and bias

Expand Down Expand Up @@ -155,6 +164,8 @@ def __init__(
# Reduce the gradient on DP group
setattr(param, "allreduce", True)
setattr(param, "sequence_parallel", self.config.sequence_parallel)
if parallel_mode == "duplicated":
setattr(param, "tensor_model_parallel", False)

def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None):
"""Sharding along axis 0, bias sharded"""
Expand Down
75 changes: 75 additions & 0 deletions tests/unit_tests/post_training/test_modelopt_module_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,20 @@
mcore_gpt_load_te_state_dict_pre_hook,
)
from megatron.core.post_training.modelopt.hybrid.model_specs import get_hybrid_stack_modelopt_spec
from megatron.core.post_training.modelopt.layers import Linear, Norm
from megatron.core.ssm.gated_delta_net import GatedDeltaNet
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed
from megatron.core.transformer import TransformerConfig
from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexer, DSAttention
from megatron.core.transformer.identity_op import IdentityOp
from megatron.core.transformer.multi_latent_attention import MLASelfAttention
from megatron.core.transformer.multi_token_prediction import (
MultiTokenPredictionBlock,
MultiTokenPredictionLayer,
)
from megatron.core.transformer.transformer_config import MLATransformerConfig
from megatron.core.transformer.transformer_layer import TransformerLayer
from megatron.core.utils import get_te_version
from tests.unit_tests.dist_checkpointing import TempNamedDir
from tests.unit_tests.test_utilities import Utils
Expand Down Expand Up @@ -308,3 +319,67 @@ def test_get_hybrid_stack_modelopt_spec_use_default_te_spec():
"""Test that use_default_te_spec=True returns the standard hybrid_stack_spec."""
spec = get_hybrid_stack_modelopt_spec(use_default_te_spec=True)
assert spec is hybrid_stack_spec


def test_get_hybrid_stack_modelopt_spec_local_feature_specs():
"""The local ModelOpt HybridStack spec covers all HybridModel layer families."""
spec = get_hybrid_stack_modelopt_spec()
submodules = spec.submodules

gdn_layer = submodules.gdn_layer
assert gdn_layer.module is TransformerLayer
assert gdn_layer.submodules.input_layernorm is Norm
assert gdn_layer.submodules.self_attention.module is GatedDeltaNet
assert gdn_layer.submodules.self_attention.submodules.in_proj is ColumnParallelLinear
assert gdn_layer.submodules.self_attention.submodules.out_norm is Norm
assert gdn_layer.submodules.self_attention.submodules.out_proj is RowParallelLinear

dsa_layer = submodules.dsa_layer
assert dsa_layer.module is TransformerLayer
assert dsa_layer.submodules.input_layernorm is Norm
assert dsa_layer.submodules.self_attention.module is MLASelfAttention
assert dsa_layer.submodules.self_attention.submodules.q_layernorm is IdentityOp
assert dsa_layer.submodules.self_attention.submodules.kv_layernorm is IdentityOp
dsa_attention = dsa_layer.submodules.self_attention.submodules.core_attention
assert dsa_attention.module is DSAttention
indexer = dsa_attention.submodules.indexer
assert indexer.module is DSAIndexer
assert indexer.submodules.linear_wq_b is Linear
assert "parallel_mode" in inspect.signature(indexer.submodules.linear_wq_b).parameters
assert indexer.submodules.linear_wk is Linear
assert indexer.submodules.k_norm is Norm
assert indexer.submodules.linear_weights_proj is Linear

mtp_block_spec = submodules.mtp_block_spec
assert mtp_block_spec.module is MultiTokenPredictionBlock
mtp_layer_spec = mtp_block_spec.submodules.layer_specs[0]
assert mtp_layer_spec.module is MultiTokenPredictionLayer
assert mtp_layer_spec.submodules.enorm is Norm
assert mtp_layer_spec.submodules.hnorm is Norm
assert mtp_layer_spec.submodules.eh_proj is ColumnParallelLinear
assert mtp_layer_spec.submodules.layer_norm is Norm


def test_get_hybrid_stack_modelopt_spec_remaps_gdn_layernorm():
"""GDN local spec can load checkpoints saved from the fused TE GDN spec."""
spec = get_hybrid_stack_modelopt_spec(remap_te_layernorm=True)
assert spec.submodules.gdn_layer.submodules.sharded_state_dict_keys_map == {
'input_layernorm.': 'self_attention.in_proj.layer_norm_'
}


def test_modelopt_linear_accepts_duplicated_parallel_mode():
"""ModelOpt Linear supports duplicated TELinear-compatible construction."""
config = TransformerConfig(
num_layers=1, hidden_size=4, num_attention_heads=1, use_cpu_initialization=True
)
linear = Linear(
4, 4, config=config, init_method=config.init_method, bias=False, parallel_mode="duplicated"
)

assert linear.parallel_mode == "duplicated"
assert linear.tp_group is None
assert linear.weight.tensor_model_parallel is False

with pytest.raises(ValueError, match="only supports parallel_mode"):
Linear(4, 4, config=config, init_method=config.init_method, parallel_mode="column")
Loading