Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 30 additions & 4 deletions test/registered/unit/models/test_kimi_k3_bfa_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,7 @@

import torch

from sglang.srt.models.kimi_k3 import (
KimiK3DeltaAttention,
_get_k3_dense_weight,
)
from sglang.srt.models.kimi_k3 import KimiK3DeltaAttention, _get_k3_dense_weight
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase

Expand Down Expand Up @@ -41,6 +38,7 @@ def fused_qkvg_proj(x):

owner = SimpleNamespace(
use_full_rank_gate=True,
_qkvg_w=qkvg_w,
_bfa_w=_randn(_BFA_W_ROWS, _H).contiguous(),
_bfa_f_b_w=_randn(1536, _N_FA).contiguous(),
_bfa_fa_size=_N_FA,
Expand All @@ -64,6 +62,34 @@ def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA is not available")

def test_fused_projections_match_unfused_math(self):
owner = _make_owner(with_stream=False)
qkv_w, gate_w = torch.split(owner._qkvg_w, owner.split_sizes)
fa_w = owner._bfa_w[: owner._bfa_fa_size]
beta_w = owner._bfa_w[
owner._bfa_fa_size : owner._bfa_fa_size + owner._bfa_b_size
]

for num_tokens in (1, 32):
with self.subTest(num_tokens=num_tokens):
x = torch.randn(num_tokens, _H, device="cuda", dtype=torch.bfloat16)
fused = _run(owner, x)
fa = torch.nn.functional.linear(x, fa_w)
unfused = (
torch.nn.functional.linear(x, qkv_w),
torch.nn.functional.linear(x, beta_w),
torch.nn.functional.linear(fa, owner._bfa_f_b_w),
torch.nn.functional.linear(x, gate_w),
)
for got, ref, name in zip(
fused, unfused, ("qkv", "beta", "forget_gate", "gate")
):
# The merged and standalone BF16 GEMMs may select different
# accumulation schedules; allow roughly one output ULP.
torch.testing.assert_close(
got, ref, rtol=1e-2, atol=2e-2, msg=lambda msg: f"{name}: {msg}"
)

def test_capture_replay_matches_serial(self):
torch.manual_seed(0)
for T in (1, 4, 12):
Expand Down
Loading