112 lines
3.6 KiB
Python
112 lines
3.6 KiB
Python
import pytest
|
|
import torch
|
|
|
|
import sglang.srt.layers.mhc as mhc
|
|
from sglang.srt.layers.mhc import mhc_fused_post_pre, mhc_post, mhc_pre
|
|
|
|
|
|
@pytest.mark.parametrize("hidden_size", [4096, 7168])
|
|
@pytest.mark.parametrize("num_tokens", [0, 1, 8, 17, 32, 64])
|
|
@pytest.mark.parametrize("use_norm", [False, True])
|
|
def test_mhc_fused_post_pre_matches_unfused(
|
|
monkeypatch, hidden_size, num_tokens, use_norm
|
|
):
|
|
if not torch.cuda.is_available():
|
|
pytest.skip("CUDA is required for TileLang mHC kernels")
|
|
|
|
monkeypatch.setattr(mhc, "is_dsa_prefill_cp_round_robin_split", lambda: False)
|
|
torch.manual_seed(0)
|
|
device = torch.device("cuda")
|
|
hc_mult = 4
|
|
hc_mult3 = hc_mult * 2 + hc_mult * hc_mult
|
|
hc_hidden_size = hc_mult * hidden_size
|
|
|
|
x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) * 0.1
|
|
residual = (
|
|
torch.randn(
|
|
num_tokens, hc_mult, hidden_size, device=device, dtype=torch.bfloat16
|
|
)
|
|
* 0.1
|
|
)
|
|
post_prev = torch.rand(num_tokens, hc_mult, 1, device=device, dtype=torch.float32)
|
|
comb_prev = (
|
|
torch.rand(num_tokens, hc_mult, hc_mult, device=device, dtype=torch.float32)
|
|
* 0.25
|
|
)
|
|
fn = (
|
|
torch.randn(hc_mult3, hc_hidden_size, device=device, dtype=torch.float32) * 0.01
|
|
)
|
|
hc_scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32)
|
|
hc_base = torch.zeros(hc_mult3, device=device, dtype=torch.float32)
|
|
norm_weight = (
|
|
torch.ones(hidden_size, device=device, dtype=torch.bfloat16)
|
|
if use_norm
|
|
else None
|
|
)
|
|
norm_eps = 1e-6 if use_norm else None
|
|
|
|
rms_eps = 1e-6
|
|
hc_eps = 1e-6
|
|
sinkhorn_repeat = 2
|
|
|
|
residual_ref = post_ref = comb_ref = layer_ref = None
|
|
if num_tokens > 0:
|
|
residual_ref = mhc_post(x, residual, post_prev, comb_prev)
|
|
post_ref, comb_ref, layer_ref = mhc_pre(
|
|
residual_ref,
|
|
fn,
|
|
hc_scale,
|
|
hc_base,
|
|
rms_eps,
|
|
hc_eps,
|
|
hc_eps,
|
|
2.0,
|
|
sinkhorn_repeat,
|
|
norm_weight=norm_weight,
|
|
norm_eps=norm_eps,
|
|
)
|
|
residual_out, post_out, comb_out, layer_out = mhc_fused_post_pre(
|
|
x,
|
|
residual,
|
|
post_prev,
|
|
comb_prev,
|
|
fn,
|
|
hc_scale,
|
|
hc_base,
|
|
rms_eps,
|
|
hc_eps,
|
|
hc_eps,
|
|
2.0,
|
|
sinkhorn_repeat,
|
|
norm_weight=norm_weight,
|
|
norm_eps=norm_eps,
|
|
)
|
|
|
|
torch.cuda.synchronize()
|
|
if num_tokens == 0:
|
|
assert residual_out.shape == residual.shape
|
|
assert post_out.shape == (0, hc_mult, 1)
|
|
assert comb_out.shape == (0, hc_mult, hc_mult)
|
|
assert layer_out.shape == (0, hidden_size)
|
|
assert residual_out.dtype == torch.bfloat16
|
|
assert post_out.dtype == torch.float32
|
|
assert comb_out.dtype == torch.float32
|
|
assert layer_out.dtype == torch.bfloat16
|
|
return
|
|
|
|
assert residual_ref is not None
|
|
assert post_ref is not None
|
|
assert comb_ref is not None
|
|
assert layer_ref is not None
|
|
assert residual_out.shape == residual_ref.shape
|
|
assert post_out.shape == post_ref.shape
|
|
assert comb_out.shape == comb_ref.shape
|
|
assert layer_out.shape == layer_ref.shape
|
|
|
|
torch.testing.assert_close(residual_out, residual_ref, atol=0, rtol=0)
|
|
torch.testing.assert_close(post_out, post_ref, atol=1e-3, rtol=1e-3)
|
|
torch.testing.assert_close(comb_out, comb_ref, atol=1e-3, rtol=1e-3)
|
|
layer_atol = 2e-2 if use_norm else 2e-3
|
|
layer_rtol = 2e-2 if use_norm else 2e-3
|
|
torch.testing.assert_close(layer_out, layer_ref, atol=layer_atol, rtol=layer_rtol)
|