From d6b54d8c3dc7d73d8125587495242cfacd778fd8 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Wed, 5 Aug 2026 15:51:11 +0800 Subject: [PATCH] WIP: align Qwen3.5 local Megatron modules --- ...rimental_attention_variant_module_specs.py | 58 ++++++---- megatron/core/ssm/gated_delta_net.py | 4 +- megatron/core/transformer/torch_norm.py | 55 ++++++++-- ...rimental_attention_variant_module_specs.py | 100 ++++++++++++++++++ .../unit_tests/transformer/test_torch_norm.py | 62 +++++++++++ 5 files changed, 250 insertions(+), 29 deletions(-) create mode 100644 tests/unit_tests/transformer/test_torch_norm.py 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 a76fe6e3a23..02f8a3f6521 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -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 ( @@ -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 @@ -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( @@ -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 diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py index 06eb0763e57..280d1d19f64 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py @@ -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}.", diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index 5948ae600f9..74b45afd9b1 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -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 @@ -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") diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index 0a454b5d7ff..a403b0c8dec 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -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) diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py new file mode 100644 index 00000000000..5535516f44d --- /dev/null +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -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)