Cute-DSL FP8 MQA logits (#25220)

Co-authored-by: Mindy Li <11663212+limin2021@users.noreply.github.com>
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-07-06 23:34:07 -07:00
committed by GitHub
co-authored by Mindy Li Brayden Zhong
parent 2d9f0b3317
commit fefc1743a9
12 changed files with 3440 additions and 37 deletions
@@ -0,0 +1,234 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import sys
import pytest
import torch
from sglang.jit_kernel.dsa import cutedsl_paged_mqa_logits, pick_dsl_expand
from sglang.srt.layers.attention.dsa.utils import (
fp8_mqa_logits_ceil_to_ue8m0,
fp8_mqa_logits_make_fused_kv,
)
from sglang.srt.utils import is_sm100_supported
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=180, suite="nightly-4-gpu-b200", nightly=True)
BLOCK_KV = 64
HEAD_DIM = 128
def _ref_fp8_paged_mqa_logits(
q_fp8,
kv_fp8,
kv_scales,
weights,
context_lens,
block_table,
max_model_len,
block_kv,
):
B, next_n, H, D = q_fp8.shape
device = q_fp8.device
logits = torch.full(
(B * next_n, max_model_len), float("-inf"), device=device, dtype=torch.float32
)
q_f32 = q_fp8.float()
for b in range(B):
ctx_len = context_lens[b].item()
q_positions = torch.arange(ctx_len - next_n, ctx_len, device=device)
w = weights[b * next_n : (b + 1) * next_n, :]
for blk_idx in range((ctx_len + block_kv - 1) // block_kv):
phys_blk = block_table[b, blk_idx].item()
k_f32 = kv_fp8[phys_blk].float()
scales = kv_scales[phys_blk]
k_positions = torch.arange(
blk_idx * block_kv, (blk_idx + 1) * block_kv, device=device
)
mask = (k_positions[None, :] < ctx_len) & (
k_positions[None, :] <= q_positions[:, None]
)
qk = torch.matmul(q_f32[b].permute(1, 0, 2), k_f32.T)
qk = torch.where(mask[None, :, :], qk, torch.zeros(1, device=device))
qk = torch.relu(qk)
weighted = (w.T[:, :, None] * qk).sum(dim=0)
weighted = weighted * scales[None, :]
start_pos = blk_idx * block_kv
end_pos = start_pos + block_kv
logits[b * next_n : (b + 1) * next_n, start_pos:end_pos] = torch.where(
mask,
weighted,
torch.tensor(float("-inf"), device=device, dtype=torch.float32),
)
return logits
def _generate_test_data(
batch_size,
next_n,
num_heads,
avg_context_len,
max_model_len,
device="cuda",
):
torch.manual_seed(42)
torch.cuda.manual_seed(42)
context_lens = torch.randint(
max(BLOCK_KV, int(0.7 * avg_context_len)),
int(1.3 * avg_context_len) + 1,
(batch_size,),
dtype=torch.int32,
device="cpu",
).clamp(max=max_model_len)
max_blocks_per_seq = (max_model_len + BLOCK_KV - 1) // BLOCK_KV
total_blocks = ((context_lens + BLOCK_KV - 1) // BLOCK_KV).sum().item()
num_phys_blocks = total_blocks + batch_size * 2
block_table = torch.full(
(batch_size, max_blocks_per_seq), 0, dtype=torch.int32, device=device
)
blk_offset = 0
for i in range(batch_size):
n_blks = (context_lens[i].item() + BLOCK_KV - 1) // BLOCK_KV
block_table[i, :n_blks] = torch.arange(
blk_offset, blk_offset + n_blks, dtype=torch.int32, device=device
)
blk_offset += n_blks
q_bf16 = torch.randn(batch_size, next_n, num_heads, HEAD_DIM, device=device)
q_fp8 = q_bf16.to(torch.float8_e4m3fn)
kv_bf16 = torch.randn(num_phys_blocks, BLOCK_KV, HEAD_DIM, device=device)
kv_amax = kv_bf16.abs().float().amax(dim=-1, keepdim=True).clamp(1e-4)
kv_scale = fp8_mqa_logits_ceil_to_ue8m0(kv_amax / 448.0).squeeze(-1)
kv_fp8 = (kv_bf16 / kv_scale.unsqueeze(-1)).to(torch.float8_e4m3fn)
weights = torch.randn(
batch_size * next_n, num_heads, device=device, dtype=torch.float32
)
kv_fused = fp8_mqa_logits_make_fused_kv(kv_fp8, kv_scale, BLOCK_KV, HEAD_DIM)
return {
"q_fp8": q_fp8,
"kv_fp8": kv_fp8,
"kv_scales": kv_scale,
"kv_fused": kv_fused,
"weights": weights,
"context_lens": context_lens.to(device),
"block_table": block_table,
}
def _assert_matches_ref(logits, ref_logits, context_lens, B, next_n, max_model_len):
device = logits.device
positions = torch.arange(max_model_len, device=device).unsqueeze(0)
row_indices = torch.arange(B * next_n, device=device) // next_n
next_n_offset = torch.arange(B * next_n, device=device) % next_n
end_pos = context_lens[row_indices] - next_n + next_n_offset
mask = positions <= end_pos.unsqueeze(1)
logits_masked = logits.float().masked_fill(~mask, 0)
ref_masked = ref_logits.float().masked_fill(~mask, 0)
torch.testing.assert_close(logits_masked, ref_masked, atol=5e-5, rtol=1e-5)
def _run_cutedsl_paged_mqa_logits(
data, batch_size, next_n, num_heads, max_model_len, is_target_verify
):
"""Mirrors the CUTEDSL dispatch in
sglang.srt.layers.attention.dsa.dsa_indexer.Indexer._get_topk_paged."""
import deep_gemm
num_sms = torch.cuda.get_device_properties(0).multi_processor_count
if is_target_verify and next_n >= 2:
dsl_expand_factor, dsl_atom = pick_dsl_expand(
next_n,
batch_size=batch_size,
max_ctx=max_model_len,
num_sms=num_sms,
kernel_atoms=(1, 2, 3, 4),
num_heads=num_heads,
)
else:
dsl_expand_factor, dsl_atom = 1, 1
context_lens = data["context_lens"]
expanded_ctx = (
context_lens.unsqueeze(-1)
- next_n
+ torch.arange(1, next_n + 1, device=context_lens.device, dtype=torch.int32)
).flatten()
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
expanded_ctx.unsqueeze(-1), BLOCK_KV, num_sms
)
block_tables_expanded = data["block_table"].repeat_interleave(next_n, dim=0)
return cutedsl_paged_mqa_logits(
data["q_fp8"].view(batch_size * next_n, num_heads, HEAD_DIM),
data["kv_fused"],
data["weights"],
context_lens,
block_tables_expanded,
schedule_metadata,
max_model_len,
q_offset=batch_size * next_n,
B=batch_size,
next_n=next_n,
is_target_verify=is_target_verify,
dsl_expand_factor=dsl_expand_factor,
dsl_atom=dsl_atom,
blocksize=BLOCK_KV,
sm_count=num_sms,
get_paged_mqa_logits_metadata_fn=deep_gemm.get_paged_mqa_logits_metadata,
)
@pytest.mark.skipif(
not is_sm100_supported(),
reason="CuTe DSL FP8 Paged MQA Logits only supports SM 100 family.",
)
@pytest.mark.parametrize("batch_size", [1, 2, 4, 8])
@pytest.mark.parametrize("next_n", [1, 2, 3, 4, 5, 6])
@pytest.mark.parametrize("num_heads", [32, 64])
@pytest.mark.parametrize("avg_ctx", [128, 1024, 4096, 16384])
def test_cutedsl_paged_mqa_logits(batch_size, next_n, num_heads, avg_ctx):
max_model_len = max(avg_ctx * 2, 2048)
data = _generate_test_data(batch_size, next_n, num_heads, avg_ctx, max_model_len)
logits = _run_cutedsl_paged_mqa_logits(
data,
batch_size,
next_n,
num_heads,
max_model_len,
is_target_verify=next_n >= 2,
)
ref_logits = _ref_fp8_paged_mqa_logits(
data["q_fp8"],
data["kv_fp8"],
data["kv_scales"],
data["weights"],
data["context_lens"],
data["block_table"],
max_model_len,
BLOCK_KV,
)
_assert_matches_ref(
logits, ref_logits, data["context_lens"], batch_size, next_n, max_model_len
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -0,0 +1,233 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import sys
import pytest
import torch
from sglang.jit_kernel.dsa import (
deepgemm_paged_mqa_logits_native,
deepgemm_paged_mqa_logits_split,
)
from sglang.srt.layers.attention.dsa.utils import (
fp8_mqa_logits_ceil_to_ue8m0,
fp8_mqa_logits_make_fused_kv,
)
from sglang.srt.utils import is_sm90_supported, is_sm100_supported
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=40, suite="nightly-4-gpu-b200", nightly=True)
BLOCK_KV = 64
HEAD_DIM = 128
def _ref_fp8_paged_mqa_logits(
q_fp8,
kv_fp8,
kv_scales,
weights,
context_lens,
block_table,
max_model_len,
block_kv,
):
B, next_n, H, D = q_fp8.shape
device = q_fp8.device
logits = torch.full(
(B * next_n, max_model_len), float("-inf"), device=device, dtype=torch.float32
)
q_f32 = q_fp8.float()
for b in range(B):
ctx_len = context_lens[b].item()
q_positions = torch.arange(ctx_len - next_n, ctx_len, device=device)
w = weights[b * next_n : (b + 1) * next_n, :]
for blk_idx in range((ctx_len + block_kv - 1) // block_kv):
phys_blk = block_table[b, blk_idx].item()
k_f32 = kv_fp8[phys_blk].float()
scales = kv_scales[phys_blk]
k_positions = torch.arange(
blk_idx * block_kv, (blk_idx + 1) * block_kv, device=device
)
mask = (k_positions[None, :] < ctx_len) & (
k_positions[None, :] <= q_positions[:, None]
)
qk = torch.matmul(q_f32[b].permute(1, 0, 2), k_f32.T)
qk = torch.where(mask[None, :, :], qk, torch.zeros(1, device=device))
qk = torch.relu(qk)
weighted = (w.T[:, :, None] * qk).sum(dim=0)
weighted = weighted * scales[None, :]
start_pos = blk_idx * block_kv
end_pos = start_pos + block_kv
logits[b * next_n : (b + 1) * next_n, start_pos:end_pos] = torch.where(
mask,
weighted,
torch.tensor(float("-inf"), device=device, dtype=torch.float32),
)
return logits
def _generate_test_data(
batch_size,
next_n,
num_heads,
avg_context_len,
max_model_len,
device="cuda",
):
torch.manual_seed(42)
torch.cuda.manual_seed(42)
context_lens = torch.randint(
max(BLOCK_KV, int(0.7 * avg_context_len)),
int(1.3 * avg_context_len) + 1,
(batch_size,),
dtype=torch.int32,
device="cpu",
).clamp(max=max_model_len)
max_blocks_per_seq = (max_model_len + BLOCK_KV - 1) // BLOCK_KV
total_blocks = ((context_lens + BLOCK_KV - 1) // BLOCK_KV).sum().item()
num_phys_blocks = total_blocks + batch_size * 2
block_table = torch.full(
(batch_size, max_blocks_per_seq), 0, dtype=torch.int32, device=device
)
blk_offset = 0
for i in range(batch_size):
n_blks = (context_lens[i].item() + BLOCK_KV - 1) // BLOCK_KV
block_table[i, :n_blks] = torch.arange(
blk_offset, blk_offset + n_blks, dtype=torch.int32, device=device
)
blk_offset += n_blks
q_bf16 = torch.randn(batch_size, next_n, num_heads, HEAD_DIM, device=device)
q_fp8 = q_bf16.to(torch.float8_e4m3fn)
kv_bf16 = torch.randn(num_phys_blocks, BLOCK_KV, HEAD_DIM, device=device)
kv_amax = kv_bf16.abs().float().amax(dim=-1, keepdim=True).clamp(1e-4)
kv_scale = fp8_mqa_logits_ceil_to_ue8m0(kv_amax / 448.0).squeeze(-1)
kv_fp8 = (kv_bf16 / kv_scale.unsqueeze(-1)).to(torch.float8_e4m3fn)
weights = torch.randn(
batch_size * next_n, num_heads, device=device, dtype=torch.float32
)
kv_fused = fp8_mqa_logits_make_fused_kv(kv_fp8, kv_scale, BLOCK_KV, HEAD_DIM)
return {
"q_fp8": q_fp8,
"kv_fp8": kv_fp8,
"kv_scales": kv_scale,
"kv_fused": kv_fused,
"weights": weights,
"context_lens": context_lens.to(device),
"block_table": block_table,
}
def _assert_matches_ref(logits, ref_logits, context_lens, B, next_n, max_model_len):
device = logits.device
positions = torch.arange(max_model_len, device=device).unsqueeze(0)
row_indices = torch.arange(B * next_n, device=device) // next_n
next_n_offset = torch.arange(B * next_n, device=device) % next_n
end_pos = context_lens[row_indices] - next_n + next_n_offset
mask = positions <= end_pos.unsqueeze(1)
logits_masked = logits.float().masked_fill(~mask, 0)
ref_masked = ref_logits.float().masked_fill(~mask, 0)
torch.testing.assert_close(logits_masked, ref_masked, atol=5e-5, rtol=1e-5)
def _run_deepgemm_paged_mqa_logits(data, batch_size, next_n, num_heads, max_model_len):
"""Mirrors the DEEPGEMM dispatch in
sglang.srt.layers.attention.dsa.dsa_indexer.Indexer._get_topk_paged:
next_n>=2 (target-verify) goes through the native wrapper, everything
else goes through the split wrapper."""
import deep_gemm
num_sms = torch.cuda.get_device_properties(0).multi_processor_count
if next_n >= 2:
ctx_lens_2d = (
data["context_lens"].unsqueeze(-1)
- next_n
+ torch.arange(
1, next_n + 1, device=data["context_lens"].device, dtype=torch.int32
)
)
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
ctx_lens_2d, BLOCK_KV, num_sms
)
block_tables_expanded = data["block_table"].repeat_interleave(next_n, dim=0)
return deepgemm_paged_mqa_logits_native(
deep_gemm.fp8_paged_mqa_logits,
data["q_fp8"].view(batch_size * next_n, num_heads, HEAD_DIM),
data["kv_fused"],
data["weights"],
ctx_lens_2d,
block_tables_expanded,
schedule_metadata,
max_model_len,
q_offset=batch_size * next_n,
B=batch_size,
next_n=next_n,
)
ctx_lens_2d = data["context_lens"].unsqueeze(-1)
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
ctx_lens_2d, BLOCK_KV, num_sms
)
return deepgemm_paged_mqa_logits_split(
deep_gemm.fp8_paged_mqa_logits,
data["q_fp8"].squeeze(1),
data["kv_fused"],
data["weights"],
ctx_lens_2d,
data["block_table"],
schedule_metadata,
max_model_len,
q_offset=batch_size,
)
@pytest.mark.skipif(
not (is_sm90_supported() or is_sm100_supported()),
reason="DeepGEMM fp8_paged_mqa_logits requires SM90 (Hopper) or newer.",
)
@pytest.mark.parametrize("batch_size", [1, 2, 4, 8])
@pytest.mark.parametrize("next_n", [1, 2, 3, 4, 5, 6])
@pytest.mark.parametrize("num_heads", [32, 64])
@pytest.mark.parametrize("avg_ctx", [128, 1024, 4096, 16384])
def test_deepgemm_paged_mqa_logits(batch_size, next_n, num_heads, avg_ctx):
max_model_len = max(avg_ctx * 2, 2048)
data = _generate_test_data(batch_size, next_n, num_heads, avg_ctx, max_model_len)
logits = _run_deepgemm_paged_mqa_logits(
data, batch_size, next_n, num_heads, max_model_len
)
ref_logits = _ref_fp8_paged_mqa_logits(
data["q_fp8"],
data["kv_fp8"],
data["kv_scales"],
data["weights"],
data["context_lens"],
data["block_table"],
max_model_len,
BLOCK_KV,
)
_assert_matches_ref(
logits, ref_logits, data["context_lens"], batch_size, next_n, max_model_len
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -260,6 +260,7 @@ class MockModelRunner:
"dsa_prefill_backend": "flashmla_sparse",
"dsa_decode_backend": "fa3",
"dsa_topk_backend": "sgl-kernel",
"dsa_paged_mqa_logits_backend": "auto",
},
)()
self.hisparse_coordinator = None