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
2 changes: 2 additions & 0 deletions gpt_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@ def _get_transformer_layer_spec(use_te, config):
args.moe_grouped_gemm,
args.qk_layernorm,
args.multi_latent_attention,
args.experimental_attention_variant,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
qk_l2_norm=args.qk_l2_norm,
use_kitchen=config.use_kitchen,
Expand All @@ -141,6 +142,7 @@ def _get_transformer_layer_spec(use_te, config):
args.moe_grouped_gemm,
args.qk_layernorm,
args.multi_latent_attention,
args.experimental_attention_variant,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
normalization=args.normalization,
use_kitchen=config.use_kitchen,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.

from typing import Optional

from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add
from megatron.core.models.backends import BackendSpecProvider
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.multi_latent_attention import (
MLASelfAttention,
MLASelfAttentionSubmodules,
)
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.transformer_layer import TransformerLayer, TransformerLayerSubmodules


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,
num_experts: Optional[int] = None,
mlp: Optional[ModuleSpec] = 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."

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()
)

# Because TransformerEngine does not support sparse attention yet, we use local
# implementation whether the backend is TransformerEngine or not.
core_attention = ModuleSpec(
module=DSAttention,
submodules=DSAttentionSubmodules(
indexer=ModuleSpec(
module=DSAIndexer,
submodules=DSAIndexerSubmodules(
linear_wq_b=backend.linear(),
linear_wk=backend.linear(),
k_norm=backend.layer_norm(rms_norm=False, for_qk=True),
linear_weights_proj=backend.linear(),
),
)
),
)

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=core_attention,
linear_proj=backend.row_parallel_linear(),
q_layernorm=IdentityOp,
kv_layernorm=IdentityOp,
),
)

return ModuleSpec(
module=TransformerLayer,
submodules=TransformerLayerSubmodules(
input_layernorm=backend.layer_norm(),
self_attention=attention,
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,
),
)


def get_experimental_attention_variant_module_spec_for_backend(
backend: BackendSpecProvider,
experimental_attention_variant: Optional[str] = None,
qk_layernorm: Optional[bool] = False,
qk_l2_norm: Optional[bool] = False,
multi_latent_attention: Optional[bool] = False,
num_experts: Optional[int] = None,
mlp: Optional[ModuleSpec] = None,
) -> ModuleSpec:
"""Helper function to get module spec for Attention"""
if 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,
num_experts=num_experts,
mlp=mlp,
)
else:
raise ValueError(
f"Invalid experimental attention variant: {experimental_attention_variant}"
)
Comment thread
santhnm2 marked this conversation as resolved.
27 changes: 27 additions & 0 deletions megatron/core/models/gpt/gpt_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,9 @@
InferenceSpecProvider,
LocalSpecProvider,
)
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
get_experimental_attention_variant_module_spec_for_backend,
)
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.enums import AttnMaskType, LayerType
Expand Down Expand Up @@ -176,6 +179,7 @@ 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,
qk_l2_norm: Optional[bool] = False,
Expand All @@ -192,6 +196,9 @@ 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 MLA. 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.
Expand Down Expand Up @@ -232,6 +239,17 @@ def get_gpt_layer_with_transformer_engine_spec(
use_te_activation_func=use_te_activation_func,
)

if experimental_attention_variant is not None:
return get_experimental_attention_variant_module_spec_for_backend(
backend=backend,
experimental_attention_variant=experimental_attention_variant,
qk_layernorm=qk_layernorm,
qk_l2_norm=qk_l2_norm,
multi_latent_attention=multi_latent_attention,
num_experts=num_experts,
mlp=mlp,
)

if multi_latent_attention:
assert qk_l2_norm is False, "qk_l2_norm is not supported with MLA."
linear_q_up_proj = (
Expand Down Expand Up @@ -310,6 +328,7 @@ 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,
Expand All @@ -325,6 +344,9 @@ 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 MLA. 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.
Expand All @@ -334,6 +356,11 @@ def get_gpt_layer_local_spec(
ModuleSpec: Module specification with Megatron-Core modules
"""

if experimental_attention_variant is not None:
raise NotImplementedError(
"Experimental attention variant is not supported with local spec yet."
)

if use_kitchen:
assert HAVE_KITCHEN
backend = KitchenSpecProvider(
Expand Down
Loading
Loading