Files
sglang/test/registered/attention/test_chunk_gated_delta_rule.py
2026-03-23 00:18:45 -07:00

275 lines
9.4 KiB
Python

import unittest
import torch
from sglang.srt.layers.attention.fla.chunk import chunk_gated_delta_rule
from sglang.srt.layers.attention.fla.fused_recurrent import (
fused_recurrent_gated_delta_rule,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, suite="stage-b-test-1-gpu-large")
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
class TestChunkGatedDeltaRule(unittest.TestCase):
"""Test chunk_gated_delta_rule against token-by-token fused_recurrent reference."""
ATOL = 2e-2
RTOL = 1e-2
def _run_reference(self, pool_init, cache_indices, q, k, v, g, beta):
"""Per-batch token-by-token reference using fused_recurrent_gated_delta_rule.
initial_state shape: [N, H, V, K] (native layout on this branch).
"""
B = cache_indices.shape[0]
T_per_seq = q.shape[1] // B
pool = pool_init.clone()
h_cur = pool[cache_indices].contiguous().clone()
o_list = []
for b in range(B):
sl = slice(b * T_per_seq, (b + 1) * T_per_seq)
o_b, h_b = fused_recurrent_gated_delta_rule(
q=q[0, sl].unsqueeze(0),
k=k[0, sl].unsqueeze(0),
v=v[0, sl].unsqueeze(0),
g=g[0, sl].unsqueeze(0),
beta=beta[0, sl].unsqueeze(0),
initial_state=h_cur[b : b + 1],
output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
o_list.append(o_b)
h_cur[b] = h_b[0]
pool[cache_indices] = h_cur
return torch.cat(o_list, dim=1), pool
def _run_chunk(self, pool_init, cache_indices, q, k, v, g, beta, cu_seqlens):
"""Run chunk_gated_delta_rule with native [V, K] pool."""
pool = pool_init.clone()
o, _, _ = chunk_gated_delta_rule(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=pool,
initial_state_indices=cache_indices,
cu_seqlens=cu_seqlens,
head_first=False,
use_qk_l2norm_in_kernel=True,
)
return o, pool
def _check_shape(
self, B, T_per_seq, H, K, V, pool_size, sequential_indices=False, seed=42
):
"""Run correctness check for one (B, T_per_seq, H, K, V, pool_size) config."""
device = "cuda"
dtype = torch.bfloat16
T = B * T_per_seq
torch.manual_seed(seed)
if sequential_indices:
cache_indices = torch.arange(B, dtype=torch.int32, device=device)
else:
perm = torch.randperm(pool_size, device=device)[:B]
cache_indices = perm.to(torch.int32)
pool_init = (
torch.randn(pool_size, H, V, K, dtype=torch.float32, device=device) * 0.1
)
cu_seqlens = torch.zeros(B + 1, dtype=torch.long, device=device)
cu_seqlens[1:] = (
torch.arange(1, B + 1, dtype=torch.long, device=device) * T_per_seq
)
q = torch.randn(1, T, H, K, dtype=dtype, device=device)
k = torch.randn(1, T, H, K, dtype=dtype, device=device)
v = torch.randn(1, T, H, V, dtype=dtype, device=device)
g = torch.nn.functional.logsigmoid(
torch.randn(1, T, H, dtype=dtype, device=device)
)
beta = torch.sigmoid(torch.randn(1, T, H, dtype=dtype, device=device))
o_ref, pool_ref = self._run_reference(
pool_init, cache_indices, q, k, v, g, beta
)
o_new, pool_new = self._run_chunk(
pool_init, cache_indices, q, k, v, g, beta, cu_seqlens
)
self.assertTrue(
torch.allclose(
o_ref.float(), o_new.float(), atol=self.ATOL, rtol=self.RTOL
),
f"Output mismatch: max_diff="
f"{(o_ref.float() - o_new.float()).abs().max().item():.2e}",
)
ref_slots = pool_ref[cache_indices].contiguous()
new_slots = pool_new[cache_indices].contiguous()
self.assertTrue(
torch.allclose(
ref_slots.float(), new_slots.float(), atol=self.ATOL, rtol=self.RTOL
),
f"State mismatch: max_diff="
f"{(ref_slots.float() - new_slots.float()).abs().max().item():.2e}",
)
# ------------------------------------------------------------------
# Production-style configs (Qwen3-Next)
# ------------------------------------------------------------------
def test_production_nt1(self):
self._check_shape(B=4, T_per_seq=64, H=16, K=128, V=128, pool_size=32)
def test_production_nt2(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=128, V=128, pool_size=32)
def test_production_nt4(self):
self._check_shape(B=4, T_per_seq=256, H=16, K=128, V=128, pool_size=32)
# ------------------------------------------------------------------
# Batch size sweep
# ------------------------------------------------------------------
def test_batch_1(self):
self._check_shape(B=1, T_per_seq=128, H=16, K=128, V=128, pool_size=32)
def test_batch_2(self):
self._check_shape(B=2, T_per_seq=128, H=16, K=128, V=128, pool_size=32)
def test_batch_8(self):
self._check_shape(B=8, T_per_seq=128, H=16, K=128, V=128, pool_size=64)
def test_batch_16(self):
self._check_shape(B=16, T_per_seq=64, H=16, K=128, V=128, pool_size=128)
def test_batch_32(self):
self._check_shape(B=32, T_per_seq=32, H=16, K=128, V=128, pool_size=256)
# ------------------------------------------------------------------
# Head count sweep
# ------------------------------------------------------------------
def test_heads_4(self):
self._check_shape(B=4, T_per_seq=128, H=4, K=128, V=128, pool_size=32)
def test_heads_8(self):
self._check_shape(B=4, T_per_seq=128, H=8, K=128, V=128, pool_size=32)
def test_heads_32(self):
self._check_shape(B=4, T_per_seq=128, H=32, K=128, V=128, pool_size=32)
def test_heads_64(self):
self._check_shape(B=4, T_per_seq=128, H=64, K=128, V=128, pool_size=32)
# ------------------------------------------------------------------
# K != V (exercises that [V,K] != [K,V] byte-order matters)
# ------------------------------------------------------------------
def test_dim_64x64(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=64, V=64, pool_size=32)
def test_dim_k_lt_v(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=64, V=128, pool_size=32)
def test_dim_k_gt_v(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=128, V=64, pool_size=32)
def test_dim_256x256(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=256, V=256, pool_size=32)
# ------------------------------------------------------------------
# Short sequences (T < chunk_size=64)
# ------------------------------------------------------------------
def test_seqlen_1(self):
self._check_shape(B=4, T_per_seq=1, H=16, K=128, V=128, pool_size=32)
def test_seqlen_7(self):
self._check_shape(B=4, T_per_seq=7, H=16, K=128, V=128, pool_size=32)
def test_seqlen_16(self):
self._check_shape(B=4, T_per_seq=16, H=16, K=128, V=128, pool_size=32)
def test_seqlen_32(self):
self._check_shape(B=4, T_per_seq=32, H=16, K=128, V=128, pool_size=32)
# ------------------------------------------------------------------
# Multi-chunk and large pool
# ------------------------------------------------------------------
def test_multi_chunk_nt8(self):
self._check_shape(B=4, T_per_seq=512, H=16, K=128, V=128, pool_size=32)
def test_large_pool(self):
self._check_shape(B=4, T_per_seq=128, H=16, K=128, V=128, pool_size=512)
# ------------------------------------------------------------------
# Combined stress
# ------------------------------------------------------------------
def test_stress(self):
self._check_shape(B=32, T_per_seq=128, H=32, K=128, V=128, pool_size=256)
# ------------------------------------------------------------------
# Sequential-index variants (pool_size == B, indices = 0..B-1)
# ------------------------------------------------------------------
def test_seq_idx_b4(self):
self._check_shape(
B=4,
T_per_seq=128,
H=16,
K=128,
V=128,
pool_size=4,
sequential_indices=True,
)
def test_seq_idx_b8(self):
self._check_shape(
B=8,
T_per_seq=128,
H=16,
K=128,
V=128,
pool_size=8,
sequential_indices=True,
)
def test_seq_idx_h32(self):
self._check_shape(
B=4,
T_per_seq=128,
H=32,
K=128,
V=128,
pool_size=4,
sequential_indices=True,
)
def test_seq_idx_h64(self):
self._check_shape(
B=4,
T_per_seq=128,
H=64,
K=128,
V=128,
pool_size=4,
sequential_indices=True,
)
def test_seq_idx_stress(self):
self._check_shape(
B=32,
T_per_seq=128,
H=32,
K=128,
V=128,
pool_size=32,
sequential_indices=True,
)
if __name__ == "__main__":
unittest.main()