[AMD] perf: compact Triton extend-attention for ragged prefill (AMD/HIP-only) (#29677)

This commit is contained in:
valechen
2026-08-06 14:46:10 -07:00
committed by GitHub
parent dd7e4c91e2
commit 18e6c61c21
9 changed files with 449 additions and 6 deletions
@@ -10,6 +10,7 @@ from sglang.kernels.ops.attention.decode_attention import (
decode_attention_fwd_normal,
)
from sglang.kernels.ops.attention.extend_attention import (
_compact_extend_q_tiles_per_head,
build_unified_kv_indices,
extend_attention_fwd,
extend_attention_fwd_unified,
@@ -19,6 +20,7 @@ from sglang.kernels.ops.attention.prefill_attention import (
context_attention_fwd,
)
from sglang.srt.utils import get_device
from sglang.srt.utils.common import temp_set_env
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase, is_in_amd_ci
@@ -338,6 +340,131 @@ class TestTritonAttention(CustomTestCase):
ea._get_block_sizes_for_extend_attention(576, 576)[3:], (64, 64, 4)
)
def test_compact_extend_attention_tile_count(self):
self.assertEqual(
_compact_extend_q_tiles_per_head(
batch_size=16,
max_len_extend=1000,
total_extend_tokens=1015,
block_m=64,
extend_seq_lens_cpu=[1] * 15 + [1000],
),
31,
)
self.assertEqual(
_compact_extend_q_tiles_per_head(
batch_size=2,
max_len_extend=4224,
total_extend_tokens=5376,
block_m=64,
extend_seq_lens_cpu=[1152, 4224],
),
84,
)
self.assertIsNone(
_compact_extend_q_tiles_per_head(
batch_size=4,
max_len_extend=64,
total_extend_tokens=256,
block_m=64,
extend_seq_lens_cpu=[64, 64, 64, 64],
)
)
def test_extend_attention_compact_grid(self):
dtype = torch.bfloat16
device = get_device()
B, H_Q, H_KV, D = 4, 8, 2, 64
b_seq_len_prefix = torch.tensor(
[8, 16, 32, 64], dtype=torch.int32, device=device
)
b_seq_len_extend = torch.tensor(
[1, 7, 13, 129], dtype=torch.int32, device=device
)
b_seq_len = b_seq_len_prefix + b_seq_len_extend
b_start_loc = torch.zeros((B,), dtype=torch.int32, device=device)
b_start_loc[1:] = torch.cumsum(b_seq_len[:-1], 0)
b_start_loc_extend = torch.zeros((B,), dtype=torch.int32, device=device)
b_start_loc_extend[1:] = torch.cumsum(b_seq_len_extend[:-1], 0)
kv_indptr = torch.zeros((B + 1,), dtype=torch.int32, device=device)
kv_indptr[1 : B + 1] = torch.cumsum(b_seq_len_prefix, dim=0)
kv_indices = torch.empty(
(int(b_seq_len_prefix.sum().item()),), dtype=torch.int32, device=device
)
for i in range(B):
kv_indices[int(kv_indptr[i]) : int(kv_indptr[i + 1])] = torch.arange(
int(b_start_loc[i].item()),
int(b_start_loc[i].item()) + int(b_seq_len_prefix[i].item()),
device=device,
)
total_token_num = int(b_seq_len.sum().item())
extend_token_num = int(b_seq_len_extend.sum().item())
k_buffer = torch.empty(
(total_token_num, H_KV, D), dtype=dtype, device=device
).normal_(mean=0.1, std=0.2)
v_buffer = torch.empty(
(total_token_num, H_KV, D), dtype=dtype, device=device
).normal_(mean=0.1, std=0.2)
k_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype, device=device)
v_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype, device=device)
q_extend = torch.empty((extend_token_num, H_Q, D), dtype=dtype, device=device)
for i in range(B):
extend_start_in_buffer = b_start_loc[i] + b_seq_len_prefix[i]
extend_end_in_buffer = b_start_loc[i] + b_seq_len[i]
extend_start = b_start_loc_extend[i]
extend_end = b_start_loc_extend[i] + b_seq_len_extend[i]
k_extend[extend_start:extend_end] = k_buffer[
extend_start_in_buffer:extend_end_in_buffer
]
v_extend[extend_start:extend_end] = v_buffer[
extend_start_in_buffer:extend_end_in_buffer
]
q_extend[extend_start:extend_end] = torch.empty(
(int(b_seq_len_extend[i].item()), H_Q, D),
dtype=dtype,
device=device,
).normal_(mean=0.1, std=0.2)
max_len_extend = int(b_seq_len_extend.max().item())
qo_indptr = torch.zeros((B + 1,), dtype=torch.int32, device=device)
qo_indptr[1 : B + 1] = torch.cumsum(b_seq_len_extend, dim=0)
extend_seq_lens_cpu = b_seq_len_extend.cpu().tolist()
o_legacy = torch.empty_like(q_extend)
o_compact = torch.empty_like(q_extend)
for output, use_compact in ((o_legacy, False), (o_compact, True)):
with temp_set_env(
allow_sglang=True,
SGLANG_TRITON_COMPACT_EXTEND_ATTENTION=str(use_compact),
):
extend_attention_fwd(
q_extend,
k_extend,
v_extend,
output,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask=None,
is_causal=True,
mask_indptr=None,
max_len_extend=max_len_extend,
k_scale=1.0,
v_scale=1.0,
extend_seq_lens_cpu=extend_seq_lens_cpu,
)
self.assertTrue(
torch.allclose(o_legacy, o_compact, rtol=1e-2, atol=1e-3),
f"compact grid output differs from legacy grid. "
f"Max diff: {(o_legacy - o_compact).abs().max()}",
)
def _test_extend_attention_sliding_window_once(
self, B, N_CTX, H_Q, H_KV, D, WINDOW_SIZE
):
@@ -234,6 +234,20 @@ class TestVerifySplitKV(CustomTestCase):
can_handle(q, k, v, kb, vb, qo, kvp, kvi, None, True, None, mle + 1)
)
def test_fallback_mla_head_dim_mismatch(self):
# MLA (DeepSeek) has head_dim != v_head_dim (576 vs 512): the shared
# latent KV / absorbed layout is not something the split-KV verify
# kernel is built for -- it GPU-faults on that shape. can_handle() must
# reject it so the backend falls back to extend_attention_fwd.
q, k, v, kb, vb, qo, kvp, kvi, mle = _build_verify_inputs(
[512, 512], 4, 16, 1, 576, 512, torch.bfloat16, "cuda"
)
self.assertEqual(q.shape[2], 576)
self.assertEqual(v.shape[2], 512)
self.assertFalse(
can_handle(q, k, v, kb, vb, qo, kvp, kvi, None, True, None, mle)
)
if __name__ == "__main__":
unittest.main()
@@ -2,8 +2,13 @@ import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import sglang.srt.managers.schedule_policy as schedule_policy
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder
from sglang.srt.managers.schedule_policy import (
AddReqResult,
PrefillAdder,
estimate_prefill_extend_tile_metrics,
)
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefResult,
IncLockRefResult,
@@ -775,6 +780,66 @@ class TestPrefillAdder(CustomTestCase):
req.set_extend_range.assert_called_once_with(0, 200)
self.assertIn(req, adder.can_run_list)
def _adder_with_extend_lens(self, extend_lens):
adder = PrefillAdder.__new__(PrefillAdder)
adder.can_run_list = [
SimpleNamespace(extend_input_len=length) for length in extend_lens
]
# BLOCK_M is auto-detected from the attention backend in production; the
# __new__ helper bypasses __init__, so set it explicitly. 64 matches the
# block_m the tile-count assertions below are computed against.
adder.prefill_tile_block_m = 64
return adder
def test_estimate_prefill_extend_tile_metrics(self):
metrics = estimate_prefill_extend_tile_metrics([1, 7, 13, 129], block_m=64)
self.assertEqual(metrics["q_tiles_per_request"], [1, 1, 1, 3])
self.assertEqual(metrics["legacy_q_tiles_per_head"], 12)
self.assertEqual(metrics["compact_q_tiles_per_head"], 6)
self.assertEqual(metrics["saved_q_tiles_per_head"], 6)
self.assertEqual(metrics["saved_q_tile_ratio"], 0.5)
def test_compact_prefill_tile_budget_admits_more_than_legacy(self):
adder = self._adder_with_extend_lens([1, 7, 13])
# The tile-budget admission is gated on HIP in production; force the gate
# on so this vendor-neutral admission-math check runs on any CI runner.
with (
patch.object(schedule_policy, "_IS_HIP", True),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET", 6),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET_MODE", "compact"),
):
self.assertIsNone(adder._check_prefill_tile_budget(129))
with (
patch.object(schedule_policy, "_IS_HIP", True),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET", 6),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET_MODE", "legacy"),
):
self.assertEqual(adder._check_prefill_tile_budget(129), AddReqResult.OTHER)
def test_prefill_tile_budget_always_allows_first_request(self):
adder = self._adder_with_extend_lens([])
with (
patch.object(schedule_policy, "_IS_HIP", True),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET", 1),
):
self.assertIsNone(adder._check_prefill_tile_budget(4096))
def test_prefill_tile_budget_disabled_on_non_hip(self):
# AMD-only: on non-HIP vendors the tile-budget admission must be a no-op
# even when the budget env is set, so scheduler behavior is unchanged.
adder = self._adder_with_extend_lens([1, 7, 13])
with (
patch.object(schedule_policy, "_IS_HIP", False),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET", 6),
patch.object(schedule_policy, "PREFILL_TILE_BUDGET_MODE", "legacy"),
):
self.assertIsNone(adder._check_prefill_tile_budget(129))
if __name__ == "__main__":
unittest.main()