Skip to content
Draft
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
9 changes: 5 additions & 4 deletions megatron/core/inference/contexts/dynamic_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,6 @@
)
from megatron.core.inference.utils import device_memory_summary, tensor_swap
from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb
from megatron.core.models.hybrid.hybrid_layer_allocation import (
Symbols,
get_layer_maps_from_layer_type_list,
)
Comment on lines -33 to -36

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why was this moved?

from megatron.core.package_info import __version__ as mcore_version
from megatron.core.transformer import MLATransformerConfig, TransformerConfig
from megatron.core.transformer.enums import InferenceCudaGraphScope
Expand Down Expand Up @@ -441,6 +437,11 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC
# For hybrid models, the layer map converts the global layer index to the
# corresponding attention layer index or Mamba layer index depending on the
# layer type.
from megatron.core.models.hybrid.hybrid_layer_allocation import (
Symbols,
get_layer_maps_from_layer_type_list,
)

attention_layer_map, dsa_layer_map, gdn_layer_map, mamba_layer_map = (
operator.itemgetter(
Symbols.ATTENTION, Symbols.DS_ATTENTION, Symbols.GDN, Symbols.MAMBA
Expand Down
108 changes: 76 additions & 32 deletions megatron/core/models/hybrid/hybrid_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,15 +21,22 @@
from megatron.core.fp8_utils import get_fp8_context
from megatron.core.inference.contexts import BaseInferenceContext
from megatron.core.inference.utils import InferenceMode
from megatron.core.models.hybrid.hybrid_layer_allocation import HybridLayerConfig
from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.recompute import checkpointed_forward
from megatron.core.ssm.gated_delta_net import GDNLayerConfig
from megatron.core.ssm.mamba_layer import MambaLayerConfig
from megatron.core.ssm.mlp_layer import MLPLayerConfig
from megatron.core.transformer import TransformerConfig
from megatron.core.transformer.attention import AttentionLayerConfig
from megatron.core.transformer.cuda_graphs import annotate_first_last_layer
from megatron.core.transformer.experimental_attention_variant.dsa import DSALayerConfig
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.moe.moe_layer import MoELayerConfig
from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention, MLALayerConfig
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 Down Expand Up @@ -59,11 +66,11 @@ class HybridStack(MegatronModule):
Args:
config (TransformerConfig): the model configuration
submodules (HybridStackSubmodules): the submodules for the stack
layer_config_list (list): per-layer configs for this pipeline segment. When
provided by HybridModel, pipeline stage selection has already been done
via '|' separators in the pattern.
pre_process (bool, optional): whether to include an embedding layer.
Defaults to True.
layer_type_list (list, optional): pre-computed list of layer type symbols for
this pipeline segment. When provided (by HybridModel), pipeline stage
selection has already been done via '|' separators in the pattern.
pp_layer_offset (int, optional): the global layer offset for this pipeline
segment. Defaults to 0.
post_layer_norm (bool, optional): whether to include a final layer norm.
Expand All @@ -81,8 +88,8 @@ def __init__(
self,
config: TransformerConfig,
submodules: HybridStackSubmodules,
layer_config_list: list[HybridLayerConfig],
pre_process: bool = True,
layer_type_list: Optional[list[str]] = None,
pp_layer_offset: int = 0,
post_layer_norm: bool = True,
post_process: bool = True,
Expand Down Expand Up @@ -111,97 +118,116 @@ def __init__(
self.input_tensor = None
self.pg_collection = pg_collection

assert layer_type_list is not None, (
"layer_type_list must be provided. It should be pre-computed from "
"--hybrid-layer-pattern by HybridModel."
)
self.layer_type_list = layer_type_list
self.layer_config_list = layer_config_list
# Retain the symbol list for inference metadata and checkpoint mapping consumers.
self.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):
for i, layer_config in enumerate(self.layer_config_list):
layer_number = i + 1 + pp_layer_offset
if self.config.fp8:
quant_init_context = get_fp8_context(self.config, i + pp_layer_offset, is_init=True)
elif self.config.fp4:
quant_init_context = get_fp4_context(self.config, i + pp_layer_offset, is_init=True)
tp_comm_overlap = layer_config.tp_comm_overlap
if layer_config.fp8:
quant_init_context = get_fp8_context(
layer_config, i + pp_layer_offset, is_init=True
)
elif layer_config.fp4:
quant_init_context = get_fp4_context(
layer_config, i + pp_layer_offset, is_init=True
)
else:
quant_init_context = nullcontext()
with quant_init_context:
if layer_type == LayerSymbols.MAMBA:
if type(layer_config) is MambaLayerConfig:
layer_type = LayerSymbols.MAMBA
layer = build_module(
submodules.mamba_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pp_layer_offset=pp_layer_offset,
pg_collection=pg_collection,
name=(name + f".layers.{i}") if name is not None else None,
)
elif layer_type == LayerSymbols.ATTENTION:
elif type(layer_config) is AttentionLayerConfig:
layer_type = LayerSymbols.ATTENTION
layer = build_module(
submodules.attention_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
is_mtp_layer=is_mtp_layer,
add_layer_offset=False,
pp_layer_offset=pp_layer_offset,
name=(name + f".layers.{i}") if name is not None else None,
)
elif layer_type == LayerSymbols.DS_ATTENTION:
elif type(layer_config) is DSALayerConfig:
layer_type = LayerSymbols.DS_ATTENTION
layer = build_module(
submodules.dsa_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
is_mtp_layer=is_mtp_layer,
add_layer_offset=False,
pp_layer_offset=pp_layer_offset,
name=(name + f".layers.{i}") if name is not None else None,
)
elif layer_type == LayerSymbols.MLA:
elif type(layer_config) is MLALayerConfig:
layer_type = LayerSymbols.MLA
layer = build_module(
submodules.mla_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
is_mtp_layer=is_mtp_layer,
add_layer_offset=False,
pp_layer_offset=pp_layer_offset,
)
elif layer_type == LayerSymbols.MLP:
elif type(layer_config) is MLPLayerConfig:
layer_type = LayerSymbols.MLP
layer = build_module(
submodules.mlp_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
add_layer_offset=False,
name=(name + f".layers.{i}") if name is not None else None,
)
elif layer_type == LayerSymbols.MOE:
elif type(layer_config) is MoELayerConfig:
layer_type = LayerSymbols.MOE
layer = build_module(
submodules.moe_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
add_layer_offset=False,
name=(name + f".layers.{i}") if name is not None else None,
)
elif layer_type == LayerSymbols.GDN:
elif type(layer_config) is GDNLayerConfig:
layer_type = LayerSymbols.GDN
layer = build_module(
submodules.gdn_layer,
config=self.config,
config=layer_config,
layer_number=layer_number,
pg_collection=pg_collection,
# Set to False as we do not want to change offset.
add_layer_offset=False,
name=(name + f".layers.{i}") if name is not None else None,
)
else:
raise ValueError("unexpected layer_type")
raise ValueError(
f"Unexpected hybrid layer config type: {type(layer_config).__name__}"
)

# Some layer builders disable unsupported TP overlap by mutating their config.
# Preserve the shared-config behavior for configs that had the same initial value.
self.synchronize_shared_config_mutation(
"tp_comm_overlap", tp_comm_overlap, layer_config.tp_comm_overlap
)
self.layer_type_list.append(layer_type)
self.layers.append(layer)

if self.config.cuda_graph_impl == "local":
Expand All @@ -218,6 +244,24 @@ def __init__(
eps=self.config.layernorm_epsilon,
)

def synchronize_shared_config_mutation(
self, attribute: str, old_value: object, new_value: object
) -> None:
"""Propagate a legacy shared-config mutation to configs cloned from it.

Plain lists are treated as independently supplied configs and are not synchronized.
Within a legacy-derived list, configs that have already diverged from the old shared
value are also left unchanged.
"""
if old_value == new_value or not getattr(
self.layer_config_list, "synchronize_shared_config_mutations", False
):
return

for config_to_update in [self.config, *self.layer_config_list]:
if getattr(config_to_update, attribute) == old_value:
setattr(config_to_update, attribute, new_value)

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.
Expand Down Expand Up @@ -359,10 +403,10 @@ def get_inner_quant_context(config, layer_number):
use_inner_quantization_context=(use_inner_fp8_context or use_fp4_context),
)
else:
for layer in self.layers:
for layer_config, layer in zip(self.layer_config_list, self.layers):
# Layers have 1-indexed layer numbers attribute.
inner_quant_context = get_inner_quant_context(
self.config, layer.layer_number - 1
layer_config, layer.layer_number - 1
)
with inner_quant_context:
if isinstance(layer, TransformerLayer):
Expand Down
Loading