358 lines
12 KiB
Python
358 lines
12 KiB
Python
import itertools
|
|
import unittest
|
|
|
|
import torch
|
|
|
|
from sglang.srt.layers.layernorm import (
|
|
Gemma3RMSNorm,
|
|
GemmaRMSNorm,
|
|
LayerNorm,
|
|
RMSNorm,
|
|
)
|
|
from sglang.test.test_utils import CustomTestCase
|
|
|
|
|
|
class TestRMSNorm(CustomTestCase):
|
|
DTYPES = [torch.half, torch.bfloat16]
|
|
NUM_TOKENS = [7, 83, 4096]
|
|
HIDDEN_SIZES = [768, 769, 770, 771, 5120, 5124, 5125, 5126, 8192, 8199]
|
|
ADD_RESIDUAL = [False, True]
|
|
SEEDS = [0]
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
if not torch.cuda.is_available():
|
|
raise unittest.SkipTest("CUDA is not available")
|
|
torch.set_default_device("cuda")
|
|
|
|
def _run_rms_norm_test(self, num_tokens, hidden_size, add_residual, dtype, seed):
|
|
torch.manual_seed(seed)
|
|
|
|
layer = RMSNorm(hidden_size).to(dtype=dtype)
|
|
layer.weight.data.normal_(mean=1.0, std=0.1)
|
|
scale = 1 / (2 * hidden_size)
|
|
x = torch.randn(num_tokens, hidden_size, 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_rms_norm(self):
|
|
for params in itertools.product(
|
|
self.NUM_TOKENS,
|
|
self.HIDDEN_SIZES,
|
|
self.ADD_RESIDUAL,
|
|
self.DTYPES,
|
|
self.SEEDS,
|
|
):
|
|
with self.subTest(
|
|
num_tokens=params[0],
|
|
hidden_size=params[1],
|
|
add_residual=params[2],
|
|
dtype=params[3],
|
|
seed=params[4],
|
|
):
|
|
self._run_rms_norm_test(*params)
|
|
|
|
|
|
class TestGemmaRMSNorm(CustomTestCase):
|
|
DTYPES = [torch.half, torch.bfloat16]
|
|
NUM_TOKENS = [7, 83, 4096]
|
|
HIDDEN_SIZES = [768, 769, 770, 771, 5120, 5124, 5125, 5126, 8192, 8199]
|
|
ADD_RESIDUAL = [False, True]
|
|
SEEDS = [0]
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
if not torch.cuda.is_available():
|
|
raise unittest.SkipTest("CUDA is not available")
|
|
torch.set_default_device("cuda")
|
|
|
|
def _run_gemma_rms_norm_test(
|
|
self, num_tokens, hidden_size, add_residual, dtype, seed
|
|
):
|
|
torch.manual_seed(seed)
|
|
|
|
layer = GemmaRMSNorm(hidden_size).to(dtype=dtype)
|
|
layer.weight.data.normal_(mean=1.0, std=0.1)
|
|
scale = 1 / (2 * hidden_size)
|
|
x = torch.randn(num_tokens, hidden_size, 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-3, rtol=1e-3))
|
|
self.assertTrue(torch.allclose(out[1], ref_out[1], atol=1e-3, rtol=1e-3))
|
|
else:
|
|
self.assertTrue(torch.allclose(out, ref_out, atol=1e-3, rtol=1e-3))
|
|
|
|
def test_gemma_rms_norm(self):
|
|
for params in itertools.product(
|
|
self.NUM_TOKENS,
|
|
self.HIDDEN_SIZES,
|
|
self.ADD_RESIDUAL,
|
|
self.DTYPES,
|
|
self.SEEDS,
|
|
):
|
|
with self.subTest(
|
|
num_tokens=params[0],
|
|
hidden_size=params[1],
|
|
add_residual=params[2],
|
|
dtype=params[3],
|
|
seed=params[4],
|
|
):
|
|
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]
|
|
NUM_TOKENS = [7, 83, 1024]
|
|
HIDDEN_SIZES = [128, 512, 1536, 5120, 5124, 5125, 5126, 7168]
|
|
USE_AFFINE = [False, True]
|
|
USE_BIAS = [False, True]
|
|
SEEDS = [0]
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
if not torch.cuda.is_available():
|
|
raise unittest.SkipTest("CUDA is not available")
|
|
torch.set_default_device("cuda")
|
|
|
|
def _run_layer_norm_test(
|
|
self, num_tokens, hidden_size, use_affine, use_bias, dtype, seed, param_dtype
|
|
):
|
|
torch.manual_seed(seed)
|
|
|
|
layer = LayerNorm(
|
|
hidden_size, elementwise_affine=use_affine, bias=use_bias, dtype=param_dtype
|
|
)
|
|
if use_affine:
|
|
layer.weight.data.normal_(mean=1.0, std=0.1)
|
|
if use_bias:
|
|
layer.bias.data.normal_(mean=0.0, std=0.1)
|
|
|
|
scale = 1 / (2 * hidden_size)
|
|
x = torch.randn(num_tokens, hidden_size, dtype=dtype) * scale
|
|
|
|
with torch.inference_mode():
|
|
ref_out = layer.forward_native(x)
|
|
out = layer(x)
|
|
|
|
self.assertTrue(torch.allclose(out, ref_out, atol=1e-2, rtol=1e-3))
|
|
|
|
if (
|
|
use_affine
|
|
and use_bias
|
|
and not (dtype == torch.bfloat16 and param_dtype == torch.float32)
|
|
):
|
|
layer.dtype = torch.float32
|
|
layer.weight.data = layer.weight.data.to(torch.float32)
|
|
layer.bias.data = layer.bias.data.to(torch.float32)
|
|
with torch.inference_mode():
|
|
cuda_out = layer(x.to(torch.bfloat16)).to(x.dtype)
|
|
|
|
self.assertTrue(torch.allclose(cuda_out, ref_out, atol=2e-2, rtol=1e-3))
|
|
|
|
def test_layer_norm(self):
|
|
for params in itertools.product(
|
|
self.NUM_TOKENS,
|
|
self.HIDDEN_SIZES,
|
|
self.USE_AFFINE,
|
|
self.USE_BIAS,
|
|
self.DTYPES,
|
|
self.SEEDS,
|
|
self.PARAM_DTYPES,
|
|
):
|
|
with self.subTest(
|
|
num_tokens=params[0],
|
|
hidden_size=params[1],
|
|
use_affine=params[2],
|
|
use_bias=params[3],
|
|
dtype=params[4],
|
|
seed=params[5],
|
|
param_dtype=params[6],
|
|
):
|
|
self._run_layer_norm_test(*params)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main(verbosity=2)
|