Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
31 commits
Select commit Hold shift + click to select a range
2c7f33e
Initial draft version
kunlunl Nov 6, 2025
501ad65
Run through single GPU version
kunlunl Nov 10, 2025
c3754a7
Fix the usage of attention_mask
kunlunl Nov 10, 2025
451a36b
Update mask & indexer loss
kunlunl Nov 10, 2025
186dcdc
Minor changes about code style
kunlunl Nov 11, 2025
5271406
Fix TP
kunlunl Nov 17, 2025
6e6fb50
Fix attn mask and norm type
kunlunl Nov 18, 2025
237164e
Format
kunlunl Nov 18, 2025
72ba916
Resolve minor comments
kunlunl Nov 18, 2025
0e77355
Add unit test
kunlunl Nov 18, 2025
8851f4e
Fix UT
kunlunl Nov 20, 2025
7cdedf1
Use hidden_size to initialize indexer when q_lora_rank is None
kunlunl Nov 24, 2025
c915e4a
Add fused tilelang kernels
kunlunl Nov 25, 2025
de64232
Add indexer loss tracker
kunlunl Nov 25, 2025
1d6f4d3
Fix sparse indexer loss
kunlunl Nov 26, 2025
8495e03
Address minor comments
kunlunl Nov 27, 2025
ad04d58
Temporarily delete fused kernels
kunlunl Nov 27, 2025
9dae1ce
Fix lint error
kunlunl Nov 27, 2025
9cc697b
Make variable/class names more specific (SparseAttention -> DSA)
kunlunl Nov 27, 2025
b8f5c46
Merge branch 'dev' into kunlunl/deepseek_v3.2
kunlunl Nov 27, 2025
0d7e1d1
Fix lint error
kunlunl Nov 27, 2025
6630c7f
Merge linear-attention-type and sparse-attention-type to experimental…
kunlunl Nov 28, 2025
4b7f3c5
Merge branch 'dev' into kunlunl/deepseek_v3.2
kunlunl Nov 28, 2025
b1e309e
Rename test_sparse_attention to test_attention_variant_dsa
kunlunl Nov 28, 2025
6fe2f35
Minor fixes
kunlunl Nov 28, 2025
33684d2
Fix qk norm spec in dsa spec
kunlunl Nov 28, 2025
a0b6fd9
Add fast-hadamard-transform to pyproject.toml
kunlunl Nov 28, 2025
f22344a
Remove fast-hadamard-transform
kunlunl Dec 1, 2025
ce99e9f
Minor fix for args
kunlunl Dec 1, 2025
23ba310
Merge branch 'dev' into kunlunl/deepseek_v3.2
kunlunl Dec 1, 2025
cbfa053
Add mock hadamard_transformer
kunlunl Dec 1, 2025
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
7 changes: 4 additions & 3 deletions gpt_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,8 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None):
else:
use_te = args.transformer_impl == "transformer_engine"

if args.num_experts or (args.linear_attention_type is not None):
linear_attention_variants = ["gated_delta_net"]
if args.num_experts or args.experimental_attention_variant in linear_attention_variants:
# Define the decoder block spec
transformer_layer_spec = get_gpt_decoder_block_spec(
config,
Expand Down Expand Up @@ -114,7 +115,7 @@ def _get_transformer_layer_spec(use_te, config):
args.moe_grouped_gemm,
args.qk_layernorm,
args.multi_latent_attention,
args.linear_attention_type,
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 @@ -126,7 +127,7 @@ def _get_transformer_layer_spec(use_te, config):
args.moe_grouped_gemm,
args.qk_layernorm,
args.multi_latent_attention,
args.linear_attention_type,
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,132 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.

from typing import Optional

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.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


def get_gated_delta_net_module_spec_for_backend(
backend: BackendSpecProvider, normalization: Optional[str] = None
) -> ModuleSpec:
"""Helper function to get module spec for Linear Attention"""
rms_norm = normalization == "RMSNorm"
attention = ModuleSpec(
module=GatedDeltaNet,
submodules=GatedDeltaNetSubmodules(
in_proj=backend.column_parallel_layer_norm_linear(),
out_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False),
out_proj=backend.row_parallel_linear(),
),
metainfo={"fuse_input_layernorm": True},
)
return attention


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,
) -> 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."

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_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.
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(),
),
)
),
)

# 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_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,
),
metainfo={"fuse_input_layernorm": False},
)

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,
) -> 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
)
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,
)
else:
raise ValueError(
f"Invalid experimental attention variant: {experimental_attention_variant}"
)
52 changes: 31 additions & 21 deletions megatron/core/models/gpt/gpt_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@

