Add fused EH norm for DeepSeek NextN (#29667)
This commit is contained in:
@@ -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()
|
||||
@@ -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"]))
|
||||
Reference in New Issue
Block a user