Skip to content
Draft
10 changes: 10 additions & 0 deletions megatron/core/models/backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
TEColumnParallelGroupedLinear,
TERowParallelGroupedLinear,
)
from megatron.core.post_training.modelopt.layers import Linear
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.transformer.dot_product_attention import DotProductAttention
from megatron.core.transformer.mlp import MLPSubmodules, TEActivationFunctionBuilder
Expand Down Expand Up @@ -99,6 +100,15 @@ def activation_func(self) -> TEActivationFunctionBuilder | None:
class LocalSpecProvider(BackendSpecProvider):
"""A protocol for providing Local submodules used in Spec building."""

def linear(self) -> type:
"""TP-replicated local Linear (modelopt Linear, not TELinear).

DSA indexer / MLA down-projections call backend.linear(). TESpecProvider
still returns TELinear; this method is the TE-off counterpart so a
LocalSpecProvider DSA spec does not re-enter Transformer Engine.
"""
return Linear

def column_parallel_linear(self) -> type:
"""Which column parallel linear module the backend uses"""
return ColumnParallelLinear
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
)
from megatron.core.transformer.identity_op import IdentityOp
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.torch_norm import WrappedTorchNorm
from megatron.core.transformer.transformer_block import (
TransformerBlockSubmodules,
get_num_layers_to_build,
Expand Down Expand Up @@ -57,6 +58,15 @@
##########


def _get_standalone_norm(
config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False
):
rms_norm = config.normalization == "RMSNorm"
if rms_norm and config.norm_accuracy_compatible:
return WrappedTorchNorm
return backend.layer_norm(rms_norm=rms_norm, for_qk=for_qk)


def get_gated_delta_net_module_spec(
config: TransformerConfig, backend: BackendSpecProvider = None
) -> ModuleSpec:
Expand All @@ -65,12 +75,11 @@ def get_gated_delta_net_module_spec(
if backend is None:
backend = _get_backend_spec_provider(config=config)

rms_norm = config.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_norm=_get_standalone_norm(config, backend),
out_proj=backend.row_parallel_linear(),
),
metainfo={"fuse_input_layernorm": True},
Expand All @@ -82,7 +91,9 @@ def get_dsa_module_spec_for_backend(
config: TransformerConfig, backend: BackendSpecProvider = None
) -> ModuleSpec:
"""Helper function to get module spec for Sparse Attention."""
assert config.multi_latent_attention, "Currently only MLA supports sparse attention."
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."

# Because TransformerEngine does not support sparse attention yet, we use local
Expand All @@ -102,12 +113,12 @@ def get_dsa_module_spec_for_backend(
),
)

# Adjust for RMS norm.
rms_norm = config.normalization == "RMSNorm"
# DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm
# with linear projection and have to use unfused qk layernorm.
qk_norm = (
backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp
_get_standalone_norm(config, backend, for_qk=True)
if config.qk_layernorm
else IdentityOp
)

