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..a4848f6cb2d 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -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 AccuracyCompatibleRMSNorm from megatron.core.transformer.transformer_block import ( TransformerBlockSubmodules, get_num_layers_to_build, @@ -107,7 +108,13 @@ def get_dsa_module_spec_for_backend( # 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 + ( + AccuracyCompatibleRMSNorm + if config.norm_accuracy_compatible + else backend.layer_norm(rms_norm=rms_norm, for_qk=True) + ) + if config.qk_layernorm + else IdentityOp ) attention = ModuleSpec( diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 984840b3a87..63a1aa51b02 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -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: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index dde238635c2..c7f3a212558 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -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() @@ -99,10 +100,15 @@ 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] @@ -110,6 +116,25 @@ def _unfused_absorbed_dsa_fn( 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, @@ -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: @@ -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, @@ -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() @@ -182,6 +209,7 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, + accuracy_compatible=accuracy_compatible, ) @@ -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. @@ -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 diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 03317b65f1c..e25339b5ff1 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -104,6 +104,14 @@ def gating(self, input: torch.Tensor): router_dtype = torch.float32 elif self.config.moe_router_dtype == 'fp64': router_dtype = torch.float64 + if self.config.router_accuracy_compatible: + inp_shape = input.shape + logits = torch.mm( + input.reshape(-1, inp_shape[-1]).float(), self.weight.float().t() + ) + if self.bias is not None: + logits = logits + self.bias.float() + return logits.view(*inp_shape[:-1], -1) logits = router_gating_linear(input, self.weight, self.bias, router_dtype) return logits diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index b20514ce6a4..536c3dfd15e 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -30,7 +30,7 @@ from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module -from megatron.core.transformer.torch_norm import LayerNormBuilder +from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm, LayerNormBuilder from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module @@ -577,7 +577,9 @@ class MultiTokenPredictionLayerSubmodules: def get_mtp_layer_spec( - mtp_model_layer_spec: ModuleSpec, use_transformer_engine: bool + mtp_model_layer_spec: ModuleSpec, + use_transformer_engine: bool, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -587,11 +589,14 @@ def get_mtp_layer_spec( return get_mtp_layer_spec_for_backend( mtp_model_layer_spec, backend=TESpecProvider() if use_transformer_engine else LocalSpecProvider(), + config=config, ) def get_mtp_layer_spec_for_backend( - mtp_model_layer_spec: ModuleSpec, backend: BackendSpecProvider + mtp_model_layer_spec: ModuleSpec, + backend: BackendSpecProvider, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -599,7 +604,11 @@ def get_mtp_layer_spec_for_backend( ModuleSpec: Module specification with modules from the backend. """ column_parallel_linear_impl: type = backend.column_parallel_linear() - layer_norm_impl = backend.layer_norm() + layer_norm_impl = ( + AccuracyCompatibleRMSNorm + if config is not None and config.norm_accuracy_compatible + else backend.layer_norm() + ) mtp_layer_spec = ModuleSpec( module=MultiTokenPredictionLayer, submodules=MultiTokenPredictionLayerSubmodules( diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index 5948ae600f9..c0ddf0763fc 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -24,6 +24,54 @@ def __call__( ) -> LayerNormInterface: ... +class _AccuracyCompatibleRMSNormFunction(torch.autograd.Function): + """RMSNorm core with a stable fp32 backward and canonical zero gradients.""" + + @staticmethod + def forward(ctx, x: torch.Tensor, eps: float) -> torch.Tensor: + variance = x.pow(2).mean(dim=-1, keepdim=True) + inv_rms = torch.rsqrt(variance + eps) + ctx.save_for_backward(x, inv_rms) + return x * inv_rms + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + x, inv_rms = ctx.saved_tensors + dot = (grad_output * x).sum(dim=-1, keepdim=True) + correction_scale = dot * (-0.5) * inv_rms.pow(3) / x.shape[-1] + correction = (correction_scale * x) * 2.0 + grad_input = grad_output * inv_rms + correction + grad_input = torch.where(grad_input == 0, torch.zeros_like(grad_input), grad_input) + return grad_input, None + + +class AccuracyCompatibleRMSNorm(torch.nn.Module, LayerNormInterface): + """RMSNorm with explicit fp32 reduction and one output cast.""" + + def __init__( + self, + normalized_shape: int | None = None, + eps: float = 1e-5, + *, + hidden_size: int | None = None, + config: TransformerConfig | None = None, + **kwargs, + ): + super().__init__() + normalized_shape = hidden_size if normalized_shape is None else normalized_shape + if normalized_shape is None: + raise ValueError("normalized_shape or hidden_size is required") + self.normalized_shape = (normalized_shape,) + self.eps = eps + dtype = config.params_dtype if config is not None else None + self.weight = torch.nn.Parameter(torch.ones(normalized_shape, dtype=dtype)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x_float = x.float() + output = _AccuracyCompatibleRMSNormFunction.apply(x_float, self.eps) + return (output * self.weight.float()).to(x.dtype) + + class WrappedTorchNorm: """ A conditional wrapper to initialize an instance of PyTorch's @@ -56,6 +104,8 @@ def __new__( if config.normalization == "LayerNorm": norm_cls = torch.nn.LayerNorm elif config.normalization == "RMSNorm": + if config.norm_accuracy_compatible: + return AccuracyCompatibleRMSNorm(normalized_shape=hidden_size, eps=eps) assert is_torch_min_version( "2.4.0a0" ), 'Torch RMSNorm requires PyTorch version >= 2.4.0' diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index bbcf413baee..b2fb6da4b05 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -186,6 +186,16 @@ class TransformerConfig(ModelParallelConfig): ) """Epsilon value for any LayerNorm/RMSNorm operations.""" + norm_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}} + ) + """Use explicit fp32 normalization formulas instead of native norm kernels for alignment.""" + + router_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}} + ) + """Use an explicit fp32 router GEMM instead of the fused Transformer Engine path.""" + layernorm_zero_centered_gamma: bool = field( default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} ) @@ -318,6 +328,11 @@ class TransformerConfig(ModelParallelConfig): ``none`` disables fused DSA kernels. Explicit ``tilelang`` or ``cudnn`` enables only that backend. Unsupported DSA layouts continue to use the PyTorch fallback.""" + dsa_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}} + ) + """Use the full-score DSA fallback with explicit softmax backward for alignment.""" + dsa_indexer_rope_interleaved: bool = False """Whether DSA indexer RoPE should use MLA-style interleaving.""" 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..5fca6010a30 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 @@ -65,6 +65,7 @@ def _make_config(**overrides): defaults = dict( num_layers=4, normalization="RMSNorm", + norm_accuracy_compatible=False, qk_layernorm=False, multi_latent_attention=False, qk_l2_norm=False, @@ -369,6 +370,23 @@ def test_qk_layernorm_enabled(self, normalization): assert spec.submodules.q_layernorm is spec.submodules.kv_layernorm backend.layer_norm.assert_any_call(rms_norm=expected_rms, for_qk=True) + def test_accuracy_compatible_qk_rmsnorm(self): + """Verify DSA q/kv norms can use the explicit fp32 RMSNorm path.""" + from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm + + backend = _make_backend() + cfg = _make_config( + multi_latent_attention=True, + qk_l2_norm=False, + qk_layernorm=True, + normalization="RMSNorm", + norm_accuracy_compatible=True, + ) + spec = self._call(cfg=cfg, backend=backend) + + assert spec.submodules.q_layernorm is AccuracyCompatibleRMSNorm + assert spec.submodules.kv_layernorm is AccuracyCompatibleRMSNorm + def test_qk_layernorm_disabled(self): """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" backend = _make_backend() diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py index 642aeeb126f..67ae7530216 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py @@ -27,6 +27,7 @@ DSAttention, DSAttentionSubmodules, FusedDSAIndexerLoss, + _AccuracyCompatibleSoftmax, _run_sparse_attention, _validate_nonpacked_cp_uniform_length, compute_dsa_indexer_loss, @@ -68,6 +69,60 @@ def mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor return x * scale +class TestAccuracyCompatibleDSA: + """Test the opt-in full-score DSA alignment path.""" + + def test_explicit_softmax_backward_matches_formula(self): + logits = torch.randn(2, 3, 5, device="cuda", requires_grad=True) + valid_mask = torch.ones_like(logits, dtype=torch.bool) + valid_mask[..., -1] = False + grad_output = torch.randn_like(logits) + + probabilities = _AccuracyCompatibleSoftmax.apply(logits, valid_mask) + probabilities.backward(grad_output) + expected = probabilities.detach() * ( + grad_output - (grad_output * probabilities.detach()).sum(dim=-1, keepdim=True) + ) + expected = expected.masked_fill(~valid_mask, 0.0) + + assert torch.equal(logits.grad, expected) + assert torch.equal(probabilities[..., -1], torch.zeros_like(probabilities[..., -1])) + + def test_accuracy_compatible_switch_defaults_off(self, monkeypatch): + query = torch.randn(8, 1, 2, 8, device="cuda", dtype=torch.bfloat16) + key = torch.randn_like(query) + value = torch.randn(8, 1, 2, 4, device="cuda", dtype=torch.bfloat16) + indices = torch.arange(8, device="cuda").view(1, 8, 1) + original = unfused_dsa_fn + calls = [] + + def capture(*args, **kwargs): + calls.append(kwargs.get("accuracy_compatible")) + return original(*args, **kwargs) + + monkeypatch.setattr( + "megatron.core.transformer.experimental_attention_variant.dsa.unfused_dsa_fn", + capture, + ) + common = dict( + absorbed_mla=False, + query=query, + key=key, + value=value, + up_v_weight=None, + topk_indices=indices, + softmax_scale=query.size(-1) ** -0.5, + mask=None, + varlen_starts=None, + varlen_ends=None, + key_positions=None, + ) + _run_sparse_attention(config=SimpleNamespace(), **common) + _run_sparse_attention(config=SimpleNamespace(dsa_accuracy_compatible=True), **common) + + assert calls == [False, True] + + class TestDSAIndexShareHelpers: """Test cross-layer top-k sharing helpers.""" diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index 9f33dd01920..da05d34b937 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -62,6 +62,42 @@ def test_constructor(self): num_weights = sum([p.numel() for p in self.router.parameters()]) assert num_weights == 12 * 4, num_weights + @pytest.mark.internal + def test_router_accuracy_compatible_gating(self): + hidden_states = torch.randn( + (3, 1, self.router.config.hidden_size), device="cuda", dtype=torch.bfloat16 + ) + self.router.config.router_accuracy_compatible = True + + logits = self.router.gating(hidden_states) + expected = torch.mm( + hidden_states.reshape(-1, hidden_states.shape[-1]).float(), + self.router.weight.float().t(), + ).view(3, 1, -1) + + assert logits.dtype == torch.float32 + assert torch.equal(logits, expected) + + @pytest.mark.internal + def test_default_router_gating_stays_native(self, monkeypatch): + expected = torch.randn((3, 1, self.router.config.num_moe_experts)) + called = False + + def fake_router_gating_linear(inp, weight, bias, router_dtype): + nonlocal called + called = True + return expected + + monkeypatch.setattr( + "megatron.core.transformer.moe.router.router_gating_linear", + fake_router_gating_linear, + ) + hidden_states = torch.randn((3, 1, self.router.config.hidden_size), dtype=torch.bfloat16) + + assert self.router.config.router_accuracy_compatible is False + assert self.router.gating(hidden_states) is expected + assert called + @pytest.mark.internal @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.parametrize("moe_router_pre_softmax", [(True), (False)]) diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index c3c3944e007..9014d0896fb 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -29,6 +29,7 @@ process_mtp_loss, roll_tensor, ) +from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import get_batch_on_this_cp_rank, is_te_min_version, unwrap_model from megatron.training.argument_utils import gpt_config_from_args, hybrid_config_from_args @@ -82,6 +83,34 @@ def _create_config_and_mtp_block_spec(self, tp, cp, use_te=False): ) return config, mtp_block_spec + def test_accuracy_compatible_norms_override_te_mtp_norms(self): + """Accuracy mode must route all MTP-owned norms through the explicit RMSNorm.""" + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + config = TransformerConfig( + mtp_num_layers=1, + num_layers=1, + hidden_size=64, + num_attention_heads=8, + normalization="RMSNorm", + norm_accuracy_compatible=True, + use_cpu_initialization=True, + ) + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec() + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, + spec=transformer_layer_spec, + use_transformer_engine=True, + ) + mtp_layer_spec = mtp_block_spec.layer_specs[0] + + assert mtp_layer_spec.submodules.enorm is AccuracyCompatibleRMSNorm + assert mtp_layer_spec.submodules.hnorm is AccuracyCompatibleRMSNorm + assert mtp_layer_spec.submodules.layer_norm is AccuracyCompatibleRMSNorm + final_norm = mtp_layer_spec.submodules.layer_norm( + config=config, hidden_size=config.hidden_size, eps=config.layernorm_epsilon + ) + assert isinstance(final_norm, AccuracyCompatibleRMSNorm) + def test_mtp_detach_heads_config(self): """Test that mtp_detach_heads config defaults to False.""" config = TransformerConfig( 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..d3e69a59c6e --- /dev/null +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -0,0 +1,59 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import torch + +from megatron.core.transformer.torch_norm import ( + AccuracyCompatibleRMSNorm, + WrappedTorchNorm, +) +from megatron.core.transformer.transformer_config import TransformerConfig + + +def _config(**overrides): + values = { + "num_layers": 1, + "hidden_size": 64, + "num_attention_heads": 4, + "normalization": "RMSNorm", + } + values.update(overrides) + return TransformerConfig(**values) + + +def test_accuracy_compatible_rmsnorm_matches_explicit_formula(): + config = _config(norm_accuracy_compatible=True) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5).cuda().bfloat16() + x = torch.randn(2, 3, 64, device="cuda", dtype=torch.bfloat16) + + output = norm(x) + x_float = x.float() + expected = ( + x_float + * torch.rsqrt(x_float.pow(2).mean(dim=-1, keepdim=True) + 1e-5) + * norm.weight.float() + ).to(torch.bfloat16) + + assert isinstance(norm, AccuracyCompatibleRMSNorm) + assert torch.equal(output, expected) + + +def test_accuracy_compatible_rmsnorm_canonicalizes_zero_input_gradients(): + config = _config(norm_accuracy_compatible=True) + norm = WrappedTorchNorm(config=config, hidden_size=4, eps=1e-5).cuda().bfloat16() + with torch.no_grad(): + norm.weight.copy_(torch.tensor([-1.0, 1.0, -2.0, 2.0], device="cuda")) + x = torch.tensor( + [[[1.0, -1.0, 2.0, -2.0]]], device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + + norm(x).backward(torch.zeros_like(x)) + + assert torch.equal(x.grad, torch.zeros_like(x.grad)) + assert torch.equal(x.grad.view(torch.uint16), torch.zeros_like(x.grad.view(torch.uint16))) + + +def test_default_rmsnorm_stays_native(): + config = _config(norm_accuracy_compatible=False) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) + + assert isinstance(norm, torch.nn.RMSNorm)