diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 626381b65..5bb45333a 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -1317,6 +1317,19 @@ class Gemma3RMSNorm(BaseFusedOp): return gemma_rmsnorm(x, self.weight.data, self.eps) return self.forward_native(x) + def forward_xpu(self, x, residual: Optional[torch.Tensor] = None): + if residual is not None and x.dim() == 2: + # The decoder residual is token-major and contiguous. The fused + # kernel updates both tensors in place: x becomes the normalized + # output and residual becomes x + residual for the next layer. + gemma_fused_add_rmsnorm(x, residual, self.weight.data, self.eps) + return x, residual + # The XPU kernel flattens leading dims internally, so 2D/3D/4D inputs + # can all go through it directly without a Python-side reshape. + elif residual is None and x.dim() in (2, 3, 4): + return gemma_rmsnorm(x, self.weight.data, self.eps) + return self.forward_native(x, residual) + def forward_musa(self, x, residual: Optional[torch.Tensor] = None): # sgl_kernel's gemma norm ops are built for MUSA; follow the CUDA path. return self.forward_cuda(x, residual) diff --git a/test/manual/layers/test_layernorm.py b/test/manual/layers/test_layernorm.py index 299e5dcff..717a6c018 100644 --- a/test/manual/layers/test_layernorm.py +++ b/test/manual/layers/test_layernorm.py @@ -3,7 +3,12 @@ import unittest import torch -from sglang.srt.layers.layernorm import GemmaRMSNorm, LayerNorm, RMSNorm +from sglang.srt.layers.layernorm import ( + Gemma3RMSNorm, + GemmaRMSNorm, + LayerNorm, + RMSNorm, +) from sglang.test.test_utils import CustomTestCase @@ -109,6 +114,173 @@ class TestGemmaRMSNorm(CustomTestCase): self._run_gemma_rms_norm_test(*params) +class TestGemma3RMSNorm(CustomTestCase): + """Covers 2D/3D/4D inputs, including the non-contiguous ("unflatten") + shape that Gemma3's q_norm/k_norm actually feed in: a head slice cut out + of a wider qkv-style tensor via .split()+.unflatten(), so the leading + dims are not flattenable to 2D.""" + + ADD_RESIDUAL = [False, True] + SEEDS = [0] + + # (batch_size, hidden_size, dtype) combos for the 2D case. + SHAPE_DTYPE_2D = [ + (batch_size, hidden_size, torch.float16) + for batch_size in [1, 19] + for hidden_size in [1152, 2560, 3840, 5376] + ] + [ + (19, 1024, torch.bfloat16), + (19, 1024, torch.float32), + (2, 32768, torch.float16), + ] + + BATCH_SIZES_3D = [1, 4] + SEQ_LENS_3D = [1, 74] + # hidden_size=1 exercises the "other dim == 1" shape (excluding the leading + # batch dim) that bypasses the flattenable fast-path check but still + # yields a correct result. + HIDDEN_SIZES_3D = [512, 1024] + DTYPES_3D = [torch.float16] + + NUM_TOKENS_4D = [1, 7] + # num_heads=1 and head_dim=1 exercise the "other dim == 1" shapes + # (excluding the leading batch/token dims) that bypass the flattenable + # fast-path check but still yield a correct result. + NUM_HEADS_4D = [1, 4, 8] + HEAD_DIMS_4D = [64, 128] + DTYPES_4D = [torch.float16] + + @classmethod + def setUpClass(cls): + if not (torch.cuda.is_available() or torch.xpu.is_available()): + raise unittest.SkipTest("Neither CUDA nor XPU is available") + device = "cuda" if torch.cuda.is_available() else "xpu" + torch.set_default_device(device) + + def _run_gemma3_rms_norm_test( + self, shape, add_residual, dtype, seed, non_contiguous=False + ): + torch.manual_seed(seed) + hidden_size = shape[-1] + layer = Gemma3RMSNorm(hidden_size).to(dtype=dtype) + layer.weight.data.normal_(mean=0.0, std=0.1) + scale = 1 / (2 * hidden_size) + + if non_contiguous: + *lead, num_heads, head_dim = shape + total_heads = num_heads + 3 + full = torch.randn(*lead, total_heads * head_dim, dtype=dtype) * scale + x = full[..., : num_heads * head_dim].unflatten(-1, (num_heads, head_dim)) + else: + x = torch.randn(*shape, dtype=dtype) * scale + + residual = torch.randn_like(x) * scale if add_residual else None + + with torch.inference_mode(): + ref_out = layer.forward_native(x, residual) + out = layer(x, residual) + + if add_residual: + self.assertTrue(torch.allclose(out[0], ref_out[0], atol=1e-2, rtol=1e-2)) + self.assertTrue(torch.allclose(out[1], ref_out[1], atol=1e-2, rtol=1e-2)) + else: + self.assertTrue(torch.allclose(out, ref_out, atol=1e-2, rtol=1e-2)) + + def test_gemma3_rms_norm_2d(self): + for (batch_size, hidden_size, dtype), add_residual, seed in itertools.product( + self.SHAPE_DTYPE_2D, self.ADD_RESIDUAL, self.SEEDS + ): + with self.subTest( + batch_size=batch_size, + hidden_size=hidden_size, + add_residual=add_residual, + dtype=dtype, + ): + self._run_gemma3_rms_norm_test( + (batch_size, hidden_size), add_residual, dtype, seed + ) + + def test_gemma3_rms_norm_3d(self): + for ( + batch_size, + seq_len, + hidden_size, + dtype, + add_residual, + seed, + ) in itertools.product( + self.BATCH_SIZES_3D, + self.SEQ_LENS_3D, + self.HIDDEN_SIZES_3D, + self.DTYPES_3D, + self.ADD_RESIDUAL, + self.SEEDS, + ): + with self.subTest( + batch_size=batch_size, + seq_len=seq_len, + hidden_size=hidden_size, + add_residual=add_residual, + dtype=dtype, + ): + self._run_gemma3_rms_norm_test( + (batch_size, seq_len, hidden_size), add_residual, dtype, seed + ) + + def test_gemma3_rms_norm_4d(self): + for ( + num_tokens, + num_heads, + head_dim, + dtype, + add_residual, + seed, + ) in itertools.product( + self.NUM_TOKENS_4D, + self.NUM_HEADS_4D, + self.HEAD_DIMS_4D, + self.DTYPES_4D, + self.ADD_RESIDUAL, + self.SEEDS, + ): + with self.subTest( + num_tokens=num_tokens, + num_heads=num_heads, + head_dim=head_dim, + add_residual=add_residual, + dtype=dtype, + ): + self._run_gemma3_rms_norm_test( + (1, num_tokens, num_heads, head_dim), add_residual, dtype, seed + ) + + def test_gemma3_rms_norm_3d_unflatten(self): + for head_dim, add_residual, dtype, seed in itertools.product( + [64, 128], self.ADD_RESIDUAL, self.DTYPES_3D, self.SEEDS + ): + with self.subTest( + head_dim=head_dim, add_residual=add_residual, dtype=dtype + ): + self._run_gemma3_rms_norm_test( + (19, 4, head_dim), add_residual, dtype, seed, non_contiguous=True + ) + + def test_gemma3_rms_norm_4d_unflatten(self): + for head_dim, add_residual, dtype, seed in itertools.product( + [64, 128], self.ADD_RESIDUAL, self.DTYPES_4D, self.SEEDS + ): + with self.subTest( + head_dim=head_dim, add_residual=add_residual, dtype=dtype + ): + self._run_gemma3_rms_norm_test( + (1, 19, 4, head_dim), + add_residual, + dtype, + seed, + non_contiguous=True, + ) + + class TestLayerNorm(CustomTestCase): DTYPES = [torch.half, torch.bfloat16] PARAM_DTYPES = [torch.bfloat16, torch.float32]