[CPU] add fused_qk_gemma_norm and refactor norm kernel implementation (#30216)

This commit is contained in:
Ma Mingfei
2026-07-07 08:52:59 +08:00
committed by GitHub
parent 6c1fb8a937
commit 30fb0dd851
6 changed files with 983 additions and 1131 deletions
+220 -258
View File
@@ -1,26 +1,29 @@
import itertools
import unittest
import sys
from typing import Optional, Tuple, Union
import pytest
import torch
from utils import make_non_contiguous, parametrize, precision
from utils import make_non_contiguous, precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
torch.manual_seed(1234)
DTYPES = [torch.float16, torch.bfloat16]
DTYPE_IDS = ["float16", "bfloat16"]
eps = 1e-6
class TestNorm(CustomTestCase):
class TestNorm:
def _forward_native(
self,
x: torch.Tensor,
weight: torch.Tensor,
variance_epsilon: float = 1e-6,
variance_epsilon: float = eps,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
orig_dtype = x.dtype
@@ -41,7 +44,7 @@ class TestNorm(CustomTestCase):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps)
def _gemma3_rmsnorm_native(
self, x: torch.Tensor, weight: torch.Tensor, variance_epsilon: float = 1e-6
self, x: torch.Tensor, weight: torch.Tensor, variance_epsilon: float = eps
):
output = self._norm(x.float(), variance_epsilon)
output = output * (1.0 + weight.float())
@@ -51,7 +54,7 @@ class TestNorm(CustomTestCase):
self,
x: torch.Tensor,
weight: torch.Tensor,
variance_epsilon: float = 1e-6,
variance_epsilon: float = eps,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
orig_dtype = x.dtype
@@ -66,144 +69,91 @@ class TestNorm(CustomTestCase):
x = x.to(orig_dtype)
return x if residual is None else (x, residual)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_norm(self, m, n, dtype):
@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS)
@pytest.mark.parametrize("hidden_size", [2048, 512])
@pytest.mark.parametrize("batch_size", [32, 121])
def test_l2norm(self, batch_size, hidden_size, dtype):
x = torch.randn([m, n], dtype=dtype)
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
out = torch.ops.sgl_kernel.rmsnorm_cpu(x, weight, variance_epsilon)
ref_out = self._forward_native(x, weight, variance_epsilon)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
ref_x = x.clone()
residual = torch.randn([m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
torch.ops.sgl_kernel.fused_add_rmsnorm_cpu(
x, residual, weight, variance_epsilon
)
ref_x, ref_residual = self._forward_native(
ref_x, weight, variance_epsilon, ref_residual
)
torch.testing.assert_close(x, ref_x, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
@parametrize(
l=[1, 2],
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_norm_3d(self, l, m, n, dtype):
x = torch.randn([l, m, n], dtype=dtype)
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
out = torch.ops.sgl_kernel.rmsnorm_cpu(x, weight, variance_epsilon)
ref_out = self._forward_native(x, weight, variance_epsilon)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
ref_x = x.clone()
residual = torch.randn([l, m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
torch.ops.sgl_kernel.fused_add_rmsnorm_cpu(
x, residual, weight, variance_epsilon
)
ref_x, ref_residual = self._forward_native(
ref_x, weight, variance_epsilon, ref_residual
)
torch.testing.assert_close(x, ref_x, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_l2norm(self, m, n, dtype):
x = torch.randn([m, n], dtype=dtype)
hidden_size = x.size(-1)
x = torch.randn([batch_size, hidden_size], dtype=dtype)
fake_ones_weight = torch.ones(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
out = torch.ops.sgl_kernel.l2norm_cpu(x, variance_epsilon)
ref_out = self._forward_native(x, fake_ones_weight, variance_epsilon)
out = torch.ops.sgl_kernel.l2norm_cpu(x, eps)
ref_out = self._forward_native(x, fake_ones_weight, eps)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_gemma_rmsnorm(self, m, n, dtype):
@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS)
@pytest.mark.parametrize("hidden_size", [2048, 512])
@pytest.mark.parametrize("batch_size", [32, 121])
@pytest.mark.parametrize("seq_len", [None, 2], ids=["2d", "3d"])
def test_rmsnorm(self, seq_len, batch_size, hidden_size, dtype):
x = torch.randn([m, n], dtype=dtype)
if seq_len is None:
x = torch.randn([batch_size, hidden_size], dtype=dtype)
else:
x = torch.randn([batch_size, seq_len, hidden_size], dtype=dtype)
x = make_non_contiguous(x)
hidden_size = x.size(-1)
residual = torch.randn(x.shape, dtype=dtype)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
out = torch.ops.sgl_kernel.gemma_rmsnorm_cpu(x, weight, variance_epsilon)
ref_out = self._gemma_rmsnorm_native(x, weight, variance_epsilon)
out = torch.ops.sgl_kernel.rmsnorm_cpu(x, weight, eps)
ref_out = self._forward_native(x, weight, eps)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
ref_x = x.clone()
residual = torch.randn([m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu(
x, residual, weight, variance_epsilon
)
torch.ops.sgl_kernel.fused_add_rmsnorm_cpu(x, residual, weight, eps)
ref_x, ref_residual = self._forward_native(ref_x, weight, eps, ref_residual)
torch.testing.assert_close(x, ref_x, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize("hidden_size", [2048, 256, 33])
@pytest.mark.parametrize("batch_size", [32, 121])
def test_gemma_rmsnorm(self, batch_size, hidden_size, dtype):
x = torch.randn([batch_size, hidden_size], dtype=dtype)
x = make_non_contiguous(x)
weight = torch.randn(hidden_size, dtype=dtype)
out = torch.ops.sgl_kernel.gemma_rmsnorm_cpu(x, weight, eps)
ref_out = self._gemma_rmsnorm_native(x, weight, eps)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
ref_x = x.clone()
residual = torch.randn([batch_size, hidden_size], dtype=dtype)
ref_residual = residual.clone()
torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu(x, residual, weight, eps)
ref_x, ref_residual = self._gemma_rmsnorm_native(
ref_x, weight, variance_epsilon, ref_residual
ref_x, weight, eps, ref_residual
)
torch.testing.assert_close(x, ref_x, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_gemma3_rmsnorm(self, m, n, dtype):
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize("hidden_size", [128, 256])
@pytest.mark.parametrize("batch_size", [32, 121])
def test_gemma3_rmsnorm(self, batch_size, hidden_size, dtype):
x_list = [
torch.randn([m, n], dtype=dtype),
torch.randn([1, m, 2, n], dtype=dtype),
torch.randn([batch_size, hidden_size], dtype=dtype),
torch.randn([batch_size, 16, 2, hidden_size], dtype=dtype),
]
for x in x_list:
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
out = torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, weight, variance_epsilon)
ref_out = self._gemma3_rmsnorm_native(x, weight, variance_epsilon)
out = torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, weight, eps)
ref_out = self._gemma3_rmsnorm_native(x, weight, eps)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
@@ -212,7 +162,7 @@ class TestNorm(CustomTestCase):
self,
x: torch.Tensor,
weight: torch.Tensor,
variance_epsilon: float = 1e-6,
variance_epsilon: float = eps,
scale_shift: float = 0.0,
with_scale: bool = True,
):
@@ -221,53 +171,41 @@ class TestNorm(CustomTestCase):
output = output * (weight.float() + scale_shift)
return output.type_as(x)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
)
def test_gemma4_rmsnorm(self, m, n, dtype):
for scale_shift, with_scale in [
(0.0, True),
(1.0, True),
(0.0, False),
(1.0, False),
]:
x_list = [
torch.randn([m, n], dtype=dtype),
torch.randn([4, m, n], dtype=dtype),
]
# Add non-block-contiguous 3D input
base = torch.randn([4, 2 * m, n], dtype=dtype)
x_list.append(base[:, :m, :])
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize("hidden_size", [128, 2048])
@pytest.mark.parametrize("batch_size", [32, 121])
@pytest.mark.parametrize("scale_shift", [0.0, 1.0], ids=["shift0.0", "shift1.0"])
@pytest.mark.parametrize("with_scale", [True, False], ids=["scale", "no-scale"])
def test_gemma4_rmsnorm(
self, batch_size, hidden_size, dtype, scale_shift, with_scale
):
x_list = [
torch.randn([batch_size, hidden_size], dtype=dtype),
torch.randn([batch_size, 4, hidden_size], dtype=dtype),
]
for x in x_list:
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
for x in x_list:
x = make_non_contiguous(x)
weight = torch.randn(hidden_size, dtype=dtype)
out = torch.ops.sgl_kernel.gemma4_rmsnorm_cpu(
x, weight, variance_epsilon, scale_shift, with_scale
)
ref_out = self._gemma4_rmsnorm_native(
x, weight, variance_epsilon, scale_shift, with_scale
)
out = torch.ops.sgl_kernel.gemma4_rmsnorm_cpu(
x, weight, eps, scale_shift, with_scale
)
ref_out = self._gemma4_rmsnorm_native(
x, weight, eps, scale_shift, with_scale
)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
class TestFusedRMSNormGated(CustomTestCase):
M = [4096, 1024]
N = [4096, 4096 + 13]
dtype = [torch.float16, torch.bfloat16]
class TestFusedRMSNormGated:
def _forward_native(
self,
hidden_states: torch.Tensor,
weight: torch.Tensor,
variance_epsilon: float = 1e-6,
variance_epsilon: float = eps,
gate: Optional[torch.Tensor] = None,
) -> torch.Tensor:
input_dtype = hidden_states.dtype
@@ -280,37 +218,29 @@ class TestFusedRMSNormGated(CustomTestCase):
return hidden_states.to(input_dtype)
def _norm_test(self, m, n, dtype):
x = torch.randn([m, n], dtype=dtype)
@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS)
@pytest.mark.parametrize("hidden_size", [64, 1024 + 13])
@pytest.mark.parametrize("batch_size", [32, 121])
def test_fused_rmsnorm_gated(self, batch_size, hidden_size, dtype):
x = torch.randn([batch_size, hidden_size], dtype=dtype)
x = make_non_contiguous(x)
batch_size = x.size(0)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
gate = torch.randn([batch_size, hidden_size], dtype=dtype)
out = torch.ops.sgl_kernel.fused_rmsnorm_gated_cpu(
x, weight, gate, variance_epsilon
)
ref_out = self._forward_native(x, weight, variance_epsilon, gate)
out = torch.ops.sgl_kernel.fused_rmsnorm_gated_cpu(x, weight, gate, eps)
ref_out = self._forward_native(x, weight, eps, gate)
atol = rtol = precision[ref_out.dtype] * 2
atol = rtol = precision[ref_out.dtype]
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
def test_norm(self):
for params in itertools.product(self.M, self.N, self.dtype):
with self.subTest(m=params[0], n=params[1], dtype=params[2]):
self._norm_test(*params)
class TestLayerNorm(CustomTestCase):
class TestLayerNorm:
def _forward_native(
self,
x: torch.Tensor,
weight: torch.Tensor,
variance_epsilon: float,
variance_epsilon: float = eps,
residual: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
@@ -328,109 +258,141 @@ class TestLayerNorm(CustomTestCase):
x = x.to(orig_dtype)
return x if residual is None else (x, residual)
@parametrize(
m=[4096, 1024],
n=[4096, 4109],
dtype=[torch.float16, torch.bfloat16],
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize("batch_size", [32, 121])
@pytest.mark.parametrize("hidden_size", [128, 4096, 533])
@pytest.mark.parametrize("has_bias", [False, True], ids=["no-bias", "bias"])
def test_layernorm(
self,
batch_size: int,
hidden_size: int,
has_bias: bool,
dtype: torch.dtype,
) -> None:
x_list = [
torch.randn([batch_size, hidden_size], dtype=dtype),
torch.randn([batch_size, 3, hidden_size], dtype=dtype),
]
for x in x_list:
x = make_non_contiguous(x)
weight = torch.randn(hidden_size, dtype=dtype)
bias = torch.randn(hidden_size, dtype=dtype) if has_bias else None
ln_out = torch.ops.sgl_kernel.layernorm_cpu(x, weight, bias, eps)
ref_ln_out = self._forward_native(x, weight, eps, residual=None, bias=bias)
atol = rtol = precision[ref_ln_out.dtype]
torch.testing.assert_close(ln_out, ref_ln_out, atol=atol, rtol=rtol)
residual = torch.randn(x.shape, dtype=dtype)
ref_residual = residual.clone()
add_ln_out = torch.ops.sgl_kernel.fused_add_layernorm_cpu(
x, residual, weight, bias, eps
)
ref_add_ln_out, ref_residual = self._forward_native(
x, weight, eps, residual=ref_residual, bias=bias
)
torch.testing.assert_close(add_ln_out, ref_add_ln_out, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
class TestFusedQKGemmaRMSNorm:
def _gemma_rmsnorm_per_head_native(
self,
x: torch.Tensor,
weight: torch.Tensor,
head_dim: int,
variance_epsilon: float = eps,
) -> torch.Tensor:
orig_dtype = x.dtype
x_f = x.to(torch.float32).reshape(-1, head_dim)
variance = x_f.pow(2).mean(dim=-1, keepdim=True)
x_f = x_f * torch.rsqrt(variance + variance_epsilon)
x_f = x_f * (1.0 + weight.to(torch.float32))
return x_f.to(orig_dtype).reshape_as(x)
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize(
"batch_size,num_head,num_head_kv,head_dim",
[
(8, 4, 2, 128),
(17, 8, 2, 64),
(5, 3, 1, 96),
],
)
def test_norm_input_2d(self, m: int, n: int, dtype: torch.dtype) -> None:
x = torch.randn([m, n], dtype=dtype)
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
bias = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
def test_fused_qk_gemma_rmsnorm(
self, batch_size: int, num_head: int, num_head_kv: int, head_dim: int, dtype
):
q = torch.randn([batch_size, num_head * head_dim], dtype=dtype)
k = torch.randn([batch_size, num_head_kv * head_dim], dtype=dtype)
ln_out = torch.ops.sgl_kernel.layernorm_cpu(x, weight, None, variance_epsilon)
ref_ln_out = self._forward_native(x, weight, variance_epsilon)
# Keep last dim contiguous but make base storage non-contiguous to stress stride handling.
q = make_non_contiguous(q)
k = make_non_contiguous(k)
atol = rtol = precision[ref_ln_out.dtype]
torch.testing.assert_close(ln_out, ref_ln_out, atol=atol, rtol=rtol)
q_weight = torch.randn(head_dim, dtype=dtype)
k_weight = torch.randn(head_dim, dtype=dtype)
ln_out = torch.ops.sgl_kernel.layernorm_cpu(x, weight, bias, variance_epsilon)
ref_ln_out = self._forward_native(
x, weight, variance_epsilon, residual=None, bias=bias
)
torch.testing.assert_close(ln_out, ref_ln_out, atol=atol, rtol=rtol)
residual = torch.randn([m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
add_ln_out = torch.ops.sgl_kernel.fused_add_layernorm_cpu(
x, residual, weight, None, variance_epsilon
)
ref_add_ln_out, ref_residual = self._forward_native(
x, weight, variance_epsilon, residual=ref_residual
q_out, k_out = torch.ops.sgl_kernel.fused_qk_gemma_rmsnorm_cpu(
q, k, q_weight, k_weight, eps, head_dim
)
torch.testing.assert_close(add_ln_out, ref_add_ln_out, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
ref_q_out = self._gemma_rmsnorm_per_head_native(q, q_weight, head_dim, eps)
ref_k_out = self._gemma_rmsnorm_per_head_native(k, k_weight, head_dim, eps)
residual = torch.randn([m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
atol = rtol = precision[ref_q_out.dtype]
torch.testing.assert_close(q_out, ref_q_out, atol=atol, rtol=rtol)
torch.testing.assert_close(k_out, ref_k_out, atol=atol, rtol=rtol)
add_ln_out = torch.ops.sgl_kernel.fused_add_layernorm_cpu(
x, residual, weight, bias, variance_epsilon
)
ref_add_ln_out, ref_residual = self._forward_native(
x, weight, variance_epsilon, residual=ref_residual, bias=bias
)
torch.testing.assert_close(add_ln_out, ref_add_ln_out, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
@parametrize(
l=[4096, 1024],
m=[1, 4],
n=[4096, 4109, 2304],
dtype=[torch.float16, torch.bfloat16],
@pytest.mark.parametrize("dtype", [torch.bfloat16], ids=["bfloat16"])
@pytest.mark.parametrize(
"batch_size,num_head,num_head_kv,head_dim",
[
(8, 4, 2, 128),
(17, 8, 2, 64),
(5, 3, 1, 96),
],
)
def test_norm_input_3d(self, l: int, m: int, n: int, dtype: torch.dtype) -> None:
x = torch.randn([l, m, n], dtype=dtype)
x = make_non_contiguous(x)
hidden_size = x.size(-1)
weight = torch.randn(hidden_size, dtype=dtype)
bias = torch.randn(hidden_size, dtype=dtype)
variance_epsilon = 1e-6
ln_out = torch.ops.sgl_kernel.layernorm_cpu(x, weight, None, variance_epsilon)
ref_ln_out = self._forward_native(x, weight, variance_epsilon)
atol = rtol = precision[ref_ln_out.dtype]
torch.testing.assert_close(ln_out, ref_ln_out, atol=atol, rtol=rtol)
ln_out = torch.ops.sgl_kernel.layernorm_cpu(x, weight, bias, variance_epsilon)
ref_ln_out = self._forward_native(
x, weight, variance_epsilon, residual=None, bias=bias
def test_fused_qk_gemma_rmsnorm_with_gate(
self, batch_size: int, num_head: int, num_head_kv: int, head_dim: int, dtype
):
q = torch.randn([batch_size, num_head, head_dim], dtype=dtype)
gate = torch.randn([batch_size, num_head, head_dim], dtype=dtype)
q_gate = torch.cat((q, gate), dim=-1).reshape(
batch_size, num_head * head_dim * 2
)
torch.testing.assert_close(ln_out, ref_ln_out, atol=atol, rtol=rtol)
k = torch.randn([batch_size, num_head_kv * head_dim], dtype=dtype)
residual = torch.randn([l, m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
q_gate = make_non_contiguous(q_gate)
k = make_non_contiguous(k)
add_ln_out = torch.ops.sgl_kernel.fused_add_layernorm_cpu(
x, residual, weight, None, variance_epsilon
)
ref_add_ln_out, ref_residual = self._forward_native(
x, weight, variance_epsilon, ref_residual
q_weight = torch.randn(head_dim, dtype=dtype)
k_weight = torch.randn(head_dim, dtype=dtype)
q_out, k_out, gate_out = (
torch.ops.sgl_kernel.fused_qk_gemma_rmsnorm_with_gate_cpu(
q_gate, k, q_weight, k_weight, eps, head_dim, num_head
)
)
torch.testing.assert_close(add_ln_out, ref_add_ln_out, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
ref_q_out = self._gemma_rmsnorm_per_head_native(q, q_weight, head_dim, eps)
ref_k_out = self._gemma_rmsnorm_per_head_native(k, k_weight, head_dim, eps)
residual = torch.randn([l, m, hidden_size], dtype=dtype)
ref_residual = residual.clone()
add_ln_out = torch.ops.sgl_kernel.fused_add_layernorm_cpu(
x, residual, weight, bias, variance_epsilon
atol = rtol = precision[ref_q_out.dtype]
torch.testing.assert_close(
q_out, ref_q_out.reshape(-1, head_dim), atol=atol, rtol=rtol
)
ref_add_ln_out, ref_residual = self._forward_native(
x, weight, variance_epsilon, residual=ref_residual, bias=bias
torch.testing.assert_close(
k_out, ref_k_out.reshape(-1, head_dim), atol=atol, rtol=rtol
)
torch.testing.assert_close(
gate_out, gate.reshape(-1, head_dim), atol=atol, rtol=rtol
)
torch.testing.assert_close(add_ln_out, ref_add_ln_out, atol=atol, rtol=rtol)
torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol)
if __name__ == "__main__":
unittest.main()
sys.exit(pytest.main([__file__]))