Add fused EH norm for DeepSeek NextN (#29667)

This commit is contained in:
Mohammad Miadh Angkad
2026-07-01 02:46:21 -07:00
committed by GitHub
parent 8ee200972e
commit 07ca24372b
5 changed files with 386 additions and 12 deletions
@@ -0,0 +1,63 @@
from __future__ import annotations
import torch
from sglang.jit_kernel.benchmark import marker
from sglang.jit_kernel.fused_eh_norm import fused_eh_norm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=6, suite="base-b-kernel-benchmark-1-gpu-large")
EPS = 1e-6
def reference(
x: torch.Tensor,
prev: torch.Tensor,
ew: torch.Tensor,
hw: torch.Tensor,
eps: float,
) -> torch.Tensor:
xf = x.float()
pf = prev.float()
x_var = xf.pow(2).mean(dim=-1, keepdim=True)
p_var = pf.pow(2).mean(dim=-1, keepdim=True)
return torch.cat(
(
(xf * torch.rsqrt(x_var + eps) * ew.float()).to(x.dtype),
(pf * torch.rsqrt(p_var + eps) * hw.float()).to(prev.dtype),
),
dim=-1,
)
FN_MAP = {
"jit": fused_eh_norm,
"torch": reference,
}
@marker.parametrize("dtype", [torch.bfloat16, torch.float16], [torch.bfloat16])
@marker.parametrize("hidden_size", [6144, 7168], [7168])
@marker.parametrize("num_tokens", [1, 4, 6, 8, 16, 32, 128, 512], [1, 6])
@marker.benchmark("impl", ["jit", "torch"])
def benchmark(num_tokens: int, hidden_size: int, dtype: torch.dtype, impl: str):
x = torch.randn(num_tokens, hidden_size, device="cuda", dtype=dtype)
prev = torch.randn_like(x)
ew = torch.randn(hidden_size, device="cuda", dtype=dtype)
hw = torch.randn(hidden_size, device="cuda", dtype=dtype)
expected = reference(x, prev, ew, hw, EPS)
actual = fused_eh_norm(x, prev, ew, hw, EPS)
torch.testing.assert_close(actual.float(), expected.float(), rtol=1e-2, atol=1e-2)
return marker.do_bench(
FN_MAP[impl],
input_args=(x, prev, ew, hw, EPS),
memory_args=(x, prev, ew, hw),
memory_output="out",
)
if __name__ == "__main__":
benchmark.run()
+122
View File
@@ -0,0 +1,122 @@
import sys
import pytest
import torch
from sglang.jit_kernel.fused_eh_norm import fused_eh_norm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=45, suite="base-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available() or torch.version.cuda is None,
reason="fused_eh_norm requires CUDA",
)
def _reference(
inputs_embeds: torch.Tensor,
previous_hidden: torch.Tensor,
enorm_weight: torch.Tensor,
hnorm_weight: torch.Tensor,
eps: float,
) -> torch.Tensor:
embeds = inputs_embeds.float()
prev = previous_hidden.float()
embeds_var = embeds.pow(2).mean(dim=-1, keepdim=True)
prev_var = prev.pow(2).mean(dim=-1, keepdim=True)
return torch.cat(
(
(embeds * torch.rsqrt(embeds_var + eps) * enorm_weight.float()).to(
inputs_embeds.dtype
),
(prev * torch.rsqrt(prev_var + eps) * hnorm_weight.float()).to(
previous_hidden.dtype
),
),
dim=-1,
)
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize("hidden_size", [6144, 7168])
@pytest.mark.parametrize("num_tokens", [1, 6, 128])
def test_fused_eh_norm_matches_reference(
dtype: torch.dtype, hidden_size: int, num_tokens: int
):
torch.manual_seed(0)
eps = 1e-6
inputs_embeds = torch.randn(num_tokens, hidden_size, device="cuda", dtype=dtype)
previous_hidden = torch.randn_like(inputs_embeds)
enorm_weight = torch.randn(hidden_size, device="cuda", dtype=dtype)
hnorm_weight = torch.randn(hidden_size, device="cuda", dtype=dtype)
actual = fused_eh_norm(
inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, eps
)
expected = _reference(
inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, eps
)
torch.testing.assert_close(actual.float(), expected.float(), rtol=1e-2, atol=1e-2)
def test_fused_eh_norm_zero_tokens():
hidden_size = 7168
inputs_embeds = torch.empty(0, hidden_size, device="cuda", dtype=torch.bfloat16)
previous_hidden = torch.empty_like(inputs_embeds)
enorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
hnorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
actual = fused_eh_norm(
inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, 1e-6
)
assert actual.shape == (0, hidden_size * 2)
assert actual.dtype == inputs_embeds.dtype
assert actual.device == inputs_embeds.device
def test_fused_eh_norm_row_strided_inputs():
torch.manual_seed(1)
hidden_size = 7168
eps = 1e-6
base = torch.randn(12, hidden_size, device="cuda", dtype=torch.bfloat16)
prev_base = torch.randn_like(base)
inputs_embeds = base[::2]
previous_hidden = prev_base[::2]
enorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
hnorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
actual = fused_eh_norm(
inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, eps
)
expected = _reference(
inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, eps
)
torch.testing.assert_close(actual.float(), expected.float(), rtol=1e-2, atol=1e-2)
def test_fused_eh_norm_rejects_unsupported_dtype():
hidden_size = 7168
inputs_embeds = torch.randn(1, hidden_size, device="cuda", dtype=torch.float32)
previous_hidden = torch.randn_like(inputs_embeds)
enorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.float32)
hnorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.float32)
with pytest.raises(RuntimeError, match="unsupported dtype"):
fused_eh_norm(inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, 1e-6)
def test_fused_eh_norm_rejects_unsupported_hidden_size():
hidden_size = 5000
inputs_embeds = torch.randn(1, hidden_size, device="cuda", dtype=torch.bfloat16)
previous_hidden = torch.randn_like(inputs_embeds)
enorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
hnorm_weight = torch.randn(hidden_size, device="cuda", dtype=torch.bfloat16)
with pytest.raises(RuntimeError, match="unsupported hidden_size"):
fused_eh_norm(inputs_embeds, previous_hidden, enorm_weight, hnorm_weight, 1e-6)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))