from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add
from megatron.core.models.backends import BackendSpecProvider, LocalSpecProvider
from megatron.core.models.gpt.linear_attention_module_specs import (
get_linear_attention_module_spec_for_backend,
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
Expand Down Expand Up @@ -78,7 +78,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,
linear_attention_type: Optional[str] = None,
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 @@ -96,7 +96,8 @@ def get_gpt_layer_with_transformer_engine_spec(
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.
linear_attention_type (str, optional): The type of linear attention. Defaults to None.
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 @@ -133,7 +134,7 @@ def get_gpt_layer_with_transformer_engine_spec(
attention = get_attention_module_spec_for_backend(
backend=backend,
sharded_state_dict_keys_map=sharded_state_dict_keys_map,
linear_attention_type=linear_attention_type,
experimental_attention_variant=experimental_attention_variant,
qk_layernorm=qk_layernorm,
qk_l2_norm=qk_l2_norm,
multi_latent_attention=multi_latent_attention,
Expand Down Expand Up @@ -166,7 +167,7 @@ def get_gpt_layer_local_spec(
moe_grouped_gemm: Optional[bool] = False,
qk_layernorm: Optional[bool] = False,
multi_latent_attention: Optional[bool] = False,
linear_attention_type: Optional[str] = None,
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 @@ -181,7 +182,8 @@ def get_gpt_layer_local_spec(
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.
linear_attention_type (str, optional): The type of linear attention. Defaults to None.
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 @@ -205,15 +207,17 @@ def get_gpt_layer_local_spec(
" and will be removed soon. Please update your code accordingly."
)

if linear_attention_type is not None:
raise NotImplementedError("Linear attention is not supported with local spec yet.")
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,
linear_attention_type=linear_attention_type,
experimental_attention_variant=experimental_attention_variant,
qk_layernorm=qk_layernorm,
qk_l2_norm=qk_l2_norm,
multi_latent_attention=multi_latent_attention,
Expand Down Expand Up @@ -278,7 +282,7 @@ def get_transformer_layer_spec_for_backend(
def get_attention_module_spec_for_backend(
backend: BackendSpecProvider,
sharded_state_dict_keys_map: dict,
linear_attention_type: Optional[str] = None,
experimental_attention_variant: Optional[str] = None,
qk_layernorm: Optional[bool] = False,
qk_l2_norm: Optional[bool] = False,
multi_latent_attention: Optional[bool] = False,
Expand All @@ -288,11 +292,17 @@ def get_attention_module_spec_for_backend(
) -> ModuleSpec:
"""Helper function to get module spec for Attention"""

if linear_attention_type is not None:
return get_linear_attention_module_spec_for_backend(
backend=backend,
linear_attention_type=linear_attention_type,
normalization=normalization,
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.
Expand Down Expand Up @@ -526,21 +536,20 @@ def get_gpt_decoder_layer_specs(
num_experts = None
moe_grouped_gemm = None
if attention_type == "linear_attention":
if config.linear_attention_type is None:
linear_attention_variants = ["gated_delta_net"]
if config.experimental_attention_variant not in linear_attention_variants:
# Skip if there is no linear attention layer in the model.
continue
linear_attention_type = config.linear_attention_type
multi_latent_attention = None
else:
linear_attention_type = None
multi_latent_attention = config.multi_latent_attention

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,
linear_attention_type=linear_attention_type,
experimental_attention_variant=config.experimental_attention_variant,
**get_layer_spec_kwargs,
)

Expand Down Expand Up @@ -583,7 +592,8 @@ def get_gpt_decoder_layer_specs(
f"current linear attention pattern: {config.linear_attention_freq}"
)
elif config.linear_attention_freq is None:
if config.linear_attention_type is None:
linear_attention_variants = ["gated_delta_net"]
if config.experimental_attention_variant not in linear_attention_variants:
linear_attention_pattern = [0] * config.num_layers
else:
linear_attention_pattern = [1] * config.num_layers
Expand Down
27 changes: 0 additions & 27 deletions megatron/core/models/gpt/linear_attention_module_specs.py

This file was deleted.

1 change: 1 addition & 0 deletions megatron/core/transformer/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,7 @@ def __init__(
self.key_hidden_size = self.hidden_size_per_attention_head
self.val_hidden_size = self.hidden_size_per_attention_head

# TODO: This is built twice when using MLA, should be refactored.
self.core_attention = build_module(
submodules.core_attention,
config=self.config,
Expand Down
Loading
Loading