attention = ModuleSpec(
Expand Down Expand Up @@ -203,7 +214,9 @@ def get_transformer_layer_with_experimental_attention_variant_spec(
experimental_attention_spec = None

if 0 in experimental_attention_pattern:
standard_attention_spec = _get_self_attention_module_spec(config=config, backend=backend)
standard_attention_spec = _get_self_attention_module_spec(
config=config, backend=backend
)
else:
standard_attention_spec = None

Expand All @@ -228,15 +241,18 @@ def get_transformer_layer_with_experimental_attention_variant_spec(
dense_mlp_layer_spec, fuse_layernorm_pre_dense = None, False

# 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
mlp = (
moe_layer_spec
if moe_layer_pattern[layer_number] == 1
else dense_mlp_layer_spec
)
fuse_pre_mlp_layernorm = (
fuse_layernorm_pre_moe
if moe_layer_pattern[layer_number] == 1
Expand All @@ -245,12 +261,12 @@ def get_transformer_layer_with_experimental_attention_variant_spec(
input_layernorm = (
IdentityOp
if attention.metainfo["fuse_input_layernorm"]
else backend.layer_norm(rms_norm=rms_norm, for_qk=False)
else _get_standalone_norm(config, backend)
)
pre_mlp_layernorm = (
IdentityOp
if fuse_pre_mlp_layernorm
else backend.layer_norm(rms_norm=rms_norm, for_qk=False)
else _get_standalone_norm(config, backend)
)

layer_specs.append(
Expand All @@ -271,7 +287,9 @@ def get_transformer_layer_with_experimental_attention_variant_spec(


def get_transformer_block_with_experimental_attention_variant_spec(
config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None
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).

Expand Down Expand Up @@ -309,17 +327,20 @@ def get_transformer_block_with_experimental_attention_variant_spec(
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)
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)

_validate_dsa_index_share_pipeline_split(config, local_layer_ids)
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)
layer_specs=layer_specs, layer_norm=_get_standalone_norm(config, backend)
)

return gpt_decoder_block_spec
Expand All @@ -336,7 +357,9 @@ def is_linear_attention_variant(experimental_attention_variant: Optional[str]) -
return experimental_attention_variant in linear_attention_variants


def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_layer_ids) -> None:
def _validate_dsa_index_share_pipeline_split(
config: TransformerConfig, local_layer_ids
) -> None:
"""Ensure DSA top-k sharing does not require top-k indices from another PP stage."""
if (
config.experimental_attention_variant != "dsa"
Expand All @@ -351,12 +374,16 @@ def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_la
for position, layer_id in enumerate(local_layer_ids):
layer_number = layer_id + 1
if not is_dsa_skip_topk_layer(
layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq
layer_number,
config.dsa_indexer_skip_topk_offset,
config.dsa_indexer_topk_freq,
):
continue

source_layer_number = source_dsa_compute_layer(
layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq
layer_number,
config.dsa_indexer_skip_topk_offset,
config.dsa_indexer_topk_freq,
)
source_layer_id = source_layer_number - 1
if (
Expand All @@ -383,7 +410,8 @@ def get_moe_layer_pattern(config: TransformerConfig) -> List[int]:
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)
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
Expand Down Expand Up @@ -471,7 +499,9 @@ def _get_self_attention_module_spec(
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
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,
Expand Down Expand Up @@ -532,7 +562,9 @@ def _get_moe_module_spec(
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
from megatron.core.models.gpt.moe_module_specs import (
get_moe_module_spec_for_backend,
)

return (
get_moe_module_spec_for_backend(
Expand Down
2 changes: 1 addition & 1 deletion megatron/core/models/gpt/gpt_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -773,7 +773,7 @@ def get_gpt_mtp_block_spec_for_backend(
raise ValueError(f"Invalid spec: {spec}")

mtp_layer_spec = get_mtp_layer_spec_for_backend(
mtp_model_layer_spec=transformer_layer_spec, backend=backend
mtp_model_layer_spec=transformer_layer_spec, backend=backend, config=config
)
mtp_num_layers = config.mtp_num_layers if config.mtp_num_layers else 0
if config.mtp_use_repeated_layer:
Expand Down
3 changes: 3 additions & 0 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -556,6 +556,9 @@ def _get_megatron_optimizer_based_on_param_groups(
# on source of optimizer (Torch or TE/Apex)
if USING_PYTORCH_OPTIMIZER:
adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam
elif config.native_unfused_adamw:
adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam
kwargs.update({"foreach": False, "fused": False})
else:
kwargs["adam_w_mode"] = config.decoupled_weight_decay
adam_cls = Adam
Expand Down
3 changes: 3 additions & 0 deletions megatron/core/optimizer/optimizer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,9 @@ class OptimizerConfig:
adam_eps: float = 1e-08
"""Term added to the denominator to improve numerical stability in Adam optimizer."""

native_unfused_adamw: bool = False
"""Use torch.optim.AdamW with foreach=False and fused=False instead of TE/Apex Adam."""

decoupled_weight_decay: bool = True
"""If true, decouples weight decay from the gradient update, equivalent to AdamW. If false,
original Adam update rule will be used. Defaults to True.
Expand Down
60 changes: 55 additions & 5 deletions megatron/core/transformer/experimental_attention_variant/dsa.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ def _unfused_absorbed_dsa_fn(
varlen_starts: Optional[torch.Tensor] = None,
varlen_ends: Optional[torch.Tensor] = None,
key_positions: Optional[torch.Tensor] = None,
accuracy_compatible: bool = False,
) -> torch.Tensor:
"""Unfused absorbed-MLA attention: output stays [sq, b, np, v_channels]."""
sq, b, np, hn = query.size()
Expand Down Expand Up @@ -99,17 +100,41 @@ def _unfused_absorbed_dsa_fn(
)

attention_scores = attention_scores + index_mask.unsqueeze(1)
valid_index_mask = torch.isfinite(index_mask)
attention_scores = dsa_masking.masked_softmax(
attention_scores.float(), valid_index_mask.unsqueeze(1).expand(b, np, sq, skv), dim=-1
)
valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv)
if accuracy_compatible:
attention_scores = _AccuracyCompatibleSoftmax.apply(
attention_scores.float(), valid_index_mask
)
else:
attention_scores = dsa_masking.masked_softmax(
attention_scores.float(), valid_index_mask, dim=-1
)

# Latent value is the first v_channels slice of absorbed key cache.
value = key[..., :v_channels].permute(1, 2, 0, 3) # [b,1,skv,v]
output = torch.matmul(attention_scores.to(value.dtype), value) # [b,np,sq,v]
return output.permute(2, 0, 1, 3).contiguous()


class _AccuracyCompatibleSoftmax(torch.autograd.Function):
"""Masked softmax with an explicit backward formula for DSA alignment."""

@staticmethod
def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor:
probabilities = torch.softmax(logits.masked_fill(~valid_mask, float("-inf")), dim=-1)
probabilities = probabilities.masked_fill(~valid_mask, 0.0)
ctx.save_for_backward(probabilities, valid_mask)
return probabilities

@staticmethod
def backward(ctx, grad_output: torch.Tensor):
probabilities, valid_mask = ctx.saved_tensors
grad_logits = probabilities * (
grad_output - (grad_output * probabilities).sum(dim=-1, keepdim=True)
)
return grad_logits.masked_fill(~valid_mask, 0.0), None


def _run_sparse_attention(
*,
absorbed_mla: bool,
Expand All @@ -127,6 +152,7 @@ def _run_sparse_attention(
topk_length: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Run sparse attention for absorbed and non-absorbed MLA paths."""
accuracy_compatible = bool(getattr(config, "dsa_accuracy_compatible", False))
if absorbed_mla:
latent_v_channels = int(getattr(config, "kv_lora_rank", 0) or 0)
if latent_v_channels <= 0:
Expand All @@ -143,7 +169,7 @@ def _run_sparse_attention(
"Received absorbed layout with explicit value tensor."
)
output = None
if dsa_kernels.use_fused_dsa_kernels(config):
if not accuracy_compatible and dsa_kernels.use_fused_dsa_kernels(config):
output = dsa_kernels.run_fused_absorbed_sparse_attention(
config,
query,
Expand All @@ -166,6 +192,7 @@ def _run_sparse_attention(
varlen_starts=varlen_starts,
varlen_ends=varlen_ends,
key_positions=key_positions,
accuracy_compatible=accuracy_compatible,
)
assert output is not None
output = torch.einsum("sbhc,hdc->sbhd", output, up_v_weight).contiguous()
Expand All @@ -182,6 +209,7 @@ def _run_sparse_attention(
varlen_starts=varlen_starts,
varlen_ends=varlen_ends,
key_positions=key_positions,
accuracy_compatible=accuracy_compatible,
)


Expand Down Expand Up @@ -1411,6 +1439,7 @@ def unfused_dsa_fn(
varlen_starts: Optional[torch.Tensor] = None,
varlen_ends: Optional[torch.Tensor] = None,
key_positions: Optional[torch.Tensor] = None,
accuracy_compatible: bool = False,
):
"""
Unfused sparse attention implementation.
Expand Down Expand Up @@ -1457,6 +1486,27 @@ def unfused_dsa_fn(
device=query.device,
)

if accuracy_compatible:
index_mask = torch.full((b, sq, skv), float("-inf"), device=query.device)
dsa_masking.scatter_topk_into_index_mask(index_mask, topk_indices)
index_mask = dsa_masking.apply_sparse_validity_to_index_mask(
index_mask,
row_mask=row_mask,
varlen_starts=varlen_starts,
varlen_ends=varlen_ends,
key_positions=key_positions,
)
valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv)
attention_scores = (
torch.matmul(query_b.float(), key_b.float().transpose(-1, -2)) * softmax_scale
)
attention_probs = _AccuracyCompatibleSoftmax.apply(
attention_scores + index_mask.unsqueeze(1), valid_index_mask
)
output = torch.matmul(attention_probs.to(value_b.dtype), value_b)
output = output.permute(2, 0, 1, 3).contiguous().view(sq, b, np * hnv)
return output.squeeze(1) if query_was_thd else output

seq_chunk_size = 512
head_chunk_size = 16
topk_chunk_size = 1024
Expand Down
Loading