diff --git a/gpt_builders.py b/gpt_builders.py index 293475b06b6..0be64edaab6 100644 --- a/gpt_builders.py +++ b/gpt_builders.py @@ -10,7 +10,8 @@ get_gpt_decoder_layer_specs, ) from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( - is_linear_attention_variant, + get_transformer_block_with_experimental_attention_variant_spec, + get_transformer_layer_with_experimental_attention_variant_spec, ) from megatron.core.models.gpt.heterogeneous.heterogeneous_layer_specs import ( get_gpt_heterogeneous_layer_spec, @@ -46,7 +47,13 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ else: use_te = args.transformer_impl == "transformer_engine" - if args.num_experts or is_linear_attention_variant(args.experimental_attention_variant): + if args.experimental_attention_variant is not None: + transformer_layer_spec = ( + get_transformer_block_with_experimental_attention_variant_spec( + config=config, vp_stage=vp_stage + ) + ) + elif args.num_experts: assert not (config.transformer_impl == "inference_optimized") # Define the decoder block spec transformer_layer_spec = get_gpt_decoder_block_spec( @@ -70,9 +77,19 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ mtp_transformer_layer_spec = import_module(args.spec) else: # Define the decoder block spec - decoder_layer_specs = get_gpt_decoder_layer_specs( - config, use_transformer_engine=use_te, normalization=args.normalization, qk_l2_norm=args.qk_l2_norm, vp_stage=vp_stage - ) + if args.experimental_attention_variant is not None: + decoder_layer_specs = ( + get_transformer_layer_with_experimental_attention_variant_spec( + config=config + ) + ) + else: + decoder_layer_specs = get_gpt_decoder_layer_specs( + config, + use_transformer_engine=use_te, + normalization=args.normalization, + qk_l2_norm=args.qk_l2_norm, + ) mtp_transformer_layer_spec = decoder_layer_specs[-1] # Use spec of the last layer in decoder block as spec of the transformer layer in MTP mtp_block_spec = get_gpt_mtp_block_spec( diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index e6d6fa03ce7..7649a0b2165 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -1,10 +1,11 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. -from typing import Optional +from typing import List, Optional +from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add from megatron.core.models.backends import BackendSpecProvider from megatron.core.ssm.gated_delta_net import GatedDeltaNet, GatedDeltaNetSubmodules -from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.experimental_attention_variant.dsa import ( DSAIndexer, DSAIndexerSubmodules, @@ -17,19 +18,50 @@ MLASelfAttentionSubmodules, ) from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.transformer_block import ( + TransformerBlockSubmodules, + get_num_layers_to_build, +) +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import ( + TransformerLayer, + TransformerLayerSubmodules, + get_transformer_layer_offset, +) +try: + import transformer_engine as te # type: ignore[import-untyped] # pylint: disable=unused-import -def is_linear_attention_variant(experimental_attention_variant: str) -> bool: - """Check if the experimental attention variant is a linear attention variant.""" - linear_attention_variants = ["gated_delta_net"] - return experimental_attention_variant in linear_attention_variants + from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider + + HAVE_TE = True +except ImportError: + HAVE_TE = False + +try: + import nvidia_kitchen # type: ignore[import-not-found] # pylint: disable=unused-import + + from megatron.core.extensions.kitchen import KitchenSpecProvider + HAVE_KITCHEN = True +except ImportError: + HAVE_KITCHEN = False -def get_gated_delta_net_module_spec_for_backend( - backend: BackendSpecProvider, normalization: Optional[str] = None + +########## +# Experimental Attention Variant Module Specs +########## + + +def get_gated_delta_net_module_spec( + config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: - """Helper function to get module spec for Linear Attention""" - rms_norm = normalization == "RMSNorm" + """Build module spec for GatedDeltaNet attention.""" + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + rms_norm = config.normalization == "RMSNorm" attention = ModuleSpec( module=GatedDeltaNet, submodules=GatedDeltaNetSubmodules( @@ -43,27 +75,22 @@ def get_gated_delta_net_module_spec_for_backend( def get_dsa_module_spec_for_backend( - backend: BackendSpecProvider, - qk_layernorm: Optional[bool] = False, - qk_l2_norm: Optional[bool] = False, - multi_latent_attention: Optional[bool] = False, - mla_down_proj_use_column_parallel: Optional[bool] = False, - normalization: Optional[str] = None, - fallback_to_eager_attn: Optional[bool] = False, + config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: """Helper function to get module spec for Sparse Attention.""" - assert multi_latent_attention, "Currently only MLA supports sparse attention." - assert qk_l2_norm is False, "qk_l2_norm is not supported with MLA." - assert fallback_to_eager_attn is False, "Fallback to eager attention is not supported with DSA." + assert config.multi_latent_attention, "Currently only MLA supports sparse attention." + assert config.qk_l2_norm is False, "qk_l2_norm is not supported with MLA." - linear_q_down_proj = ( - backend.column_parallel_linear() if mla_down_proj_use_column_parallel else backend.linear() + linear_q_up_proj = ( + backend.column_parallel_layer_norm_linear() + if config.qk_layernorm + else backend.column_parallel_linear() ) - linear_kv_down_proj = ( - backend.column_parallel_linear() if mla_down_proj_use_column_parallel else backend.linear() + linear_kv_up_proj = ( + backend.column_parallel_layer_norm_linear() + if config.qk_layernorm + else backend.column_parallel_linear() ) - linear_q_up_proj = backend.column_parallel_linear() - linear_kv_up_proj = backend.column_parallel_linear() # Because TransformerEngine does not support sparse attention yet, we use local # implementation whether the backend is TransformerEngine or not. @@ -82,23 +109,19 @@ def get_dsa_module_spec_for_backend( ), ) - # Adjust for RMS norm. - rms_norm = normalization == "RMSNorm" - qk_norm = backend.layer_norm(rms_norm=rms_norm, for_qk=True) if qk_layernorm else IdentityOp - attention = ModuleSpec( module=MLASelfAttention, params={"attn_mask_type": AttnMaskType.causal}, submodules=MLASelfAttentionSubmodules( linear_q_proj=backend.column_parallel_linear(), - linear_q_down_proj=linear_q_down_proj, + linear_q_down_proj=backend.linear(), linear_q_up_proj=linear_q_up_proj, - linear_kv_down_proj=linear_kv_down_proj, + linear_kv_down_proj=backend.linear(), linear_kv_up_proj=linear_kv_up_proj, core_attention=core_attention, linear_proj=backend.row_parallel_linear(), - q_layernorm=qk_norm, - kv_layernorm=qk_norm, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, ), metainfo={"fuse_input_layernorm": False}, ) @@ -106,33 +129,359 @@ def get_dsa_module_spec_for_backend( return attention -def get_experimental_attention_variant_module_spec_for_backend( - backend: BackendSpecProvider, - sharded_state_dict_keys_map: dict, - experimental_attention_variant: Optional[str] = None, - qk_layernorm: Optional[bool] = False, - qk_l2_norm: Optional[bool] = False, - multi_latent_attention: Optional[bool] = False, - mla_down_proj_use_column_parallel: Optional[bool] = False, - normalization: Optional[str] = None, - fallback_to_eager_attn: Optional[bool] = False, +def get_experimental_attention_variant_module_spec( + config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: - """Helper function to get module spec for Attention""" - if experimental_attention_variant == "gated_delta_net": - return get_gated_delta_net_module_spec_for_backend( - backend=backend, normalization=normalization + """Helper function to get module spec for experimental attention variant""" + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + if config.experimental_attention_variant == "gated_delta_net": + return get_gated_delta_net_module_spec(config=config, backend=backend) + else: + raise ValueError( + f"Invalid experimental attention variant: {config.experimental_attention_variant}" ) - elif experimental_attention_variant == "dsa": - return get_dsa_module_spec_for_backend( - backend=backend, - qk_layernorm=qk_layernorm, - qk_l2_norm=qk_l2_norm, - multi_latent_attention=multi_latent_attention, - mla_down_proj_use_column_parallel=mla_down_proj_use_column_parallel, - normalization=normalization, - fallback_to_eager_attn=fallback_to_eager_attn, + + +########## +# Experimental GPT Decoder Block Spec +########## + + +def get_transformer_layer_with_experimental_attention_variant_spec( + config: TransformerConfig, backend: BackendSpecProvider = None +) -> List[ModuleSpec]: + """Build transformer layer specs with experimental attention variants (e.g., linear attention). + + This function is for constructing a heterogeneous transformer that supports mixing different + attention mechanisms (experimental vs standard) and MLP types (MoE vs dense) across layers. + **Note that, this API is a experimental API in the short term, and might be deprecated in the + future. In the long run, we will move to a new design that better support hybrid models.** + + Key Design: + 1. Attention and MLP patterns: The attention pattern and MLP pattern are orthogonal + and determined independently. This allows flexible combinations (e.g., linear attention + with MoE, or standard attention with dense MLP). + - Attention pattern: derived from `config.linear_attention_freq` or + `config.experimental_attention_variant`. + - MLP pattern: derived from `config.moe_layer_freq`. + + 2. Per-Layer Spec Construction: Iterates through layers, constructing transformer + layer specs based on attention and MLP patterns. + + Args: + config: Transformer configuration containing model hyperparameters and feature flags. + + Returns: + List[ModuleSpec] containing per-layer specs. + + Note: + Currently only supports transformer_engine backend. Kitchen backend can be used as a + wrapper with TE fallback for unsupported operations. + """ + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + # Get attention patterns and specs + experimental_attention_pattern = [0] * config.num_layers + if is_linear_attention_variant(config.experimental_attention_variant): + experimental_attention_pattern = get_linear_attention_pattern(config=config) + elif config.experimental_attention_variant is not None: + experimental_attention_pattern = [1] * config.num_layers + + if 1 in experimental_attention_pattern: + experimental_attention_spec = get_experimental_attention_variant_module_spec( + config=config, backend=backend + ) + else: + experimental_attention_spec = None + + if 0 in experimental_attention_pattern: + standard_attention_spec = _get_self_attention_module_spec(config=config, backend=backend) + else: + standard_attention_spec = None + + # Get MLP patterns and specs + if config.num_moe_experts is not None: + moe_layer_pattern = get_moe_layer_pattern(config=config) + else: + moe_layer_pattern = [0] * config.num_layers + + if 1 in moe_layer_pattern: + moe_layer_spec = _get_moe_module_spec(config=config, backend=backend) + else: + moe_layer_spec = None + + if 0 in moe_layer_pattern: + dense_mlp_layer_spec = _get_dense_mlp_module_spec(config=config, backend=backend) + else: + dense_mlp_layer_spec = None + + # Get GPT decoder block layer specs + rms_norm = config.normalization == "RMSNorm" + layer_specs = [] + for layer_number in range(config.num_layers): + attention = ( + experimental_attention_spec + if experimental_attention_pattern[layer_number] == 1 + else standard_attention_spec + ) + mlp = moe_layer_spec if moe_layer_pattern[layer_number] == 1 else dense_mlp_layer_spec + input_layernorm = ( + IdentityOp + if attention.metainfo["fuse_input_layernorm"] + else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + ) + pre_mlp_layernorm = ( + IdentityOp + if mlp.metainfo["fuse_pre_mlp_layernorm"] + else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + ) + + layer_specs.append( + ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=input_layernorm, + self_attention=attention, + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=pre_mlp_layernorm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + ), + ) + ) + + return layer_specs + + +def get_transformer_block_with_experimental_attention_variant_spec( + config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None +) -> TransformerBlockSubmodules: + """Build transformer block spec with experimental attention variants (e.g., linear attention). + + This function constructs a heterogeneous transformer block that supports mixing different + attention mechanisms (experimental vs standard) and MLP types (MoE vs dense) across layers. + **Note that, this API is a experimental API in the short term, and might be deprecated in the + future. In the long run, we will move to a new design that better support hybrid models.** + + Constructing transformer layer specs by + `get_transformer_layer_with_experimental_attention_variant_spec` and then slicing the + layer specs to only include the layers that are built in this pipeline stage. + + Args: + config: Transformer configuration containing model hyperparameters and feature flags. + vp_stage: Virtual pipeline stage index for interleaved pipeline parallelism. + pp_rank: Pipeline model parallel rank. + + Returns: + TransformerBlockSubmodules containing per-layer specs and final layer norm. + + Note: + Currently only supports transformer_engine backend. Kitchen backend can be used as a + wrapper with TE fallback for unsupported operations. + """ + + backend = _get_backend_spec_provider(config=config) + + layer_specs = get_transformer_layer_with_experimental_attention_variant_spec( + config=config, backend=backend + ) + + # Slice the layer specs to only include the layers that are built in this pipeline stage. + if config.pipeline_model_parallel_layout is not None: + local_layer_ids = config.pipeline_model_parallel_layout.get_layer_id_list( + layer_type=LayerType.decoder, vp_stage=vp_stage, pp_rank=pp_rank + ) + else: + offset = get_transformer_layer_offset(config, vp_stage=vp_stage, pp_rank=pp_rank) + num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage, pp_rank=pp_rank) + local_layer_ids = range(offset, offset + num_layers_to_build) + + layer_specs = [layer_specs[layer_id] for layer_id in local_layer_ids] + + # Get GPT decoder block spec + rms_norm = config.normalization == "RMSNorm" + gpt_decoder_block_spec = TransformerBlockSubmodules( + layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False) + ) + + return gpt_decoder_block_spec + + +########## +# Utilities +########## + + +def is_linear_attention_variant(experimental_attention_variant: Optional[str]) -> bool: + """Check if the experimental attention variant is a linear attention variant.""" + linear_attention_variants = ["gated_delta_net"] + return experimental_attention_variant in linear_attention_variants + + +def get_moe_layer_pattern(config: TransformerConfig) -> List[int]: + """Parse config.moe_layer_freq to get per-layer MoE pattern (1=MoE, 0=dense). + + - int N: one MoE layer every N layers (e.g., N=2 -> [1,0,1,0,...]) + - list: use directly as the pattern.""" + + if isinstance(config.moe_layer_freq, int): + # [1,0,0,...,0,1,0,0,...,0,...] + moe_layer_pattern = [ + 1 if (i % config.moe_layer_freq == 0) else 0 for i in range(config.num_layers) + ] + elif isinstance(config.moe_layer_freq, list): + moe_layer_pattern = config.moe_layer_freq + assert len(moe_layer_pattern) == config.num_layers, ( + f"Invalid length of moe_layer_pattern: {len(moe_layer_pattern)}, " + f"expected {config.num_layers}, " + f"current moe layer pattern: {config.moe_layer_freq}" ) else: raise ValueError( - f"Invalid experimental attention variant: {experimental_attention_variant}" + f"Invalid moe_layer_freq: {type(config.moe_layer_freq)}, {config.moe_layer_freq}" + ) + return moe_layer_pattern + + +def get_linear_attention_pattern(config: TransformerConfig) -> List[int]: + """Parse config.linear_attention_freq to get per-layer attention pattern (1=LA, 0=SDPA). + + - int N: one SDPA layer every N layers (e.g., N=4 -> [1,1,1,0,1,1,1,0,...]) + - list: use directly as the pattern.""" + + if isinstance(config.linear_attention_freq, int): + linear_attention_pattern = [ + # [1,1,...,1,0,1,1,...,1,0,...] + 0 if ((i + 1) % config.linear_attention_freq == 0) else 1 + for i in range(config.num_layers) + ] + elif isinstance(config.linear_attention_freq, list): + linear_attention_pattern = config.linear_attention_freq + assert len(linear_attention_pattern) == config.num_layers, ( + f"Invalid length of linear_attention_pattern: {len(linear_attention_pattern)}, " + f"expected {config.num_layers}, " + f"current linear attention pattern: {config.linear_attention_freq}" + ) + elif config.linear_attention_freq is None: + if not is_linear_attention_variant(config.experimental_attention_variant): + linear_attention_pattern = [0] * config.num_layers + else: + # This should be caught by config validation, but raise here as a safety check + raise ValueError( + f"Linear attention type {config.experimental_attention_variant} is specified " + "but linear_attention_freq is None. " + "Please set linear_attention_freq to specify the LA/SDPA layer pattern." + ) + else: + raise ValueError( + f"Invalid linear_attention_freq: {type(config.linear_attention_freq)}," + f" {config.linear_attention_freq}" + ) + return linear_attention_pattern + + +def _get_backend_spec_provider(config: TransformerConfig) -> BackendSpecProvider: + """Get backend spec provider for experimental attention variant.""" + + assert config.transformer_impl == "transformer_engine", ( + "Experimental GPT decoder block spec only supports " + "transformer engine implementation for now." + ) + backend: BackendSpecProvider = ( + KitchenSpecProvider( + fallback=TESpecProvider(fallback_to_eager_attn=config.fallback_to_eager_attn), + use_kitchen_attention=config.use_kitchen_attention, + kitchen_attention_backend=config.kitchen_attention_backend, ) + if config.use_kitchen + else TESpecProvider() + ) + return backend + + +########## +# Spec functions for non-experimental self attention and MLP layer. +########## + + +def _get_self_attention_module_spec( + config: TransformerConfig, backend: BackendSpecProvider = None +) -> ModuleSpec: + """Get non-experimental self-attention module spec. + For hybrid models that mix experimental and non-experimental attention architectures. + + Warning: This function may be deprecated in the future.""" + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + + layer_spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=config.num_moe_experts, + moe_grouped_gemm=config.moe_grouped_gemm, + qk_layernorm=config.qk_layernorm, + multi_latent_attention=config.multi_latent_attention, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + qk_l2_norm=config.qk_l2_norm, + use_kitchen=config.use_kitchen, + use_te_activation_func=config.use_te_activation_func, + fallback_to_eager_attn=config.fallback_to_eager_attn, + use_kitchen_attention=config.use_kitchen_attention, + kitchen_attention_backend=config.kitchen_attention_backend, + ) + attn_spec = layer_spec.submodules.self_attention + if config.multi_latent_attention: + attn_spec.metainfo["fuse_input_layernorm"] = False + else: + attn_spec.metainfo["fuse_input_layernorm"] = backend.fuse_layernorm_and_linear() + + return attn_spec + + +def _get_dense_mlp_module_spec( + config: TransformerConfig, backend: BackendSpecProvider = None +) -> ModuleSpec: + """Get dense MLP module spec. + For hybrid models that mix dense MLP and experimental attention architectures. + + Warning: This function may be deprecated in the future.""" + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + from megatron.core.models.gpt.gpt_layer_specs import get_mlp_module_spec_for_backend + + mlp_spec = get_mlp_module_spec_for_backend(backend=backend, num_experts=None) + mlp_spec.metainfo["fuse_pre_mlp_layernorm"] = backend.fuse_layernorm_and_linear() + + return mlp_spec + + +def _get_moe_module_spec( + config: TransformerConfig, backend: BackendSpecProvider = None +) -> ModuleSpec: + """Get MoE module spec. + For hybrid models that mix MoE and experimental attention architectures. + + Warning: This function may be deprecated in the future.""" + + if backend is None: + backend = _get_backend_spec_provider(config=config) + + from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec_for_backend + + moe_spec = get_moe_module_spec_for_backend( + backend=backend, + num_experts=config.num_moe_experts, + moe_grouped_gemm=config.moe_grouped_gemm, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + use_te_activation_func=config.use_te_activation_func, + ) + moe_spec.metainfo["fuse_pre_mlp_layernorm"] = False + return moe_spec diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 1db3b939530..70f0a8244ca 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -9,13 +9,8 @@ InferenceSpecProvider, LocalSpecProvider, ) -from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( - get_experimental_attention_variant_module_spec_for_backend, - is_linear_attention_variant, -) from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec_for_backend from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules -from megatron.core.transformer.dot_product_attention import DotProductAttention from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.mlp import MLP, MLPSubmodules @@ -45,7 +40,7 @@ from megatron.core.utils import is_te_min_version try: - import transformer_engine as te # type: ignore[import-untyped] # pylint: disable=unused-import + import transformer_engine as te # pylint: disable=unused-import from megatron.core.extensions.transformer_engine import TEFusedMLP, TENorm from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider @@ -55,7 +50,7 @@ HAVE_TE = False try: - import nvidia_kitchen # type: ignore[import-not-found] # pylint: disable=unused-import + import nvidia_kitchen # pylint: disable=unused-import from megatron.core.extensions.kitchen import KitchenSpecProvider @@ -64,7 +59,7 @@ HAVE_KITCHEN = False try: - import apex # type: ignore[import-untyped] # pylint: disable=unused-import + import apex # pylint: disable=unused-import from megatron.core.fusions.fused_layer_norm import FusedLayerNorm @@ -181,10 +176,8 @@ def get_gpt_layer_with_transformer_engine_spec( moe_grouped_gemm: Optional[bool] = False, qk_layernorm: Optional[bool] = False, multi_latent_attention: Optional[bool] = False, - experimental_attention_variant: Optional[str] = None, fp8: Optional[str] = None, # pylint: disable=unused-argument moe_use_legacy_grouped_gemm: Optional[bool] = False, - normalization: Optional[str] = None, qk_l2_norm: Optional[bool] = False, use_te_op_fuser: Optional[bool] = False, use_kitchen: bool = False, @@ -200,15 +193,10 @@ def get_gpt_layer_with_transformer_engine_spec( num_experts (int, optional): Number of experts. Defaults to None. moe_grouped_gemm (bool, optional): To use Grouped GEMM. Defaults to False. qk_layernorm (bool, optional): To use layernorm for queries/keys. Defaults to False. - multi_latent_attention (bool, optional): To use multi-latent attention. Defaults to False. - experimental_attention_variant (str, optional): The type of experimental attention variant. - Defaults to None. fp8 (str, optional): Deprecated. For temporary Nemo compatibility. moe_use_legacy_grouped_gemm (bool, optional): Force use the legacy GroupedMLP. Defaults to False. - normalization (str, optional): The normalization to use. Defaults to None. qk_l2_norm (bool, optional): To use l2 norm for queries/keys. Defaults to False. - use_kitchen (bool, optional): To use KitchenSpecProvider. Defaults to False. use_te_op_fuser (bool, optional): Use Transformer Engine's operation-based API, which may enable certain operation fusions. Defaults to False. @@ -236,23 +224,8 @@ def get_gpt_layer_with_transformer_engine_spec( else: backend = TESpecProvider(fallback_to_eager_attn=fallback_to_eager_attn) - sharded_state_dict_keys_map = {} - - attention = get_attention_module_spec_for_backend( - backend=backend, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, - experimental_attention_variant=experimental_attention_variant, - qk_layernorm=qk_layernorm, - qk_l2_norm=qk_l2_norm, - multi_latent_attention=multi_latent_attention, - mla_down_proj_use_column_parallel=False, - normalization=normalization, - fallback_to_eager_attn=fallback_to_eager_attn, - ) - mlp = get_mlp_module_spec_for_backend( backend=backend, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, num_experts=num_experts, moe_grouped_gemm=moe_grouped_gemm, moe_use_legacy_grouped_gemm=moe_use_legacy_grouped_gemm, @@ -260,13 +233,77 @@ def get_gpt_layer_with_transformer_engine_spec( use_te_activation_func=use_te_activation_func, ) - return get_transformer_layer_spec_for_backend( - backend=backend, - attention=attention, - mlp=mlp, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, - normalization=normalization, - ) + if multi_latent_attention: + assert qk_l2_norm is False, "qk_l2_norm is not supported with MLA." + linear_q_up_proj = ( + backend.column_parallel_layer_norm_linear() + if qk_layernorm + else backend.column_parallel_linear() + ) + linear_kv_up_proj = ( + backend.column_parallel_layer_norm_linear() + if qk_layernorm + else backend.column_parallel_linear() + ) + return ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=backend.layer_norm(), + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.linear(), + linear_q_up_proj=linear_q_up_proj, + linear_kv_down_proj=backend.linear(), + linear_kv_up_proj=linear_kv_up_proj, + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + ), + ) + else: + qk_norm = backend.layer_norm(for_qk=True) + return ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=SelfAttentionSubmodules( + linear_qkv=backend.column_parallel_layer_norm_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + k_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + ), + ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + sharded_state_dict_keys_map={ + "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", + "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", + "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", + "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", + "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", + "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", + }, + ), + ) def get_gpt_layer_local_spec( @@ -274,7 +311,6 @@ def get_gpt_layer_local_spec( moe_grouped_gemm: Optional[bool] = False, qk_layernorm: Optional[bool] = False, multi_latent_attention: Optional[bool] = False, - experimental_attention_variant: Optional[str] = None, fp8: Optional[str] = None, # pylint: disable=unused-argument moe_use_legacy_grouped_gemm: Optional[bool] = False, normalization: Optional[str] = None, @@ -290,15 +326,10 @@ def get_gpt_layer_local_spec( num_experts (int, optional): Number of experts. Defaults to None. moe_grouped_gemm (bool, optional): To use Grouped GEMM. Defaults to False. qk_layernorm (bool, optional): To use layernorm for queries/keys. Defaults to False. - multi_latent_attention (bool, optional): To use multi-latent attention. Defaults to False. - experimental_attention_variant (str, optional): The type of experimental attention variant. - Defaults to None. fp8 (str, optional): Deprecated. For temporary Nemo compatibility. moe_use_legacy_grouped_gemm (bool, optional): Force use the legacy GroupedMLP. Defaults to False. - normalization (str, optional): The normalization to use. Defaults to None. qk_l2_norm (bool, optional): To use l2 norm for queries/keys. Defaults to False. - use_kitchen (bool, optional): To use KitchenSpecProvider. Defaults to False. Returns: ModuleSpec: Module specification with Megatron-Core modules @@ -313,6 +344,13 @@ def get_gpt_layer_local_spec( ) else: backend = LocalSpecProvider() + # Adjust for RMS norm. + if normalization == "RMSNorm": + layer_norm = backend.layer_norm(rms_norm=True, for_qk=False) + qk_norm = backend.layer_norm(rms_norm=True, for_qk=True) + else: + layer_norm = backend.layer_norm(rms_norm=False, for_qk=False) + qk_norm = backend.layer_norm(rms_norm=False, for_qk=True) if fp8 is not None: warnings.warn( @@ -320,25 +358,6 @@ def get_gpt_layer_local_spec( " and will be removed soon. Please update your code accordingly." ) - if experimental_attention_variant is not None: - raise NotImplementedError( - "Experimental attention variant is not supported with local spec yet." - ) - - sharded_state_dict_keys_map = {} - - attention = get_attention_module_spec_for_backend( - backend=backend, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, - experimental_attention_variant=experimental_attention_variant, - qk_layernorm=qk_layernorm, - qk_l2_norm=qk_l2_norm, - multi_latent_attention=multi_latent_attention, - mla_down_proj_use_column_parallel=True, - normalization=normalization, - fallback_to_eager_attn=False, - ) - mlp = get_mlp_module_spec_for_backend( backend=backend, num_experts=num_experts, @@ -346,170 +365,63 @@ def get_gpt_layer_local_spec( moe_use_legacy_grouped_gemm=moe_use_legacy_grouped_gemm, ) - return get_transformer_layer_spec_for_backend( - backend=backend, - attention=attention, - mlp=mlp, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, - normalization=normalization, - ) - - -def get_transformer_layer_spec_for_backend( - backend: BackendSpecProvider, - attention: ModuleSpec, - mlp: ModuleSpec, - sharded_state_dict_keys_map: Optional[dict] = None, - normalization: Optional[str] = None, -) -> ModuleSpec: - """Helper function to get module spec for TransformerLayer""" - - rms_norm = normalization == "RMSNorm" - - input_layernorm = ( - IdentityOp - if attention.metainfo["fuse_input_layernorm"] - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) - ) - pre_mlp_layernorm = ( - IdentityOp - if mlp.metainfo["fuse_pre_mlp_layernorm"] - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) - ) - - transformer_layer = ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - input_layernorm=input_layernorm, - self_attention=attention, - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=pre_mlp_layernorm, - mlp=mlp, - mlp_bda=get_bias_dropout_add, - sharded_state_dict_keys_map=sharded_state_dict_keys_map, - ), - ) - return transformer_layer - - -def get_attention_module_spec_for_backend( - backend: BackendSpecProvider, - sharded_state_dict_keys_map: dict, - experimental_attention_variant: Optional[str] = None, - qk_layernorm: Optional[bool] = False, - qk_l2_norm: Optional[bool] = False, - multi_latent_attention: Optional[bool] = False, - mla_down_proj_use_column_parallel: Optional[bool] = False, - normalization: Optional[str] = None, - fallback_to_eager_attn: Optional[bool] = False, -) -> ModuleSpec: - """Helper function to get module spec for Attention""" - - if experimental_attention_variant is not None: - return get_experimental_attention_variant_module_spec_for_backend( - backend, - sharded_state_dict_keys_map, - experimental_attention_variant, - qk_layernorm, - qk_l2_norm, - multi_latent_attention, - mla_down_proj_use_column_parallel, - normalization, - fallback_to_eager_attn, - ) - - # Adjust for RMS norm. - rms_norm = normalization == "RMSNorm" - qk_norm = backend.layer_norm(rms_norm=rms_norm, for_qk=True) - - core_attention = backend.core_attention() if not fallback_to_eager_attn else DotProductAttention if multi_latent_attention: assert qk_l2_norm is False, "qk_l2_norm is not supported with MLA." - linear_q_down_proj = ( - backend.column_parallel_linear() - if mla_down_proj_use_column_parallel - else backend.linear() - ) - linear_kv_down_proj = ( - backend.column_parallel_linear() - if mla_down_proj_use_column_parallel - else backend.linear() - ) - linear_q_up_proj = ( - backend.column_parallel_layer_norm_linear() - if qk_layernorm and backend.fuse_layernorm_and_linear() - else backend.column_parallel_linear() - ) - linear_kv_up_proj = ( - backend.column_parallel_layer_norm_linear() - if qk_layernorm and backend.fuse_layernorm_and_linear() - else backend.column_parallel_linear() - ) - qk_norm = ( - backend.layer_norm(rms_norm=rms_norm, for_qk=True) - if qk_layernorm and not backend.fuse_layernorm_and_linear() - else IdentityOp - ) - attention = ModuleSpec( - module=MLASelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=MLASelfAttentionSubmodules( - linear_q_proj=backend.column_parallel_linear(), - linear_q_down_proj=linear_q_down_proj, - linear_q_up_proj=linear_q_up_proj, - linear_kv_down_proj=linear_kv_down_proj, - linear_kv_up_proj=linear_kv_up_proj, - core_attention=core_attention, - linear_proj=backend.row_parallel_linear(), - q_layernorm=qk_norm, - kv_layernorm=qk_norm, + return ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=layer_norm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.column_parallel_linear(), + linear_q_up_proj=backend.column_parallel_linear(), + linear_kv_down_proj=backend.column_parallel_linear(), + linear_kv_up_proj=backend.column_parallel_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=qk_norm if qk_layernorm else IdentityOp, + kv_layernorm=qk_norm if qk_layernorm else IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=layer_norm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, ), - metainfo={"fuse_input_layernorm": False}, ) else: - linear_qkv = ( - backend.column_parallel_layer_norm_linear() - if backend.fuse_layernorm_and_linear() - else backend.column_parallel_linear() - ) - if qk_l2_norm: - qk_norm = L2Norm - elif qk_layernorm: - qk_norm = backend.layer_norm(rms_norm=rms_norm, for_qk=True) - else: - qk_norm = IdentityOp - attention = ModuleSpec( - module=SelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=SelfAttentionSubmodules( - linear_qkv=linear_qkv, - core_attention=core_attention, - linear_proj=backend.row_parallel_linear(), - q_layernorm=qk_norm, - k_layernorm=qk_norm, - ), - metainfo={"fuse_input_layernorm": backend.fuse_layernorm_and_linear()}, - ) - if backend.fuse_layernorm_and_linear(): - sharded_state_dict_keys_map.update( - { - "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", - "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", - "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", - "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", - "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", - "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", - } - ) - else: - sharded_state_dict_keys_map.update( - { + return ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=layer_norm, + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=SelfAttentionSubmodules( + linear_qkv=backend.column_parallel_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + k_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + ), + ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=layer_norm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + sharded_state_dict_keys_map={ "input_layernorm.": "self_attention.linear_qkv.layer_norm_", "pre_mlp_layernorm.": "mlp.linear_fc1.layer_norm_", - } - ) - - return attention + }, + ), + ) def _get_mlp_module_spec( @@ -568,7 +480,6 @@ def get_mlp_module_spec( def get_mlp_module_spec_for_backend( backend: BackendSpecProvider, - sharded_state_dict_keys_map: Optional[dict] = None, num_experts: Optional[int] = None, moe_grouped_gemm: Optional[bool] = False, moe_use_legacy_grouped_gemm: Optional[bool] = False, @@ -586,16 +497,13 @@ def get_mlp_module_spec_for_backend( if backend.fuse_layernorm_and_linear(): linear_fc1 = backend.column_parallel_layer_norm_linear() assert linear_fc1 is not None - fuse_pre_mlp_layernorm = True else: linear_fc1 = backend.column_parallel_linear() - fuse_pre_mlp_layernorm = False return ModuleSpec( module=module, submodules=MLPSubmodules( linear_fc1=linear_fc1, linear_fc2=linear_fc2, activation_func=activation_func ), - metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm}, ) else: # Mixture of experts with modules in megatron core. @@ -613,76 +521,61 @@ def get_gpt_decoder_layer_specs( use_transformer_engine: bool, normalization: Optional[str] = None, qk_l2_norm: Optional[bool] = False, - vp_stage: Optional[int] = None, - pp_rank: Optional[int] = None, ) -> TransformerBlockSubmodules: - """Helper function to get GPT block spec. - - Return a list of transformer layer spec of the current pipeline stage.""" - - get_layer_spec_kwargs = { - "qk_layernorm": config.qk_layernorm, - "moe_use_legacy_grouped_gemm": config.moe_use_legacy_grouped_gemm, - "qk_l2_norm": qk_l2_norm, - "use_kitchen": config.use_kitchen, - "normalization": normalization, - "use_kitchen_attention": config.use_kitchen_attention, - "kitchen_attention_backend": config.kitchen_attention_backend, - } + """GPT block spec.""" + assert config.experimental_attention_variant is None, ( + "Experimental attention variant is not supported with get_gpt_decoder_layer_specs, " + f"but got {config.experimental_attention_variant=}." + ) + if use_transformer_engine: - layer_norm_impl = TENorm - get_layer_spec_kwargs["use_te_activation_func"] = config.use_te_activation_func - get_layer_spec_kwargs['fallback_to_eager_attn'] = config.fallback_to_eager_attn - get_layer_spec_fn = get_gpt_layer_with_transformer_engine_spec + dense_layer_spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=None, + moe_grouped_gemm=False, + qk_layernorm=config.qk_layernorm, + multi_latent_attention=config.multi_latent_attention, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + qk_l2_norm=qk_l2_norm, + use_kitchen=config.use_kitchen, + use_te_activation_func=config.use_te_activation_func, + ) + moe_layer_spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=config.num_moe_experts, + moe_grouped_gemm=config.moe_grouped_gemm, + qk_layernorm=config.qk_layernorm, + multi_latent_attention=config.multi_latent_attention, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + qk_l2_norm=qk_l2_norm, + use_kitchen=config.use_kitchen, + use_te_activation_func=config.use_te_activation_func, + ) else: - layer_norm_impl = LNImpl - get_layer_spec_fn = get_gpt_layer_local_spec - - layer_spec_dict = {} - for mlp_type in ["dense", "moe"]: - for attention_type in ["softmax_attention", "linear_attention"]: - if mlp_type == "moe": - if config.moe_layer_freq is None: - # Skip if there is no MoE layer in the model. - continue - num_experts = config.num_moe_experts - moe_grouped_gemm = config.moe_grouped_gemm - else: - num_experts = None - moe_grouped_gemm = None - if attention_type == "linear_attention": - multi_latent_attention = None - if is_linear_attention_variant(config.experimental_attention_variant): - # There exists linear attention layer in the model. - experimental_attention_variant = config.experimental_attention_variant - else: - # Skip if there is no linear attention layer in the model. - continue - else: - multi_latent_attention = config.multi_latent_attention - if is_linear_attention_variant(config.experimental_attention_variant): - # experimental_attention_variant is a linear attention variant, - # so softmax attention is regular attention layer. - experimental_attention_variant = None - else: - # Softmax attention is an experimental attention variant. - experimental_attention_variant = config.experimental_attention_variant - - layer_spec_key = f"{mlp_type}_{attention_type}" - layer_spec_dict[layer_spec_key] = get_layer_spec_fn( - num_experts=num_experts, - moe_grouped_gemm=moe_grouped_gemm, - multi_latent_attention=multi_latent_attention, - experimental_attention_variant=experimental_attention_variant, - **get_layer_spec_kwargs, - ) + dense_layer_spec = get_gpt_layer_local_spec( + num_experts=None, + moe_grouped_gemm=False, + qk_layernorm=config.qk_layernorm, + multi_latent_attention=config.multi_latent_attention, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + normalization=normalization, + qk_l2_norm=qk_l2_norm, + use_kitchen=config.use_kitchen, + ) + moe_layer_spec = get_gpt_layer_local_spec( + num_experts=config.num_moe_experts, + moe_grouped_gemm=config.moe_grouped_gemm, + qk_layernorm=config.qk_layernorm, + multi_latent_attention=config.multi_latent_attention, + moe_use_legacy_grouped_gemm=config.moe_use_legacy_grouped_gemm, + normalization=normalization, + qk_l2_norm=qk_l2_norm, + use_kitchen=config.use_kitchen, + ) # Parse config.moe_layer_freq to determine the pattern of expert/dense layers. # 0 stands for dense layers, 1 stands for expert layers. # For integer N: Creates a pattern with one expert layer every N layers. # For string pattern: Evaluates the str directly (e.g. "[1,0,1]" for alternating expert/dense). if isinstance(config.moe_layer_freq, int): - # [1,0,0,...,0,1,0,0,...,0,...] moe_layer_pattern = [ 1 if (i % config.moe_layer_freq == 0) else 0 for i in range(config.num_layers) ] @@ -698,50 +591,15 @@ def get_gpt_decoder_layer_specs( f"Invalid moe_layer_freq: {type(config.moe_layer_freq)}, {config.moe_layer_freq}" ) - # Parse config.linear_attention_freq to determine the pattern of expert/dense layers. - # 0 stands for SDPA layers, 1 stands for LA layers. - # For integer N: Creates a pattern with (N-1) LA layers and 1 SDPA layer every N layers. - # For string pattern: Evaluates the str directly (e.g. "[1,0,1]" for alternating LA/SDPA). - if isinstance(config.linear_attention_freq, int): - linear_attention_pattern = [ - # [1,1,...,1,0,1,1,...,1,0,...] - 0 if ((i + 1) % config.linear_attention_freq == 0) else 1 - for i in range(config.num_layers) - ] - elif isinstance(config.linear_attention_freq, list): - linear_attention_pattern = config.linear_attention_freq - assert len(linear_attention_pattern) == config.num_layers, ( - f"Invalid length of linear_attention_pattern: {len(linear_attention_pattern)}, " - f"expected {config.num_layers}, " - f"current linear attention pattern: {config.linear_attention_freq}" - ) - elif config.linear_attention_freq is None: - if not is_linear_attention_variant(config.experimental_attention_variant): - linear_attention_pattern = [0] * config.num_layers - else: - linear_attention_pattern = [1] * config.num_layers - warnings.warn( - f"Linear attention type {config.experimental_attention_variant} is specified " - "but linear_attention_freq is None. " - "Setting linear_attention_pattern to [1] * config.num_layers as default." - ) - else: - raise ValueError( - f"Invalid linear_attention_freq: {type(config.linear_attention_freq)}," - f" {config.linear_attention_freq}" - ) - # Create the layer specs for the model. layer_specs = [] for layer_number in range(config.num_layers): - mlp_type = "moe" if moe_layer_pattern[layer_number] else "dense" - attention_type = ( - "linear_attention" if linear_attention_pattern[layer_number] else "softmax_attention" - ) - layer_spec_key = f"{mlp_type}_{attention_type}" - if layer_spec_key not in layer_spec_dict: - raise ValueError(f"Invalid layer spec key: {layer_spec_key}") - layer_specs.append(layer_spec_dict[layer_spec_key]) + if moe_layer_pattern[layer_number] == 1: + layer_specs.append(moe_layer_spec) + elif moe_layer_pattern[layer_number] == 0: + layer_specs.append(dense_layer_spec) + else: + raise ValueError(f"Invalid layer pattern: {moe_layer_pattern}") return layer_specs @@ -758,13 +616,16 @@ def get_gpt_decoder_block_spec( layer_specs = get_gpt_decoder_layer_specs( config, use_transformer_engine, normalization, qk_l2_norm ) + # Slice the layer specs to only include the layers that are built in this pipeline stage. # Note: MCore layer_number starts at 1 num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage, pp_rank=pp_rank) if config.pipeline_model_parallel_layout is not None: layout = config.pipeline_model_parallel_layout - assert isinstance(layout, PipelineParallelLayerLayout) + assert isinstance( + layout, PipelineParallelLayerLayout + ), f"Invalid pipeline model parallel layout: {layout}" local_layer_specs = [ layer_specs[layer_id] for layer_id in layout.get_layer_id_list( @@ -775,11 +636,11 @@ def get_gpt_decoder_block_spec( offset = get_transformer_layer_offset(config, vp_stage=vp_stage, pp_rank=pp_rank) local_layer_specs = layer_specs[offset : offset + num_layers_to_build] + # Block spec. if use_transformer_engine: layer_norm_impl = TENorm else: layer_norm_impl = LNImpl - # Block spec. block_spec = TransformerBlockSubmodules( layer_specs=local_layer_specs, layer_norm=layer_norm_impl ) @@ -796,22 +657,17 @@ def get_gpt_mtp_block_spec( ) -> MultiTokenPredictionBlockSubmodules: """GPT Multi-Token Prediction (MTP) block spec.""" if use_transformer_engine: - backend: BackendSpecProvider = ( - KitchenSpecProvider( + if config.use_kitchen: + backend: BackendSpecProvider = KitchenSpecProvider( fallback=TESpecProvider(fallback_to_eager_attn=config.fallback_to_eager_attn), use_kitchen_attention=config.use_kitchen_attention, kitchen_attention_backend=config.kitchen_attention_backend, ) - if config.use_kitchen - else TESpecProvider(fallback_to_eager_attn=config.fallback_to_eager_attn) - ) + else: + backend = TESpecProvider(fallback_to_eager_attn=config.fallback_to_eager_attn) else: backend = ( - KitchenSpecProvider( - fallback=LocalSpecProvider(), - use_kitchen_attention=config.use_kitchen_attention, - kitchen_attention_backend=config.kitchen_attention_backend, - ) + KitchenSpecProvider(fallback=LocalSpecProvider()) if config.use_kitchen else LocalSpecProvider() ) diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py index 2b0a18b433b..16dc3a79ebb 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py @@ -104,7 +104,9 @@ def __init__( """ if not HAVE_FLA: - raise ImportError("FLA is not installed. Please install it with `pip install fla`.") + raise ImportError( + "FLA is not installed. Please install it with `pip install flash-linear-attention`." + ) super().__init__(config) @@ -246,7 +248,7 @@ def reset_parameters(self): dtype=self.config.params_dtype, device=torch.cuda.current_device(), ).uniform_(*self.A_init_range) - self.A_log.data.copy_(A) + self.A_log.data.copy_(torch.log(A)) def forward( self, diff --git a/megatron/core/transformer/dot_product_attention_context_parallel.py b/megatron/core/transformer/dot_product_attention_context_parallel.py index 89659a1d743..aaf08d40ade 100644 --- a/megatron/core/transformer/dot_product_attention_context_parallel.py +++ b/megatron/core/transformer/dot_product_attention_context_parallel.py @@ -185,6 +185,9 @@ def forward(ctx, q, k, v, attention_mask, attention_dropout, softmax_scale, pg): comm.all_gather(kv_buffer_copy[1], v_0) # Prepare attention bias + assert ( + attention_mask is not None + ), "Attention mask is required for the native attention function with context parallelism" attn_bias = to_zz_mask_attn_bias( attention_mask, cp_size, nheads, nheads_k, heads_k_stride, q.device, q.dtype ) diff --git a/megatron/core/transformer/spec_utils.py b/megatron/core/transformer/spec_utils.py index 24df1add0eb..dbd2e08bccb 100644 --- a/megatron/core/transformer/spec_utils.py +++ b/megatron/core/transformer/spec_utils.py @@ -46,6 +46,7 @@ def import_module(module_path: Tuple[str]): return vars(module)[name] +# pylint: disable=missing-function-docstring def get_module(spec_or_module: Union[ModuleSpec, type], **additional_kwargs): """Retrieve the module class or function specified by a ModuleSpec or return it as is if already provided. diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index df11daeb095..6a9fc64f18d 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -194,6 +194,9 @@ class TransformerConfig(ModelParallelConfig): qk_layernorm: bool = False """Whether to apply `normalization` type of normalization to the query and key embeddings.""" + qk_l2_norm: bool = False + """Whether to apply llama 4-style qk L2 norm.""" + qk_clip: bool = False """Whether to clip the query and key weights. Needed for Muon MLA Model training.""" @@ -234,7 +237,26 @@ class TransformerConfig(ModelParallelConfig): """Type of attention variant to use. Currently support gated_delta_net and dsa.""" #################### - # attention variant: gated_delta_net + # DSA + #################### + dsa_indexer_n_heads: Optional[int] = None + """Number of DSA indexer heads.""" + + dsa_indexer_head_dim: Optional[int] = None + """Dimension per DSA indexer head.""" + + dsa_indexer_topk: Optional[int] = None + """Number of top-k tokens to select in DSA indexer.""" + + dsa_indexer_loss_coeff: Optional[float] = None + """Coefficient for the DSA indexer KL divergence loss. Set to 0 to disable indexer loss.""" + + dsa_indexer_use_sparse_loss: Optional[bool] = None + """Whether to use sparse DSA indexer loss. If True, the indexer loss will be computed using the + top-k indices.""" + + #################### + # linear attention #################### linear_attention_type: Optional[str] = None """Type of linear attention to use. @@ -262,25 +284,6 @@ class TransformerConfig(ModelParallelConfig): linear_num_value_heads: Optional[int] = None """Number of value and gate heads for the gated delta net.""" - #################### - # attention variant: dsa - #################### - dsa_indexer_n_heads: Optional[int] = None - """Number of DSA indexer heads.""" - - dsa_indexer_head_dim: Optional[int] = None - """Dimension per DSA indexer head.""" - - dsa_indexer_topk: Optional[int] = None - """Number of top-k tokens to select in DSA indexer.""" - - dsa_indexer_loss_coeff: Optional[float] = None - """Coefficient for the DSA indexer KL divergence loss. Set to 0 to disable indexer loss.""" - - dsa_indexer_use_sparse_loss: Optional[bool] = None - """Whether to use sparse DSA indexer loss. If True, the indexer loss will be computed using the - top-k indices.""" - #################### # initialization #################### diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 5f9e7350c18..8f621722e31 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2441,7 +2441,6 @@ def _add_training_args(parser): 'which only ensures bitwise identical results when the same inputs are processed in the same batch configuration. ' 'This will significantly affect speed of training and inference as the kernels are not full optimized.') - return parser @@ -3419,7 +3418,17 @@ def _add_experimental_attention_variant_args(parser): group = parser.add_argument_group(title="experimental_attention_variant") group.add_argument('--experimental-attention-variant', default=None, choices=['gated_delta_net', 'dsa'], type=str, help='Type of attention variant to use. Currently support gated_delta_net and dsa.') - + # DSA + group.add_argument('--dsa-indexer-n-heads', default=None, type=int, + help='Number of indexer heads for sparse attention. If not set, defaults to num-attention-heads.') + group.add_argument('--dsa-indexer-head-dim', default=None, type=int, + help='Dimension per indexer head for sparse attention. If not set, defaults to kv-channels.') + group.add_argument('--dsa-indexer-topk', default=None, type=int, + help='Number of top-k tokens to select in sparse attention indexer.') + group.add_argument('--dsa-indexer-loss-coeff', default=0.0, type=float, + help='Coefficient for the indexer KL divergence loss. Set to 0 to disable indexer loss.') + group.add_argument('--dsa-indexer-use-sparse-loss', action='store_true', + help='Use sparse indexer loss. If set, the indexer loss will be computed using the top-k indices.') # Linear attention group.add_argument('--linear-attention-type', default=None, choices=['gated_delta_net'], type=str, help='(Deprecated, use --experimental-attention-variant instead) Type of linear attention to use. Currently support gated_delta_net.') @@ -3442,19 +3451,6 @@ def _add_experimental_attention_variant_args(parser): help='Number of query and key heads for the gated delta net.') group.add_argument('--linear-num-value-heads', default=32, type=int, help='Number of value and gate heads for the gated delta net.') - - # DSA - group.add_argument('--dsa-indexer-n-heads', default=None, type=int, - help='Number of indexer heads for sparse attention. If not set, defaults to num-attention-heads.') - group.add_argument('--dsa-indexer-head-dim', default=None, type=int, - help='Dimension per indexer head for sparse attention. If not set, defaults to kv-channels.') - group.add_argument('--dsa-indexer-topk', default=None, type=int, - help='Number of top-k tokens to select in sparse attention indexer.') - group.add_argument('--dsa-indexer-loss-coeff', default=0.0, type=float, - help='Coefficient for the indexer KL divergence loss. Set to 0 to disable indexer loss.') - group.add_argument('--dsa-indexer-use-sparse-loss', action='store_true', - help='Use sparse indexer loss. If set, the indexer loss will be computed using the top-k indices.') - return parser def _add_heterogeneous_args(parser): diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 77b17b07e13..f7ff7cd2775 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -1472,13 +1472,13 @@ def load_checkpoint(ddp_model, optimizer, opt_param_scheduler, load_arg='load', ckpt_args = state_dict.get("args") if not hasattr(ckpt_args, "tensor_model_parallel_size"): - print_rank_0("WARNING: TP size not found in checkpoint args, using 0 as default.") + print_rank_0("WARNING: TP size not found in checkpoint args, using 1 as default.") if not hasattr(ckpt_args, "pipeline_model_parallel_size"): - print_rank_0("WARNING: PP size not found in checkpoint args, using 0 as default.") + print_rank_0("WARNING: PP size not found in checkpoint args, using 1 as default.") ckpt_tp_pp = ( - getattr(ckpt_args, "tensor_model_parallel_size", 0), - getattr(ckpt_args, "pipeline_model_parallel_size", 0), + getattr(ckpt_args, "tensor_model_parallel_size", 1), + getattr(ckpt_args, "pipeline_model_parallel_size", 1), ) run_tp_pp = ( args.tensor_model_parallel_size, diff --git a/megatron/training/training.py b/megatron/training/training.py index 845d271f62e..b3290626be0 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -332,18 +332,15 @@ def transformer_flops(): if args.moe_shared_expert_intermediate_size is None else args.moe_shared_expert_intermediate_size ) - # SwiGLU. - gated_linear_multiplier = 3 / 2 if args.swiglu else 1 - # The 12x term below comes from the following factors; for more details, see - # "APPENDIX: FLOATING-POINT OPERATIONS" in https://arxiv.org/abs/2104.04473. # - 3x: Each GEMM in the model needs to be performed 3 times (forward pass, # backward wgrad [weight gradient], backward dgrad [data gradient]). - # - 2x: GEMMs of a particular size are stacked twice in the standard Transformer model - # architectures implemented in this codebase (e.g., h->ffn_h GEMM and ffn_h->h GEMM - # in MLP layer). + forward_backward_expansion_factor = 3 # - 2x: A GEMM of a m*n tensor with a n*k tensor requires 2mnk floating-point operations. - expansion_factor = 3 * 2 * 2 + fma_expansion_factor = 2 + # - 3x (SwiGLU enabled): h->2*ffn_h GEMM and ffn_h->h GEMM are stacked. + # - 2x (SwiGLU disabled): h->ffn_h GEMM and ffn_h->h GEMM are stacked. + ffn_expansion_factor = 3 if args.swiglu else 2 if args.multi_latent_attention: assert not args.group_query_attention @@ -374,8 +371,8 @@ def transformer_flops(): + 1 ) standard_self_attn_term = ( - 3 - * 2 # fwd(1) + bwd(2) *FMA + forward_backward_expansion_factor + * fma_expansion_factor * ( ## q lora + rope + q norm q_term @@ -402,13 +399,19 @@ def transformer_flops(): query_projection_size = args.kv_channels * args.num_attention_heads key_projection_size = args.kv_channels * args.num_query_groups value_projection_size = args.kv_channels * args.num_query_groups + gate_projection_size = query_projection_size if args.attention_output_gate else 0 standard_self_attn_term = ( - 3 - * 2 # fwd(1) + bwd(2) *FMA + forward_backward_expansion_factor + * fma_expansion_factor * ( ## qkv proj args.hidden_size - * (query_projection_size + key_projection_size + value_projection_size) + * ( + query_projection_size + + key_projection_size + + value_projection_size + + gate_projection_size + ) ## core attention + query_projection_size * args.seq_length @@ -436,7 +439,12 @@ def transformer_flops(): f"current linear attention pattern: {args.linear_attention_freq}" ) elif args.linear_attention_freq is None: - linear_attention_pattern = [1] * num_layers + # This should be caught by config validation, but raise here as a safety check + raise ValueError( + f"Linear attention type {args.experimental_attention_variant} is specified " + "but linear_attention_freq is None. " + "Please set linear_attention_freq to specify the LA/SDPA layer pattern." + ) else: raise ValueError( f"Invalid linear_attention_freq: {type(args.linear_attention_freq)}," @@ -454,8 +462,8 @@ def transformer_flops(): qk_dim = qk_head_dim * num_qk_heads v_dim = v_head_dim * num_v_heads linear_self_attn_term = ( - 3 - * 2 # fwd(1) + bwd(2) *FMA + forward_backward_expansion_factor + * fma_expansion_factor * ( ## in proj args.hidden_size @@ -492,25 +500,25 @@ def transformer_flops(): * args.seq_length * ( # MLP - expansion_factor - * num_layers + forward_backward_expansion_factor + * fma_expansion_factor * args.hidden_size * ( # dense layer (deepseek v2, v3 style) - (args.ffn_hidden_size * gated_linear_multiplier) - * (num_dense_layers / num_layers) + (args.ffn_hidden_size * ffn_expansion_factor) + * num_dense_layers # routed experts - + (moe_ffn_hidden_size * num_experts_routed_to * gated_linear_multiplier) - * (num_moe_layers / num_layers) + + (moe_ffn_hidden_size * num_experts_routed_to * ffn_expansion_factor) + * num_moe_layers # Shared Experts. - + (shared_expert_ffn_hidden_size * gated_linear_multiplier) - * (num_moe_layers / num_layers) + + (shared_expert_ffn_hidden_size * ffn_expansion_factor) + * num_moe_layers ) # Self Attention + self_attn_term # MTP norms and proj - + 3 - * 2 + + forward_backward_expansion_factor + * fma_expansion_factor * mtp_num_layers * ( # MTP eh norm + final nrom @@ -519,7 +527,11 @@ def transformer_flops(): + 2 * args.hidden_size * args.hidden_size ) # Logit. - + 3 * 2 * args.hidden_size * args.padded_vocab_size * (mtp_num_layers + 1) + + forward_backward_expansion_factor + * fma_expansion_factor + * args.hidden_size + * args.padded_vocab_size + * (mtp_num_layers + 1) # MTP + final logit ) ) return total_floating_point_operations diff --git a/tests/unit_tests/post_training/test_modelopt_module_spec.py b/tests/unit_tests/post_training/test_modelopt_module_spec.py index ec80fcb1a72..dac96785bc0 100644 --- a/tests/unit_tests/post_training/test_modelopt_module_spec.py +++ b/tests/unit_tests/post_training/test_modelopt_module_spec.py @@ -173,6 +173,7 @@ def setup_method(self, method): moe_ffn_hidden_size=128, moe_shared_expert_intermediate_size=128, qk_layernorm=True, + qk_l2_norm=True, use_cpu_initialization=True, ) default_spec = get_gpt_decoder_block_spec( diff --git a/tests/unit_tests/ssm/test_gated_delta_net.py b/tests/unit_tests/ssm/test_gated_delta_net.py index 725d18fbc06..81f8eed0574 100644 --- a/tests/unit_tests/ssm/test_gated_delta_net.py +++ b/tests/unit_tests/ssm/test_gated_delta_net.py @@ -11,7 +11,10 @@ from megatron.core.models.common.embeddings.rope_utils import ( get_pos_emb_on_this_cp_rank as get_tensor_on_this_cp_rank, ) -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_module_spec, + get_transformer_block_with_experimental_attention_variant_spec, +) from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.ssm.gated_delta_net import GatedDeltaNet @@ -82,10 +85,13 @@ def setup_method(self, tp_size, sp, cp_size): tensor_model_parallel_size=tp_size, sequence_parallel=sp, context_parallel_size=cp_size, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + transformer_impl="transformer_engine", ) - gdn_submodules = get_gpt_layer_with_transformer_engine_spec( - experimental_attention_variant="gated_delta_net", normalization="RMSNorm" - ).submodules.self_attention.submodules + gdn_submodules = get_experimental_attention_variant_module_spec( + config=self.transformer_config + ).submodules self.gdn = GatedDeltaNet( self.transformer_config, @@ -159,10 +165,13 @@ def test_parallel_gated_delta_net_correctness(tmp_path_dist_ckpt, tp, sp, cp): num_attention_heads=8, activation_func=F.silu, bf16=True, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + transformer_impl="transformer_engine", ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( - experimental_attention_variant="gated_delta_net", normalization="RMSNorm" + transformer_layer_spec = get_transformer_block_with_experimental_attention_variant_spec( + config=transformer_config, vp_stage=None, pp_rank=0 ) if cp: @@ -171,5 +180,15 @@ def test_parallel_gated_delta_net_correctness(tmp_path_dist_ckpt, tp, sp, cp): atol, rtol = 5e-4, 5e-4 _test_parallel_attention_correctness( - transformer_config, transformer_layer_spec, tmp_path_dist_ckpt, tp, sp, cp + transformer_config=transformer_config, + transformer_layer_spec=transformer_layer_spec, + tmp_path_dist_ckpt=tmp_path_dist_ckpt, + atol=atol, + rtol=rtol, + tp=tp, + sp=sp, + cp=cp, + seed=123, + sequence_length=256, + micro_batch_size=4, ) diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index cd7ca916091..b5f2857d622 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -875,6 +875,7 @@ def get_tensor_on_this_rank(tensor): Utils.destroy_model_parallel() +# TODO(yuzhongw): Add test case for fallback_to_eager_attn @pytest.mark.parametrize("apply_rope_fusion", [False, True]) @pytest.mark.parametrize( ("tp", "sp", "cp"), @@ -887,25 +888,15 @@ def get_tensor_on_this_rank(tensor): ], ) @pytest.mark.parametrize("qk_layernorm", [False, True]) -@pytest.mark.parametrize("fallback_to_eager_attn", [False, True]) @pytest.mark.parametrize("output_gate", [False, True]) def test_parallel_attention_correctness( - tmp_path_dist_ckpt, - apply_rope_fusion, - tp, - sp, - cp, - qk_layernorm, - fallback_to_eager_attn, - output_gate, + tmp_path_dist_ckpt, apply_rope_fusion, tp, sp, cp, qk_layernorm, output_gate ): transformer_config = TransformerConfig( num_layers=1, hidden_size=128, num_attention_heads=4, - context_parallel_size=1, - tensor_model_parallel_size=1, - sequence_parallel=False, + normalization="RMSNorm", bf16=True, qk_layernorm=qk_layernorm, apply_rope_fusion=apply_rope_fusion, @@ -914,24 +905,20 @@ def test_parallel_attention_correctness( attention_dropout=0.0, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( - fallback_to_eager_attn=fallback_to_eager_attn, - normalization="RMSNorm", - qk_layernorm=qk_layernorm, - ) - if cp > 1: - if qk_layernorm: - atol, rtol = 2e-2, 2e-2 - else: - atol, rtol = 5e-3, 5e-3 - else: - if qk_layernorm: - atol, rtol = 1e-2, 1e-2 - else: - atol, rtol = 2e-3, 2e-3 + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(qk_layernorm=qk_layernorm) + atol, rtol = 1e-2, 1e-2 _test_parallel_attention_correctness( - transformer_config, transformer_layer_spec, tmp_path_dist_ckpt, tp, sp, cp + transformer_config, + transformer_layer_spec, + tmp_path_dist_ckpt, + atol=atol, + rtol=rtol, + tp=tp, + sp=sp, + cp=cp, + seed=123, + sequence_length=256, )