[XPU] Add xpu forward in Gemma3RMSNorm & Add test and benchmark for Gemma3RMSNorm (#36278)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user