Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
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.models.backends import BackendSpecProvider, LocalSpecProvider
from megatron.core.ssm.gated_delta_net import GatedDeltaNet, GatedDeltaNetSubmodules
from megatron.core.transformer.enums import AttnMaskType, LayerType
from megatron.core.transformer.experimental_attention_variant.absorbed_mla import (
Expand Down Expand Up @@ -66,14 +66,16 @@ def get_gated_delta_net_module_spec(
backend = _get_backend_spec_provider(config=config)

rms_norm = config.normalization == "RMSNorm"
fused_in_proj = backend.column_parallel_layer_norm_linear()
fuse_input_layernorm = backend.fuse_layernorm_and_linear() and fused_in_proj is not None
attention = ModuleSpec(
module=GatedDeltaNet,
submodules=GatedDeltaNetSubmodules(
in_proj=backend.column_parallel_layer_norm_linear(),
in_proj=fused_in_proj if fuse_input_layernorm else backend.column_parallel_linear(),
out_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False),
out_proj=backend.row_parallel_linear(),
),
metainfo={"fuse_input_layernorm": True},
metainfo={"fuse_input_layernorm": fuse_input_layernorm},
)
return attention

Expand Down Expand Up @@ -437,11 +439,15 @@ def get_linear_attention_pattern(config: TransformerConfig) -> List[int]:


def _get_backend_spec_provider(config: TransformerConfig) -> BackendSpecProvider:
"""Get backend spec provider for experimental attention variant."""
"""Get backend spec provider for an experimental attention variant."""

if config.transformer_impl == "local":
assert not config.use_kitchen, "Kitchen is not supported with the local transformer implementation."
return LocalSpecProvider()

assert config.transformer_impl == "transformer_engine", (
"Experimental GPT decoder block spec only supports "
"transformer engine implementation for now."
"Experimental GPT decoder block spec supports only the local or transformer_engine "
f"implementations, got {config.transformer_impl!r}."
)
backend: BackendSpecProvider = (
KitchenSpecProvider(
Expand Down Expand Up @@ -471,20 +477,34 @@ 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

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,
qk_l2_norm=config.qk_l2_norm,
use_kitchen=config.use_kitchen,
use_te_activation_func=config.use_te_activation_func,
use_kitchen_attention=config.use_kitchen_attention,
kitchen_attention_backend=config.kitchen_attention_backend,
mla_down_proj_fusion=getattr(config, "mla_down_proj_fusion", False),
from megatron.core.models.gpt.gpt_layer_specs import (
get_gpt_layer_local_spec,
get_gpt_layer_with_transformer_engine_spec,
)

if config.transformer_impl == "local":
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,
normalization=config.normalization,
qk_l2_norm=config.qk_l2_norm,
use_kitchen=False,
)
else:
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,
qk_l2_norm=config.qk_l2_norm,
use_kitchen=config.use_kitchen,
use_te_activation_func=config.use_te_activation_func,
use_kitchen_attention=config.use_kitchen_attention,
kitchen_attention_backend=config.kitchen_attention_backend,
mla_down_proj_fusion=getattr(config, "mla_down_proj_fusion", False),
)
attn_spec = layer_spec.submodules.self_attention
if config.multi_latent_attention:
attn_spec.metainfo["fuse_input_layernorm"] = False
Expand Down
4 changes: 2 additions & 2 deletions megatron/core/ssm/gated_delta_net.py
Original file line number Diff line number Diff line change
Expand Up @@ -641,9 +641,9 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None, tp_gr
if name == "conv1d":
# Add TP sharding for Conv1d
module_sd = module.state_dict(prefix="", keep_vars=True)
tp_sharding_map = {f"weight": 0}
tp_sharding_map = {"weight": 0}
if self.conv_bias:
tp_sharding_map[f"bias"] = 0
tp_sharding_map["bias"] = 0
module_sharded_sd = make_sharded_tensors_for_checkpoint(
module_sd,
f"{prefix}{name}.",
Expand Down
55 changes: 47 additions & 8 deletions megatron/core/transformer/torch_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,42 @@ def __call__(
) -> LayerNormInterface: ...


class TorchRMSNorm(torch.nn.Module):
"""Eager RMSNorm with optional 1-centered (zero-centered gamma) weights.

PyTorch's native ``RMSNorm`` does not support the 1-centered parameterization used by
models such as Qwen3.5. Keep the explicit fp32 reduction and multiplication order here
so the local, non-TransformerEngine path matches that model's reference implementation.
"""

def __init__(self, normalized_shape: int, eps: float, zero_centered_gamma: bool = False):
super().__init__()
self.normalized_shape = (normalized_shape,)
self.eps = eps
self.zero_centered_gamma = zero_centered_gamma
self.weight = torch.nn.Parameter(torch.empty(normalized_shape))
self.reset_parameters()

def reset_parameters(self) -> None:
if self.zero_centered_gamma:
torch.nn.init.zeros_(self.weight)
else:
torch.nn.init.ones_(self.weight)

def forward(self, x: torch.Tensor) -> torch.Tensor:
input_dtype = x.dtype
output = x.float()
output = output * torch.rsqrt(output.pow(2).mean(dim=-1, keepdim=True) + self.eps)
weight = 1.0 + self.weight.float() if self.zero_centered_gamma else self.weight.float()
return (output * weight).to(dtype=input_dtype)

def extra_repr(self) -> str:
return (
f"normalized_shape={self.normalized_shape}, eps={self.eps}, "
f"zero_centered_gamma={self.zero_centered_gamma}"
)


class WrappedTorchNorm:
"""
A conditional wrapper to initialize an instance of PyTorch's
Expand All @@ -41,27 +77,30 @@ def __new__(
zero_centered_gamma: bool = False,
normalization: str = "LayerNorm",
) -> LayerNormInterface:
assert (
not config.layernorm_zero_centered_gamma
), f"zero_centered_gamma not supported by torch LayerNorm"

assert not config.persist_layer_norm, f"persist_layer_norm not supported by torch LayerNorm"
assert not config.persist_layer_norm, "persist_layer_norm not supported by torch LayerNorm"

assert not config.sequence_parallel, f"sequence parallel not supported by torch LayerNorm"
assert not config.sequence_parallel, "sequence parallel not supported by torch LayerNorm"

assert (
not config.memory_efficient_layer_norm
), f"memory_efficient_layer_norm not supported by torch LayerNorm"
), "memory_efficient_layer_norm not supported by torch LayerNorm"

if config.normalization == "LayerNorm":
assert (
not config.layernorm_zero_centered_gamma
), "zero_centered_gamma not supported by torch LayerNorm"
norm_cls = torch.nn.LayerNorm
elif config.normalization == "RMSNorm":
assert is_torch_min_version(
"2.4.0a0"
), 'Torch RMSNorm requires PyTorch version >= 2.4.0'

if config.layernorm_zero_centered_gamma:
return TorchRMSNorm(normalized_shape=hidden_size, eps=eps, zero_centered_gamma=True)
norm_cls = torch.nn.RMSNorm
elif config.normalization == "L2Norm":
assert (
not config.layernorm_zero_centered_gamma
), "zero_centered_gamma not supported by torch L2Norm"
norm_cls = torch.nn.L2Norm
else:
raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -653,3 +653,103 @@ def test_get_transformer_block_with_experimental_attention_variant_spec(

assert isinstance(result, TransformerBlockSubmodules)
assert result.layer_specs == [fake_layer_specs[i] for i in expected_ids]


class TestLocalExperimentalBackend:
def test_gdn_keeps_input_norm_separate_without_fused_linear(self):
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
get_gated_delta_net_module_spec,
)

backend = _make_backend(fuse_layernorm=False)
backend.column_parallel_layer_norm_linear.return_value = None
spec = get_gated_delta_net_module_spec(_make_config(), backend=backend)

assert spec.metainfo == {"fuse_input_layernorm": False}
assert spec.submodules.in_proj is _FakeColumnParallelLinear
assert spec.submodules.out_norm is _FakeLayerNorm
assert spec.submodules.out_proj is _FakeRowParallelLinear

def test_backend_selector_preserves_te_and_adds_local(self):
from megatron.core.extensions.transformer_engine_spec_provider import (
TESpecProvider,
)
from megatron.core.models.backends import LocalSpecProvider
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
_get_backend_spec_provider,
)

assert isinstance(_get_backend_spec_provider(_make_config(transformer_impl="local")), LocalSpecProvider)
assert isinstance(
_get_backend_spec_provider(_make_config(transformer_impl="transformer_engine")), TESpecProvider
)

def test_mixed_gdn_and_full_attention_specs_are_entirely_local(self):
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
_get_backend_spec_provider,
get_transformer_layer_with_experimental_attention_variant_spec,
)
from megatron.core.ssm.gated_delta_net import GatedDeltaNet
from megatron.core.tensor_parallel.layers import (
ColumnParallelLinear,
RowParallelLinear,
)
from megatron.core.transformer.attention import SelfAttention
from megatron.core.transformer.dot_product_attention import DotProductAttention
from megatron.core.transformer.torch_norm import WrappedTorchNorm
from megatron.core.transformer.transformer_config import TransformerConfig

config = TransformerConfig(
num_layers=2,
hidden_size=32,
num_attention_heads=4,
num_query_groups=2,
ffn_hidden_size=64,
num_moe_experts=4,
moe_ffn_hidden_size=16,
moe_router_topk=2,
moe_layer_freq=1,
linear_attention_freq=[1, 0],
experimental_attention_variant="gated_delta_net",
linear_key_head_dim=8,
linear_value_head_dim=8,
linear_num_key_heads=4,
linear_num_value_heads=4,
normalization="RMSNorm",
qk_layernorm=True,
transformer_impl="local",
)
backend = _get_backend_spec_provider(config)
layers = get_transformer_layer_with_experimental_attention_variant_spec(config, backend=backend)

gdn_layer, full_layer = layers
gdn = gdn_layer.submodules.self_attention
full = full_layer.submodules.self_attention
assert gdn.module is GatedDeltaNet
assert gdn.metainfo == {"fuse_input_layernorm": False}
assert gdn_layer.submodules.input_layernorm is WrappedTorchNorm
assert gdn.submodules.in_proj is ColumnParallelLinear
assert gdn.submodules.out_norm is WrappedTorchNorm
assert gdn.submodules.out_proj is RowParallelLinear

assert full.module is SelfAttention
assert full_layer.submodules.input_layernorm is WrappedTorchNorm
assert full.submodules.linear_qkv is ColumnParallelLinear
assert full.submodules.core_attention is DotProductAttention
assert full.submodules.linear_proj is RowParallelLinear
assert full.submodules.q_layernorm is WrappedTorchNorm
assert full.submodules.k_layernorm is WrappedTorchNorm

selected = [
gdn_layer.submodules.input_layernorm,
gdn.submodules.in_proj,
gdn.submodules.out_norm,
gdn.submodules.out_proj,
full_layer.submodules.input_layernorm,
full.submodules.linear_qkv,
full.submodules.core_attention,
full.submodules.linear_proj,
full.submodules.q_layernorm,
full.submodules.k_layernorm,
]
assert all("transformer_engine" not in module.__module__ for module in selected)
62 changes: 62 additions & 0 deletions tests/unit_tests/transformer/test_torch_norm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.

import pytest
import torch

from megatron.core.transformer.torch_norm import TorchRMSNorm, WrappedTorchNorm
from megatron.core.transformer.transformer_config import TransformerConfig


def _config(normalization="RMSNorm", zero_centered_gamma=True):
return TransformerConfig(
num_layers=1,
hidden_size=16,
num_attention_heads=1,
normalization=normalization,
layernorm_zero_centered_gamma=zero_centered_gamma,
persist_layer_norm=False,
sequence_parallel=False,
)


@pytest.mark.parametrize("shape", [(2, 3, 8), (1, 5, 16)])
def test_zero_centered_rmsnorm_matches_explicit_qwen_formula(shape):
torch.manual_seed(1234)
norm = WrappedTorchNorm(_config(), hidden_size=shape[-1], eps=1e-6)
assert isinstance(norm, TorchRMSNorm)
assert list(norm.state_dict()) == ["weight"]
assert torch.count_nonzero(norm.weight) == 0

with torch.no_grad():
norm.weight.copy_(torch.linspace(-0.25, 0.25, shape[-1]))

x = torch.randn(shape, dtype=torch.float32, requires_grad=True)
weight = norm.weight.detach().clone().requires_grad_(True)

actual = norm(x)
reference_fp32 = x.float()
reference_fp32 = reference_fp32 * torch.rsqrt(
reference_fp32.pow(2).mean(dim=-1, keepdim=True) + 1e-6
)
expected = (reference_fp32 * (1.0 + weight.float())).to(dtype=x.dtype)
assert torch.equal(actual, expected)

grad = torch.randn_like(actual)
actual.backward(grad, retain_graph=True)
actual_x_grad = x.grad.detach().clone()
actual_weight_grad = norm.weight.grad.detach().clone()

x.grad = None
expected.backward(grad)
assert torch.equal(actual_x_grad, x.grad)
assert torch.equal(actual_weight_grad, weight.grad)


def test_non_zero_centered_rmsnorm_keeps_native_torch_module():
norm = WrappedTorchNorm(_config(zero_centered_gamma=False), hidden_size=16, eps=1e-6)
assert isinstance(norm, torch.nn.RMSNorm)


def test_zero_centered_layernorm_remains_unsupported():
with pytest.raises(AssertionError, match="zero_centered_gamma not supported by torch LayerNorm"):
WrappedTorchNorm(_config(normalization="LayerNorm"), hidden_size=16, eps=1e-6)