[XPU] Add xpu forward in Gemma3RMSNorm & Add test and benchmark for Gemma3RMSNorm (#36278)

This commit is contained in:
Miaomiao Jiang
2026-09-11 10:32:25 +08:00
committed by GitHub
parent 7e3ae15f73
commit 690428b470
2 changed files with 186 additions and 1 deletions
+13
View File
@@ -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)
+173 -1
View File
@@ -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]