[KDA] Fuse gate+cumsum and reuse chunk index for KDA (#23038)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -2,10 +2,15 @@ import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.fla.cumsum import chunk_local_cumsum
|
||||
from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import (
|
||||
fused_sigmoid_gating_delta_rule_update,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.kda import fused_kda_gate, fused_recurrent_kda
|
||||
from sglang.srt.layers.attention.fla.index import prepare_chunk_indices
|
||||
from sglang.srt.layers.attention.fla.kda import (
|
||||
fused_recurrent_kda,
|
||||
kda_gate_chunk_cumsum,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=8, suite="stage-b-test-1-gpu-large")
|
||||
@@ -95,7 +100,15 @@ class TestKDAFusedSigmoidGatingRecurrent(unittest.TestCase):
|
||||
|
||||
def run_kda(self):
|
||||
b = self.beta.float().sigmoid()
|
||||
g = fused_kda_gate(self.a, self.A_log, self.head_dim, g_bias=self.dt_bias)
|
||||
# Reference gate activation using torch ops:
|
||||
# g = -exp(A_log) * softplus(raw_g + dt_bias)
|
||||
H, K = self.local_num_heads, self.head_dim
|
||||
raw_g = self.a.float() # [1, T, H*K]
|
||||
if self.dt_bias is not None:
|
||||
raw_g = raw_g + self.dt_bias.float()
|
||||
g = -torch.exp(
|
||||
self.A_log.float().view(1, 1, H, 1)
|
||||
) * torch.nn.functional.softplus(raw_g.view(1, -1, H, K))
|
||||
initial_state = self.ssm_states[self.cache_indices].clone()
|
||||
core_attn_out, last_state = fused_recurrent_kda(
|
||||
q=self.q,
|
||||
@@ -121,5 +134,101 @@ class TestKDAFusedSigmoidGatingRecurrent(unittest.TestCase):
|
||||
self.assertTrue(torch.allclose(last_state, last_state_ref))
|
||||
|
||||
|
||||
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
|
||||
class TestKDAGateChunkCumsum(unittest.TestCase):
|
||||
"""Test kda_gate_chunk_cumsum against torch reference (gate activation + cumsum)."""
|
||||
|
||||
CHUNK_SIZE = 64
|
||||
|
||||
def _ref_gate_cumsum(self, raw_g, A_log, dt_bias, cu_seqlens, chunk_size):
|
||||
"""Reference: torch gate activation then chunk_local_cumsum."""
|
||||
B, T, H, K = raw_g.shape
|
||||
g = raw_g.float()
|
||||
if dt_bias is not None:
|
||||
g = g + dt_bias.float().view(1, 1, H, K)
|
||||
g = -torch.exp(A_log.float().view(1, 1, H, 1)) * torch.nn.functional.softplus(g)
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, chunk_size)
|
||||
if cu_seqlens is not None
|
||||
else None
|
||||
)
|
||||
return chunk_local_cumsum(
|
||||
g, chunk_size=chunk_size, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices
|
||||
)
|
||||
|
||||
def _run_case(self, B, T_per_seq, H, K, use_bias, use_varlen):
|
||||
T = B * T_per_seq
|
||||
torch.manual_seed(42)
|
||||
raw_g = torch.randn(1, T, H, K, dtype=torch.bfloat16, device="cuda")
|
||||
A_log = torch.randn(H, dtype=torch.float32, device="cuda") * 0.5
|
||||
dt_bias = (
|
||||
torch.randn(H * K, dtype=torch.float32, device="cuda") * 0.1
|
||||
if use_bias
|
||||
else None
|
||||
)
|
||||
cu_seqlens = (
|
||||
torch.arange(
|
||||
0, (B + 1) * T_per_seq, T_per_seq, dtype=torch.long, device="cuda"
|
||||
)
|
||||
if use_varlen
|
||||
else None
|
||||
)
|
||||
|
||||
out_fused = kda_gate_chunk_cumsum(
|
||||
raw_g,
|
||||
A_log=A_log,
|
||||
chunk_size=self.CHUNK_SIZE,
|
||||
dt_bias=dt_bias,
|
||||
cu_seqlens=cu_seqlens,
|
||||
)
|
||||
out_ref = self._ref_gate_cumsum(
|
||||
raw_g, A_log, dt_bias, cu_seqlens, self.CHUNK_SIZE
|
||||
)
|
||||
|
||||
max_diff = (out_fused - out_ref).abs().max().item()
|
||||
rel_diff = max_diff / (out_ref.abs().mean().item() + 1e-8)
|
||||
return max_diff, rel_diff
|
||||
|
||||
def test_varlen_with_bias(self):
|
||||
max_diff, rel_diff = self._run_case(
|
||||
B=4, T_per_seq=256, H=16, K=128, use_bias=True, use_varlen=True
|
||||
)
|
||||
self.assertLess(
|
||||
max_diff, 1e-3, f"max_diff={max_diff:.2e}, rel_diff={rel_diff:.2e}"
|
||||
)
|
||||
|
||||
def test_varlen_no_bias(self):
|
||||
max_diff, rel_diff = self._run_case(
|
||||
B=4, T_per_seq=256, H=16, K=128, use_bias=False, use_varlen=True
|
||||
)
|
||||
self.assertLess(
|
||||
max_diff, 1e-3, f"max_diff={max_diff:.2e}, rel_diff={rel_diff:.2e}"
|
||||
)
|
||||
|
||||
def test_fixed_len_with_bias(self):
|
||||
max_diff, rel_diff = self._run_case(
|
||||
B=4, T_per_seq=256, H=16, K=128, use_bias=True, use_varlen=False
|
||||
)
|
||||
self.assertLess(
|
||||
max_diff, 1e-3, f"max_diff={max_diff:.2e}, rel_diff={rel_diff:.2e}"
|
||||
)
|
||||
|
||||
def test_single_seq_long(self):
|
||||
max_diff, rel_diff = self._run_case(
|
||||
B=1, T_per_seq=2048, H=16, K=128, use_bias=True, use_varlen=True
|
||||
)
|
||||
self.assertLess(
|
||||
max_diff, 1e-3, f"max_diff={max_diff:.2e}, rel_diff={rel_diff:.2e}"
|
||||
)
|
||||
|
||||
def test_small_head_dim(self):
|
||||
max_diff, rel_diff = self._run_case(
|
||||
B=4, T_per_seq=128, H=8, K=64, use_bias=True, use_varlen=True
|
||||
)
|
||||
self.assertLess(
|
||||
max_diff, 1e-3, f"max_diff={max_diff:.2e}, rel_diff={rel_diff:.2e}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user