fix(qsa): dequantize FP8 cached prefixes in the sparse prefill kernels (#38855)
This commit is contained in:
@@ -243,6 +243,13 @@ def _sparse_gqa_chunk_prefill(
|
|||||||
mask=valid[:, None],
|
mask=valid[:, None],
|
||||||
other=0.0,
|
other=0.0,
|
||||||
)
|
)
|
||||||
|
# The chunk-prefill K/V tensors are gathered from the KV pool and can
|
||||||
|
# therefore carry the FP8 storage dtype, which Triton's dot rejects
|
||||||
|
# (`Unsupported rhs dtype fp8e4nv`). Convert to Q's dtype; the QSA
|
||||||
|
# backend writes the pool without per-tensor k/v scales, so this is a
|
||||||
|
# plain cast (no-op for BF16 pools).
|
||||||
|
keys = keys.to(q_values.dtype)
|
||||||
|
values = values.to(q_values.dtype)
|
||||||
scores = tl.where(valid[None, :], tl.dot(q_values, keys), -float("inf"))
|
scores = tl.where(valid[None, :], tl.dot(q_values, keys), -float("inf"))
|
||||||
next_max = tl.maximum(max_value, tl.max(scores, 1))
|
next_max = tl.maximum(max_value, tl.max(scores, 1))
|
||||||
alpha = tl.math.exp2(max_value - next_max)
|
alpha = tl.math.exp2(max_value - next_max)
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer
|
|||||||
from sglang.srt.layers.attention.qsa.sparse_attn import (
|
from sglang.srt.layers.attention.qsa.sparse_attn import (
|
||||||
qwen_sparse_fa2_cu_seqlens_triton,
|
qwen_sparse_fa2_cu_seqlens_triton,
|
||||||
qwen_sparse_kv_extraction_compact_triton,
|
qwen_sparse_kv_extraction_compact_triton,
|
||||||
|
sparse_gqa_fwd_interface_triton_ck,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.qwen_sparse_attn_backend import (
|
from sglang.srt.layers.attention.qwen_sparse_attn_backend import (
|
||||||
QwenSparseAttnBackend,
|
QwenSparseAttnBackend,
|
||||||
@@ -44,6 +45,60 @@ BLOCK_TOPK = TOKEN_TOPK // COMPRESS_RATIO
|
|||||||
FINAL_TOPK = TOKEN_TOPK + COMPRESS_RATIO - 1
|
FINAL_TOPK = TOKEN_TOPK + COMPRESS_RATIO - 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_qsa_chunk_prefill_accepts_fp8_cached_prefix():
|
||||||
|
if not torch.cuda.is_available() or torch.cuda.get_device_capability() < (8, 9):
|
||||||
|
pytest.skip("FP8-capable CUDA GPU required")
|
||||||
|
|
||||||
|
torch.manual_seed(42)
|
||||||
|
device = torch.device("cuda")
|
||||||
|
num_requests, num_q_heads, num_kv_heads, head_dim, topk = 2, 4, 1, 128, 16
|
||||||
|
q = torch.randn(
|
||||||
|
num_requests, num_q_heads, head_dim, dtype=torch.bfloat16, device=device
|
||||||
|
)
|
||||||
|
k = torch.randn(
|
||||||
|
num_requests * topk,
|
||||||
|
num_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
).to(torch.float8_e4m3fn)
|
||||||
|
v = torch.randn(
|
||||||
|
num_requests * topk,
|
||||||
|
num_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
).to(torch.float8_e4m3fn)
|
||||||
|
indices = torch.arange(topk, dtype=torch.int32, device=device).repeat(
|
||||||
|
num_requests, 1
|
||||||
|
)
|
||||||
|
cu_q = torch.arange(num_requests + 1, dtype=torch.int32, device=device)
|
||||||
|
cu_k = torch.arange(
|
||||||
|
0,
|
||||||
|
(num_requests + 1) * topk,
|
||||||
|
topk,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
kv_lens = torch.full((num_requests,), topk, dtype=torch.int32, device=device)
|
||||||
|
scale = head_dim**-0.5
|
||||||
|
|
||||||
|
actual = sparse_gqa_fwd_interface_triton_ck(
|
||||||
|
q, k, v, indices, cu_q, cu_k, kv_lens, scale
|
||||||
|
)
|
||||||
|
expected = sparse_gqa_fwd_interface_triton_ck(
|
||||||
|
q,
|
||||||
|
k.to(torch.bfloat16),
|
||||||
|
v.to(torch.bfloat16),
|
||||||
|
indices,
|
||||||
|
cu_q,
|
||||||
|
cu_k,
|
||||||
|
kv_lens,
|
||||||
|
scale,
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
("capability", "expected"),
|
("capability", "expected"),
|
||||||
[((12, 0), True), ((12, 1), False), ((10, 0), False)],
|
[((12, 0), True), ((12, 1), False), ((10, 0), False)],
|
||||||
|
|||||||
Reference in New Issue
Block a user