[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 gemma_rmsnorm(x, self.weight.data, self.eps)
|
||||||
return self.forward_native(x)
|
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):
|
def forward_musa(self, x, residual: Optional[torch.Tensor] = None):
|
||||||
# sgl_kernel's gemma norm ops are built for MUSA; follow the CUDA path.
|
# sgl_kernel's gemma norm ops are built for MUSA; follow the CUDA path.
|
||||||
return self.forward_cuda(x, residual)
|
return self.forward_cuda(x, residual)
|
||||||
|
|||||||
@@ -3,7 +3,12 @@ import unittest
|
|||||||
|
|
||||||
import torch
|
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
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
@@ -109,6 +114,173 @@ class TestGemmaRMSNorm(CustomTestCase):
|
|||||||
self._run_gemma_rms_norm_test(*params)
|
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):
|
class TestLayerNorm(CustomTestCase):
|
||||||
DTYPES = [torch.half, torch.bfloat16]
|
DTYPES = [torch.half, torch.bfloat16]
|
||||||
PARAM_DTYPES = [torch.bfloat16, torch.float32]
|
PARAM_DTYPES = [torch.bfloat16, torch.float32]
|
||||||
|
|||||||
Reference in New Issue
Block a user