Files
sglang/test/registered/kernels/ops/attention/test_kda_prefill.py
T

209 lines
6.7 KiB
Python

import unittest
from unittest.mock import patch
import torch
import torch.nn.functional as F
from sglang.kernels.ops.attention.fla.kda import chunk_kda
from sglang.kernels.ops.attention.linear.kda_nvidia_prefill import (
chunk_kda_fwd as nvidia_chunk_kda_fwd,
)
from sglang.kernels.ops.attention.linear.kda_ptx_prefill import (
chunk_kda_fwd as ptx_chunk_kda_fwd,
)
from sglang.srt.layers.attention.linear.kernels.kda_ptx import PtxKDAKernel
from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKernel
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=180, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
register_cuda_ci(est_time=80, stage="base-c", runner_config="4-gpu-gb300")
def _inputs(seed, seq_len=128):
generator = torch.Generator(device="cuda").manual_seed(seed)
batch_size, num_heads, head_dim = 1, 2, 128
shape = (batch_size, seq_len, num_heads, head_dim)
q = torch.randn(shape, generator=generator, device="cuda", dtype=torch.bfloat16)
k = torch.randn(shape, generator=generator, device="cuda", dtype=torch.bfloat16)
v = (
0.1
* torch.randn(
shape,
generator=generator,
device="cuda",
dtype=torch.float32,
)
).to(torch.bfloat16)
gate = torch.randn(shape, generator=generator, device="cuda", dtype=torch.bfloat16)
beta_logits = torch.randn(
shape[:-1],
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
a_log = torch.randn(
num_heads, generator=generator, device="cuda", dtype=torch.float32
)
dt_bias = torch.randn(
num_heads * head_dim,
generator=generator,
device="cuda",
dtype=torch.float32,
)
state = torch.zeros(
batch_size,
num_heads,
head_dim,
head_dim,
device="cuda",
dtype=torch.float32,
)
return q, k, v, gate, beta_logits, a_log, dt_bias, state
def _reference(q, k, v, gate, beta, a_log, dt_bias, state, fused_qk_norm):
return chunk_kda(
q=q,
k=k,
v=v,
g=gate,
beta=beta,
scale=q.shape[-1] ** -0.5,
initial_state=state,
initial_state_indices=torch.arange(
q.shape[0], device="cuda", dtype=torch.int32
),
use_qk_l2norm_in_kernel=fused_qk_norm,
A_log=a_log,
dt_bias=dt_bias,
lower_bound=-5.0,
)
class TestKdaPrefill(CustomTestCase):
@torch.inference_mode()
def test_ptx_padded_raw_beta(self):
"""Raw beta must match Triton, including final state after neutral padding."""
if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (
10,
3,
):
self.skipTest("PTX KDA prefill requires GB300")
q, k, v, gate, beta, a_log, dt_bias, state = _inputs(2, seq_len=1025)
state.fill_(0.1)
actual_state = state.clone()
inputs = dict(
q=q,
k=k,
v=v,
g=gate,
beta=beta,
cache_indices=torch.zeros(1, device="cuda", dtype=torch.int32),
query_start_loc=torch.tensor([0, 1025], device="cuda", dtype=torch.int32),
A_log=a_log,
dt_bias=dt_bias,
lower_bound=-5.0,
beta_is_raw=True,
extend_seq_lens_cpu=[1025],
)
kernel = PtxKDAKernel()
with patch.object(
kernel._triton,
"extend",
side_effect=AssertionError("PTX unexpectedly fell back to Triton"),
):
actual = kernel.extend(**inputs, ssm_states=actual_state)
# Triton may mutate inputs, so run the reference last.
expected = TritonKDAKernel().extend(**inputs, ssm_states=state)
torch.testing.assert_close(
actual.float(), expected.float(), rtol=2e-2, atol=3e-2
)
torch.testing.assert_close(actual_state, state, rtol=2e-2, atol=3e-2)
@torch.inference_mode()
def test_nvidia_prefill(self):
if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10:
self.skipTest("NVIDIA KDA prefill requires datacenter Blackwell")
q, k, v, gate, beta_logits, a_log, dt_bias, state = _inputs(0)
q = F.normalize(q.float(), dim=-1).to(torch.bfloat16)
k = F.normalize(k.float(), dim=-1).to(torch.bfloat16)
beta = torch.sigmoid(beta_logits.float()).to(torch.bfloat16)
actual, actual_state = nvidia_chunk_kda_fwd(
q=q,
k=k,
v=v,
g=gate,
beta=beta,
scale=q.shape[-1] ** -0.5,
initial_state=state.transpose(-1, -2).contiguous(),
output_final_state=True,
safe_gate=True,
lower_bound=-5.0,
use_gate_in_kernel=True,
A_log=a_log,
dt_bias=dt_bias,
)[:2]
expected = _reference(
q, k, v, gate, beta, a_log, dt_bias, state, fused_qk_norm=False
)
torch.testing.assert_close(
actual.float(), expected.float(), rtol=2e-2, atol=3e-2
)
torch.testing.assert_close(
actual_state.transpose(-1, -2),
state,
rtol=2e-2,
atol=3e-2,
)
@torch.inference_mode()
def test_ptx_prefill(self):
if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (
10,
3,
):
self.skipTest("PTX KDA prefill requires GB300")
q, k, v, gate, beta_logits, a_log, dt_bias, state = _inputs(1)
actual, actual_state = ptx_chunk_kda_fwd(
q=q,
k=k,
v=v,
g=gate,
beta=beta_logits,
scale=q.shape[-1] ** -0.5,
initial_state=state.transpose(-1, -2).contiguous(),
output_final_state=True,
safe_gate=True,
lower_bound=-5.0,
use_gate_in_kernel=True,
A_log=a_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_beta_sigmoid_in_kernel=True,
)[:2]
expected = _reference(
q,
k,
v,
gate,
torch.sigmoid(beta_logits.float()).to(torch.bfloat16),
a_log,
dt_bias,
state,
fused_qk_norm=True,
)
torch.testing.assert_close(
actual.float(), expected.float(), rtol=2e-2, atol=3e-2
)
torch.testing.assert_close(
actual_state.transpose(-1, -2),
state,
rtol=2e-2,
atol=3e-2,
)
if __name__ == "__main__":
unittest.main()