forked from Karylab-cklius/vllm
[ROCm][DSv4][Perf] Flash-decode split-K decode attention kernel (#44899)
Co-authored-by: vLLM Contributor <contributor@vllm.ai>
This commit is contained in:
co-authored by
vLLM Contributor
parent
4bc83323f2
commit
fcf5115c45
@@ -10,6 +10,25 @@ pytestmark = pytest.mark.skipif(
|
||||
not current_platform.is_rocm(), reason="Only used by ROCm"
|
||||
)
|
||||
|
||||
|
||||
def _on_gfx950() -> bool:
|
||||
if not current_platform.is_rocm():
|
||||
return False
|
||||
try:
|
||||
from vllm.platforms.rocm import _ON_GFX950
|
||||
|
||||
return bool(_ON_GFX950)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# The flash-decode split-K decode path is only tuned for AMD gfx950; other
|
||||
# architectures take the fallback decode kernel, so its tests are skipped there.
|
||||
requires_gfx950 = pytest.mark.skipif(
|
||||
not _on_gfx950(),
|
||||
reason="split-K decode kernel is only tuned for AMD gfx950",
|
||||
)
|
||||
|
||||
NOPE_HEAD_DIM = 448
|
||||
ROPE_HEAD_DIM = 64
|
||||
HEAD_DIM = NOPE_HEAD_DIM + ROPE_HEAD_DIM
|
||||
@@ -156,6 +175,20 @@ def _ref_sparse_decode_ragged(
|
||||
return out.to(torch.bfloat16)
|
||||
|
||||
|
||||
def _ragged_from_rows(
|
||||
rows: list[list[int]], device: torch.device
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Flatten per-query slot lists into ragged (indices, indptr) tensors."""
|
||||
flat = [slot for row in rows for slot in row]
|
||||
indptr = [0]
|
||||
for row in rows:
|
||||
indptr.append(indptr[-1] + len(row))
|
||||
return (
|
||||
torch.tensor(flat, dtype=torch.int32, device=device),
|
||||
torch.tensor(indptr, dtype=torch.int32, device=device),
|
||||
)
|
||||
|
||||
|
||||
def _ref_combine_topk_swa_ragged(
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
@@ -375,3 +408,110 @@ def test_combine_topk_swa_indices_ragged() -> None:
|
||||
)
|
||||
torch.testing.assert_close(actual_indptr, expected_indptr)
|
||||
torch.testing.assert_close(actual_lens, expected_lens)
|
||||
|
||||
|
||||
@requires_gfx950
|
||||
@torch.inference_mode()
|
||||
def test_decode_num_splits_heuristic(monkeypatch) -> None:
|
||||
"""Split-count heuristic added with the flash-decode split-K decode path."""
|
||||
from vllm.v1.attention.ops import rocm_aiter_mla_sparse as mod
|
||||
|
||||
# Pin the CU count so the heuristic is deterministic off-device.
|
||||
monkeypatch.setattr(mod, "_decode_cu_count", lambda: 256)
|
||||
|
||||
# A batch that already fills the device should not be split.
|
||||
assert mod._decode_num_splits(256, 1, avg_main_len=128.0, avg_extra_len=0.0) == 1
|
||||
# A tiny batch on a large device should split to add parallelism.
|
||||
assert mod._decode_num_splits(2, 1, avg_main_len=256.0, avg_extra_len=0.0) > 1
|
||||
|
||||
# The chosen count always stays within the searched [1, 16] range, and a
|
||||
# zero-length workload never splits (no work to parallelize).
|
||||
for num_queries in (1, 4, 24, 224, 1024):
|
||||
splits = mod._decode_num_splits(
|
||||
num_queries, 1, avg_main_len=512.0, avg_extra_len=128.0
|
||||
)
|
||||
assert 1 <= splits <= 16
|
||||
assert mod._decode_num_splits(2, 1, avg_main_len=0.0, avg_extra_len=0.0) >= 1
|
||||
|
||||
|
||||
@requires_gfx950
|
||||
@pytest.mark.parametrize("num_splits", [1, 2, 3, 4, 8])
|
||||
@pytest.mark.parametrize("with_extra", [True, False])
|
||||
@pytest.mark.parametrize("with_sink", [True, False])
|
||||
@torch.inference_mode()
|
||||
def test_sparse_attn_decode_split_k_kernel(
|
||||
monkeypatch, num_splits: int, with_extra: bool, with_sink: bool
|
||||
) -> None:
|
||||
"""Flash-decode split-K decode path (partial + reduce kernels).
|
||||
|
||||
This path is the gfx950 production path (``_ON_GFX950``), so the test only
|
||||
runs on gfx950. The split count is pinned so the partial/reduce kernels are
|
||||
exercised across split counts. ``num_splits=8`` drives splits past the
|
||||
shortest segment length, covering the empty-split edge case handled by the
|
||||
reduce kernel.
|
||||
"""
|
||||
from vllm.v1.attention.ops import rocm_aiter_mla_sparse as mod
|
||||
|
||||
device = torch.device("cuda")
|
||||
torch.manual_seed(7)
|
||||
block_size = 4
|
||||
num_heads = 3
|
||||
|
||||
main_rows = [[0, 2, 4, 6, 1, 3, 7, 5], [4, 1, 6, 0, 2]]
|
||||
num_queries = len(main_rows)
|
||||
q = (
|
||||
torch.randn(
|
||||
num_queries, num_heads, HEAD_DIM, dtype=torch.bfloat16, device=device
|
||||
)
|
||||
* 0.125
|
||||
)
|
||||
main_kv = torch.randn(8, HEAD_DIM, dtype=torch.bfloat16, device=device) * 0.125
|
||||
main_cache = _pack_fp8_ds_mla_cache(main_kv, block_size)
|
||||
main_indices, main_indptr = _ragged_from_rows(main_rows, device)
|
||||
|
||||
extra_rows: list[list[int]] | None = None
|
||||
extra_cache: torch.Tensor | None = None
|
||||
extra_indices: torch.Tensor | None = None
|
||||
extra_indptr: torch.Tensor | None = None
|
||||
if with_extra:
|
||||
rows = [[1, 3, 0, 5, 2, 4], [3, 0, 6]]
|
||||
extra_kv = torch.randn(7, HEAD_DIM, dtype=torch.bfloat16, device=device) * 0.125
|
||||
extra_rows = rows
|
||||
extra_cache = _pack_fp8_ds_mla_cache(extra_kv, block_size)
|
||||
extra_indices, extra_indptr = _ragged_from_rows(rows, device)
|
||||
|
||||
attn_sink = (
|
||||
torch.tensor([-0.1, 0.0, 0.1], dtype=torch.float32, device=device)
|
||||
if with_sink
|
||||
else None
|
||||
)
|
||||
scale = HEAD_DIM**-0.5
|
||||
|
||||
# Pin the split count so each parametrized value is exercised deterministically.
|
||||
monkeypatch.setattr(mod, "_decode_num_splits", lambda *args, **kwargs: num_splits)
|
||||
|
||||
actual = mod._rocm_sparse_attn_decode_ragged_triton(
|
||||
q=q,
|
||||
main_cache=main_cache,
|
||||
main_indices=main_indices,
|
||||
main_indptr=main_indptr,
|
||||
scale=scale,
|
||||
attn_sink=attn_sink,
|
||||
nope_head_dim=NOPE_HEAD_DIM,
|
||||
rope_head_dim=ROPE_HEAD_DIM,
|
||||
extra_cache=extra_cache,
|
||||
extra_indices=extra_indices,
|
||||
extra_indptr=extra_indptr,
|
||||
)
|
||||
expected = _ref_sparse_decode_ragged(
|
||||
q=q,
|
||||
main_cache=main_cache,
|
||||
main_rows=main_rows,
|
||||
scale=scale,
|
||||
attn_sink=attn_sink,
|
||||
block_size=block_size,
|
||||
extra_cache=extra_cache,
|
||||
extra_rows=extra_rows,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2)
|
||||
|
||||
@@ -1406,6 +1406,348 @@ def _sparse_attn_decode_ragged_kernel(
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _sparse_attn_decode_partial_kernel(
|
||||
q_ptr,
|
||||
main_cache_ptr,
|
||||
main_indices_ptr,
|
||||
main_indptr_ptr,
|
||||
extra_cache_ptr,
|
||||
extra_indices_ptr,
|
||||
extra_indptr_ptr,
|
||||
part_m_ptr,
|
||||
part_l_ptr,
|
||||
part_acc_ptr,
|
||||
q_stride0,
|
||||
q_stride1,
|
||||
main_cache_stride0,
|
||||
extra_cache_stride0,
|
||||
pm_stride0,
|
||||
pm_stride_s,
|
||||
pa_stride0,
|
||||
pa_stride_s,
|
||||
pa_stride_h,
|
||||
main_num_rows,
|
||||
extra_num_rows,
|
||||
main_block_size,
|
||||
extra_block_size,
|
||||
scale,
|
||||
num_heads,
|
||||
HAS_EXTRA: tl.constexpr,
|
||||
NOPE_DIM: tl.constexpr,
|
||||
NOPE_BLOCK: tl.constexpr,
|
||||
ROPE_DIM: tl.constexpr,
|
||||
IS_FNUZ: tl.constexpr,
|
||||
BLOCK_H: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
NUM_SPLITS: tl.constexpr,
|
||||
NUM_STAGES: tl.constexpr,
|
||||
):
|
||||
query_idx = tl.program_id(0)
|
||||
split_id = tl.program_id(1)
|
||||
pid_h = tl.program_id(2)
|
||||
|
||||
head_offsets = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
|
||||
head_mask = head_offsets < num_heads
|
||||
nope_offsets = tl.arange(0, NOPE_BLOCK)
|
||||
nope_mask = nope_offsets < NOPE_DIM
|
||||
rope_offsets = tl.arange(0, ROPE_DIM)
|
||||
|
||||
q_row_ptr = q_ptr + query_idx * q_stride0 + head_offsets[:, None] * q_stride1
|
||||
q_nope = tl.load(
|
||||
q_row_ptr + nope_offsets[None, :],
|
||||
mask=head_mask[:, None] & nope_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
q_rope = tl.load(
|
||||
q_row_ptr + NOPE_DIM + rope_offsets[None, :],
|
||||
mask=head_mask[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
neg_large = -3.4028234663852886e38
|
||||
m_i = tl.full((BLOCK_H,), neg_large, dtype=tl.float32)
|
||||
l_i = tl.zeros((BLOCK_H,), dtype=tl.float32)
|
||||
acc_nope = tl.zeros((BLOCK_H, NOPE_BLOCK), dtype=tl.float32)
|
||||
acc_rope = tl.zeros((BLOCK_H, ROPE_DIM), dtype=tl.float32)
|
||||
k_offsets = tl.arange(0, BLOCK_K)
|
||||
|
||||
zero_nope = tl.zeros((BLOCK_K, NOPE_BLOCK), dtype=tl.bfloat16)
|
||||
zero_rope = tl.zeros((BLOCK_K, ROPE_DIM), dtype=tl.bfloat16)
|
||||
|
||||
# Each split processes a contiguous slice of this query's main (SWA) and
|
||||
# extra (topk) segments. Slices are handled independently so a block never
|
||||
# straddles the main/extra boundary.
|
||||
main_start = tl.load(main_indptr_ptr + query_idx)
|
||||
main_end = tl.load(main_indptr_ptr + query_idx + 1)
|
||||
main_len = main_end - main_start
|
||||
main_chunk = (main_len + NUM_SPLITS - 1) // NUM_SPLITS
|
||||
main_lo = split_id * main_chunk
|
||||
main_hi = tl.minimum(main_lo + main_chunk, main_len)
|
||||
|
||||
for k_start in tl.range(main_lo, main_hi, BLOCK_K, num_stages=NUM_STAGES):
|
||||
k_pos = k_start + k_offsets
|
||||
in_range = k_pos < main_hi
|
||||
slot = tl.load(main_indices_ptr + main_start + k_pos, mask=in_range, other=-1)
|
||||
valid = in_range & (slot >= 0) & (slot < main_num_rows)
|
||||
safe_slot = tl.where(valid, slot, 0)
|
||||
|
||||
block_idx = safe_slot // main_block_size
|
||||
pos_in_block = safe_slot % main_block_size
|
||||
cache_block_ptr = main_cache_ptr + block_idx.to(tl.int64) * main_cache_stride0
|
||||
token_data_ptr = cache_block_ptr + pos_in_block * 576
|
||||
token_scale_ptr = cache_block_ptr + main_block_size * 576 + pos_in_block * 8
|
||||
|
||||
x_uint8 = tl.load(
|
||||
token_data_ptr[:, None] + nope_offsets[None, :],
|
||||
mask=valid[:, None] & nope_mask[None, :],
|
||||
other=0,
|
||||
)
|
||||
if IS_FNUZ:
|
||||
x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True)
|
||||
else:
|
||||
x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True)
|
||||
encoded_scales = tl.load(
|
||||
token_scale_ptr[:, None] + nope_offsets[None, :] // 64,
|
||||
mask=valid[:, None] & nope_mask[None, :],
|
||||
other=127,
|
||||
)
|
||||
scales = tl.exp2(encoded_scales.to(tl.float32) - 127.0)
|
||||
k_nope = x_fp8.to(tl.bfloat16) * scales.to(tl.bfloat16)
|
||||
k_nope = tl.where(valid[:, None] & nope_mask[None, :], k_nope, zero_nope)
|
||||
k_nope = tl.where(k_nope == k_nope, k_nope, zero_nope)
|
||||
|
||||
rope_ptr = (token_data_ptr + NOPE_DIM).to(tl.pointer_type(tl.bfloat16))
|
||||
k_rope = tl.load(
|
||||
rope_ptr[:, None] + rope_offsets[None, :],
|
||||
mask=valid[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
k_rope = tl.where(valid[:, None], k_rope, zero_rope)
|
||||
k_rope = tl.where(k_rope == k_rope, k_rope, zero_rope)
|
||||
|
||||
scores = tl.dot(q_nope, tl.trans(k_nope)) + tl.dot(q_rope, tl.trans(k_rope))
|
||||
scores *= scale
|
||||
scores = tl.where(head_mask[:, None] & valid[None, :], scores, neg_large)
|
||||
|
||||
m_block = tl.max(scores, axis=1)
|
||||
m_new = tl.maximum(m_i, m_block)
|
||||
alpha = tl.exp(m_i - m_new)
|
||||
p = tl.exp(scores - m_new[:, None])
|
||||
p = tl.where(head_mask[:, None] & valid[None, :], p, 0.0)
|
||||
l_new = l_i * alpha + tl.sum(p, axis=1)
|
||||
|
||||
acc_nope = acc_nope * alpha[:, None] + tl.dot(p.to(k_nope.dtype), k_nope)
|
||||
acc_rope = acc_rope * alpha[:, None] + tl.dot(p.to(k_rope.dtype), k_rope)
|
||||
m_i = m_new
|
||||
l_i = l_new
|
||||
|
||||
if HAS_EXTRA:
|
||||
extra_start = tl.load(extra_indptr_ptr + query_idx)
|
||||
extra_end = tl.load(extra_indptr_ptr + query_idx + 1)
|
||||
extra_len = extra_end - extra_start
|
||||
extra_chunk = (extra_len + NUM_SPLITS - 1) // NUM_SPLITS
|
||||
extra_lo = split_id * extra_chunk
|
||||
extra_hi = tl.minimum(extra_lo + extra_chunk, extra_len)
|
||||
|
||||
for k_start in tl.range(extra_lo, extra_hi, BLOCK_K, num_stages=NUM_STAGES):
|
||||
k_pos = k_start + k_offsets
|
||||
in_range = k_pos < extra_hi
|
||||
slot = tl.load(
|
||||
extra_indices_ptr + extra_start + k_pos, mask=in_range, other=-1
|
||||
)
|
||||
valid = in_range & (slot >= 0) & (slot < extra_num_rows)
|
||||
safe_slot = tl.where(valid, slot, 0)
|
||||
|
||||
block_idx = safe_slot // extra_block_size
|
||||
pos_in_block = safe_slot % extra_block_size
|
||||
cache_block_ptr = (
|
||||
extra_cache_ptr + block_idx.to(tl.int64) * extra_cache_stride0
|
||||
)
|
||||
token_data_ptr = cache_block_ptr + pos_in_block * 576
|
||||
token_scale_ptr = (
|
||||
cache_block_ptr + extra_block_size * 576 + pos_in_block * 8
|
||||
)
|
||||
|
||||
x_uint8 = tl.load(
|
||||
token_data_ptr[:, None] + nope_offsets[None, :],
|
||||
mask=valid[:, None] & nope_mask[None, :],
|
||||
other=0,
|
||||
)
|
||||
if IS_FNUZ:
|
||||
x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True)
|
||||
else:
|
||||
x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True)
|
||||
encoded_scales = tl.load(
|
||||
token_scale_ptr[:, None] + nope_offsets[None, :] // 64,
|
||||
mask=valid[:, None] & nope_mask[None, :],
|
||||
other=127,
|
||||
)
|
||||
scales = tl.exp2(encoded_scales.to(tl.float32) - 127.0)
|
||||
k_nope = x_fp8.to(tl.bfloat16) * scales.to(tl.bfloat16)
|
||||
k_nope = tl.where(valid[:, None] & nope_mask[None, :], k_nope, zero_nope)
|
||||
k_nope = tl.where(k_nope == k_nope, k_nope, zero_nope)
|
||||
|
||||
rope_ptr = (token_data_ptr + NOPE_DIM).to(tl.pointer_type(tl.bfloat16))
|
||||
k_rope = tl.load(
|
||||
rope_ptr[:, None] + rope_offsets[None, :],
|
||||
mask=valid[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
k_rope = tl.where(valid[:, None], k_rope, zero_rope)
|
||||
k_rope = tl.where(k_rope == k_rope, k_rope, zero_rope)
|
||||
|
||||
scores = tl.dot(q_nope, tl.trans(k_nope)) + tl.dot(
|
||||
q_rope,
|
||||
tl.trans(k_rope),
|
||||
)
|
||||
scores *= scale
|
||||
scores = tl.where(head_mask[:, None] & valid[None, :], scores, neg_large)
|
||||
|
||||
m_block = tl.max(scores, axis=1)
|
||||
m_new = tl.maximum(m_i, m_block)
|
||||
alpha = tl.exp(m_i - m_new)
|
||||
p = tl.exp(scores - m_new[:, None])
|
||||
p = tl.where(head_mask[:, None] & valid[None, :], p, 0.0)
|
||||
l_new = l_i * alpha + tl.sum(p, axis=1)
|
||||
|
||||
acc_nope = acc_nope * alpha[:, None] + tl.dot(p.to(k_nope.dtype), k_nope)
|
||||
acc_rope = acc_rope * alpha[:, None] + tl.dot(p.to(k_rope.dtype), k_rope)
|
||||
m_i = m_new
|
||||
l_i = l_new
|
||||
|
||||
# Store raw (un-normalized) partial state for this split. Softmax sink and
|
||||
# final normalization happen in the reduce kernel.
|
||||
pm_base = query_idx * pm_stride0 + split_id * pm_stride_s + head_offsets
|
||||
tl.store(part_m_ptr + pm_base, m_i, mask=head_mask)
|
||||
tl.store(part_l_ptr + pm_base, l_i, mask=head_mask)
|
||||
acc_base = (
|
||||
part_acc_ptr
|
||||
+ query_idx * pa_stride0
|
||||
+ split_id * pa_stride_s
|
||||
+ head_offsets[:, None] * pa_stride_h
|
||||
)
|
||||
tl.store(
|
||||
acc_base + nope_offsets[None, :],
|
||||
acc_nope,
|
||||
mask=head_mask[:, None] & nope_mask[None, :],
|
||||
)
|
||||
tl.store(
|
||||
acc_base + NOPE_DIM + rope_offsets[None, :],
|
||||
acc_rope,
|
||||
mask=head_mask[:, None],
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _sparse_attn_decode_reduce_kernel(
|
||||
part_m_ptr,
|
||||
part_l_ptr,
|
||||
part_acc_ptr,
|
||||
attn_sink_ptr,
|
||||
out_ptr,
|
||||
out_stride0,
|
||||
out_stride1,
|
||||
pm_stride0,
|
||||
pm_stride_s,
|
||||
pa_stride0,
|
||||
pa_stride_s,
|
||||
pa_stride_h,
|
||||
num_heads,
|
||||
HAS_ATTN_SINK: tl.constexpr,
|
||||
COMB_DIM: tl.constexpr,
|
||||
BLOCK_H: tl.constexpr,
|
||||
NUM_SPLITS: tl.constexpr,
|
||||
SPLITS_PAD: tl.constexpr,
|
||||
):
|
||||
query_idx = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
|
||||
head_offsets = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
|
||||
head_mask = head_offsets < num_heads
|
||||
comb_offsets = tl.arange(0, COMB_DIM)
|
||||
# SPLITS_PAD is NUM_SPLITS rounded up to a power of two so the parallel
|
||||
# split-axis load is a legal arange for any split count; padding lanes are
|
||||
# masked off.
|
||||
split_offsets = tl.arange(0, SPLITS_PAD)
|
||||
split_mask = split_offsets < NUM_SPLITS
|
||||
|
||||
neg_large = -3.4028234663852886e38
|
||||
|
||||
# Phase 1: load every split's running max/sum at once and reduce the max
|
||||
# in parallel (tl.max over the split axis) instead of walking the splits
|
||||
# serially. This breaks the long online-softmax dependency chain that made
|
||||
# the reduce latency-bound.
|
||||
load_mask = split_mask[:, None] & head_mask[None, :]
|
||||
pm_split = (
|
||||
part_m_ptr
|
||||
+ query_idx * pm_stride0
|
||||
+ split_offsets[:, None] * pm_stride_s
|
||||
+ head_offsets[None, :]
|
||||
)
|
||||
m_all = tl.load(pm_split, mask=load_mask, other=neg_large) # [S, H]
|
||||
l_all = tl.load(
|
||||
part_l_ptr
|
||||
+ query_idx * pm_stride0
|
||||
+ split_offsets[:, None] * pm_stride_s
|
||||
+ head_offsets[None, :],
|
||||
mask=load_mask,
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
m_comb = tl.max(m_all, axis=0) # [H]
|
||||
if HAS_ATTN_SINK:
|
||||
sink = tl.load(
|
||||
attn_sink_ptr + head_offsets, mask=head_mask, other=neg_large
|
||||
).to(tl.float32)
|
||||
m_final = tl.maximum(m_comb, sink)
|
||||
else:
|
||||
m_final = m_comb
|
||||
|
||||
w_all = tl.exp(m_all - m_final[None, :]) # [S, H]
|
||||
w_all = tl.where(load_mask, w_all, 0.0)
|
||||
l_final = tl.sum(w_all * l_all, axis=0) # [H]
|
||||
if HAS_ATTN_SINK:
|
||||
l_final = l_final + tl.exp(sink - m_final)
|
||||
denom = tl.maximum(l_final, 1.0e-30)
|
||||
|
||||
# Phase 2: weighted sum of the per-split accumulators. The combine weight
|
||||
# for each split only depends on the (already known) global max, so the
|
||||
# acc loads carry no cross-split dependency and the compiler can pipeline
|
||||
# them; only the cheap FMA into `acc` is loop-carried.
|
||||
acc = tl.zeros((BLOCK_H, COMB_DIM), dtype=tl.float32)
|
||||
for s in tl.static_range(NUM_SPLITS):
|
||||
m_s = tl.load(
|
||||
part_m_ptr + query_idx * pm_stride0 + s * pm_stride_s + head_offsets,
|
||||
mask=head_mask,
|
||||
other=neg_large,
|
||||
)
|
||||
w_s = tl.exp(m_s - m_final)
|
||||
acc_base = (
|
||||
part_acc_ptr
|
||||
+ query_idx * pa_stride0
|
||||
+ s * pa_stride_s
|
||||
+ head_offsets[:, None] * pa_stride_h
|
||||
)
|
||||
acc_s = tl.load(
|
||||
acc_base + comb_offsets[None, :],
|
||||
mask=head_mask[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
acc += w_s[:, None] * acc_s
|
||||
|
||||
out = tl.where(l_final[:, None] > 0.0, acc / denom[:, None], 0.0)
|
||||
|
||||
out_row_ptr = (
|
||||
out_ptr + query_idx * out_stride0 + head_offsets[:, None] * out_stride1
|
||||
)
|
||||
tl.store(
|
||||
out_row_ptr + comb_offsets[None, :],
|
||||
out,
|
||||
mask=head_mask[:, None],
|
||||
)
|
||||
|
||||
|
||||
def _rocm_sparse_attn_prefill_ragged_triton(
|
||||
q: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
@@ -1502,6 +1844,101 @@ def _rocm_sparse_attn_prefill_triton(
|
||||
)
|
||||
|
||||
|
||||
@functools.lru_cache
|
||||
def _decode_cu_count() -> int:
|
||||
try:
|
||||
return torch.cuda.get_device_properties(0).multi_processor_count
|
||||
except Exception:
|
||||
return 256 # For gfx950 arch, gated behind a fallback path for other archs.
|
||||
|
||||
|
||||
def _decode_partial_iters(
|
||||
avg_main_len: float, avg_extra_len: float, splits: int, block_k: int
|
||||
) -> int:
|
||||
"""BLOCK_K iterations one partial workgroup walks for ``splits`` splits.
|
||||
|
||||
Each split processes ``ceil(seg_len / splits)`` tokens of a segment, walked
|
||||
``BLOCK_K`` at a time, and the main/extra segments are handled separately.
|
||||
"""
|
||||
main_iters = (
|
||||
math.ceil(math.ceil(avg_main_len / splits) / block_k) if avg_main_len > 0 else 0
|
||||
)
|
||||
extra_iters = (
|
||||
math.ceil(math.ceil(avg_extra_len / splits) / block_k)
|
||||
if avg_extra_len > 0
|
||||
else 0
|
||||
)
|
||||
return main_iters + extra_iters
|
||||
|
||||
|
||||
def _decode_num_splits(
|
||||
num_queries: int,
|
||||
heads_blocks: int,
|
||||
avg_main_len: float = 0.0,
|
||||
avg_extra_len: float = 0.0,
|
||||
block_k: int = 32,
|
||||
) -> int:
|
||||
"""Pick a flash-decode split count to keep the GPU busy across batch sizes.
|
||||
|
||||
Decode launches only ``num_queries * heads_blocks`` workgroups otherwise,
|
||||
which severely under-fills the device for the low-concurrency regime that
|
||||
dominates latency. Splitting the KV sequence adds parallelism.
|
||||
|
||||
We model the relative partial-kernel latency for a given split count ``s``
|
||||
as ``waves * (1/s + mu)`` where ``waves = ceil(base * s / CU)`` and ``mu``
|
||||
is a small per-wave overhead penalty:
|
||||
|
||||
- ``waves / s`` captures the partial compute: each wave walks roughly
|
||||
``total_tokens / s`` tokens and there are ``waves`` of them, so dividing
|
||||
by ``s`` makes more splits cheaper *until* they spill into extra waves.
|
||||
- ``mu * waves`` charges per-wave launch/tail overhead so we do not
|
||||
over-split into many mostly-idle waves (e.g. batch 224 on 256 CUs is
|
||||
best left at 1 split rather than 8 splits across 7 waves).
|
||||
|
||||
The minimiser naturally prefers split counts that pack the device into full
|
||||
waves (``base * s`` near a multiple of ``CU``) and falls back to 1 split
|
||||
once the batch already fills the device. Ties favour the smaller split
|
||||
count (less reduce work).
|
||||
|
||||
Finally we "snap down" the chosen split count to the smallest value that
|
||||
yields the same wave count *and* the same per-workgroup BLOCK_K iteration
|
||||
count. Because latency tracks iteration count (not raw token count), extra
|
||||
splits that do not lower the iteration count add only reduce/HBM overhead
|
||||
for no parallelism gain (e.g. batch 24: s8 and s10 both walk 4 extra iters
|
||||
in one wave, so s8 is strictly better). Snapping needs the average segment
|
||||
lengths, which the caller derives sync-free from the ragged index sizes.
|
||||
"""
|
||||
base = max(1, num_queries * heads_blocks)
|
||||
# Target ~1 workgroup per CU: enough to fill the device while keeping the
|
||||
# reduce cost (which grows with split count) small. Tuned on gfx950.
|
||||
cu = max(1, _decode_cu_count())
|
||||
# Per-wave overhead penalty: higher values discourage split counts that
|
||||
# spill into extra GPU waves. Tuned on gfx950.
|
||||
mu = 0.04
|
||||
best_splits = 1
|
||||
best_cost = None
|
||||
# Search up to 16 splits; beyond that the reduce/HBM overhead dominates.
|
||||
for splits in range(1, 17):
|
||||
waves = (base * splits + cu - 1) // cu
|
||||
cost = waves * (1.0 / splits + mu)
|
||||
if best_cost is None or cost < best_cost - 1e-9:
|
||||
best_splits = splits
|
||||
best_cost = cost
|
||||
|
||||
if best_splits > 1 and (avg_main_len > 0 or avg_extra_len > 0):
|
||||
target_waves = (base * best_splits + cu - 1) // cu
|
||||
target_iters = _decode_partial_iters(
|
||||
avg_main_len, avg_extra_len, best_splits, block_k
|
||||
)
|
||||
for splits in range(1, best_splits):
|
||||
waves = (base * splits + cu - 1) // cu
|
||||
iters = _decode_partial_iters(avg_main_len, avg_extra_len, splits, block_k)
|
||||
if waves == target_waves and iters == target_iters:
|
||||
best_splits = splits
|
||||
break
|
||||
return best_splits
|
||||
|
||||
|
||||
def _rocm_sparse_attn_decode_ragged_triton(
|
||||
q: torch.Tensor,
|
||||
main_cache: torch.Tensor,
|
||||
@@ -1575,9 +2012,70 @@ def _rocm_sparse_attn_decode_ragged_triton(
|
||||
extra_indptr = torch.zeros(num_queries + 1, device=q.device, dtype=torch.int32)
|
||||
|
||||
block_h = 16
|
||||
block_k = 16 if head_dim >= 256 else 32
|
||||
out = torch.empty_like(q, dtype=torch.bfloat16)
|
||||
_sparse_attn_decode_ragged_kernel[(num_queries, triton.cdiv(num_heads, block_h))](
|
||||
heads_blocks = triton.cdiv(num_heads, block_h)
|
||||
nope_block = triton.next_power_of_2(nope_head_dim)
|
||||
comb_dim = nope_head_dim + rope_head_dim
|
||||
is_fnuz = current_platform.is_fp8_fnuz()
|
||||
|
||||
if not _ON_GFX950: # Fallback path for un-tuned architectures.
|
||||
block_k = 16 if head_dim >= 256 else 32
|
||||
_sparse_attn_decode_ragged_kernel[(num_queries, heads_blocks)](
|
||||
q,
|
||||
main_cache,
|
||||
main_indices,
|
||||
main_indptr,
|
||||
extra_cache,
|
||||
extra_indices,
|
||||
extra_indptr,
|
||||
attn_sink,
|
||||
out,
|
||||
q.stride(0),
|
||||
q.stride(1),
|
||||
out.stride(0),
|
||||
out.stride(1),
|
||||
main_cache.stride(0),
|
||||
extra_cache.stride(0),
|
||||
main_cache.shape[0] * main_cache.shape[1],
|
||||
extra_cache.shape[0] * extra_cache.shape[1],
|
||||
main_cache.shape[1],
|
||||
extra_cache.shape[1],
|
||||
scale,
|
||||
num_heads,
|
||||
HAS_ATTN_SINK=has_attn_sink,
|
||||
HAS_EXTRA=has_extra,
|
||||
NOPE_DIM=nope_head_dim,
|
||||
NOPE_BLOCK=nope_block,
|
||||
ROPE_DIM=rope_head_dim,
|
||||
IS_FNUZ=is_fnuz,
|
||||
BLOCK_H=block_h,
|
||||
BLOCK_K=block_k,
|
||||
num_warps=8,
|
||||
)
|
||||
return out
|
||||
|
||||
block_k = 32 # KV tokens walked per split-K iteration. Tuned on gfx950.
|
||||
# Average per-query segment lengths, read sync-free from the ragged index
|
||||
# sizes, let the split heuristic avoid over-splitting
|
||||
# main_indices/extra_indices are flat [nnz] int32.
|
||||
inv_q = 1.0 / max(1, num_queries)
|
||||
avg_main_len = main_indices.numel() * inv_q
|
||||
avg_extra_len = (extra_indices.numel() * inv_q) if has_extra else 0.0
|
||||
num_splits = _decode_num_splits(
|
||||
num_queries, heads_blocks, avg_main_len, avg_extra_len, block_k
|
||||
)
|
||||
|
||||
part_m = torch.empty(
|
||||
(num_queries, num_splits, num_heads), dtype=torch.float32, device=q.device
|
||||
)
|
||||
part_l = torch.empty_like(part_m)
|
||||
part_acc = torch.empty(
|
||||
(num_queries, num_splits, num_heads, comb_dim),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
|
||||
_sparse_attn_decode_partial_kernel[(num_queries, num_splits, heads_blocks)](
|
||||
q,
|
||||
main_cache,
|
||||
main_indices,
|
||||
@@ -1585,29 +2083,56 @@ def _rocm_sparse_attn_decode_ragged_triton(
|
||||
extra_cache,
|
||||
extra_indices,
|
||||
extra_indptr,
|
||||
attn_sink,
|
||||
out,
|
||||
part_m,
|
||||
part_l,
|
||||
part_acc,
|
||||
q.stride(0),
|
||||
q.stride(1),
|
||||
out.stride(0),
|
||||
out.stride(1),
|
||||
main_cache.stride(0),
|
||||
extra_cache.stride(0),
|
||||
part_m.stride(0),
|
||||
part_m.stride(1),
|
||||
part_acc.stride(0),
|
||||
part_acc.stride(1),
|
||||
part_acc.stride(2),
|
||||
main_cache.shape[0] * main_cache.shape[1],
|
||||
extra_cache.shape[0] * extra_cache.shape[1],
|
||||
main_cache.shape[1],
|
||||
extra_cache.shape[1],
|
||||
scale,
|
||||
num_heads,
|
||||
HAS_ATTN_SINK=has_attn_sink,
|
||||
HAS_EXTRA=has_extra,
|
||||
NOPE_DIM=nope_head_dim,
|
||||
NOPE_BLOCK=triton.next_power_of_2(nope_head_dim),
|
||||
NOPE_BLOCK=nope_block,
|
||||
ROPE_DIM=rope_head_dim,
|
||||
IS_FNUZ=current_platform.is_fp8_fnuz(),
|
||||
IS_FNUZ=is_fnuz,
|
||||
BLOCK_H=block_h,
|
||||
BLOCK_K=block_k,
|
||||
num_warps=8,
|
||||
NUM_SPLITS=num_splits,
|
||||
NUM_STAGES=1,
|
||||
num_warps=4,
|
||||
)
|
||||
|
||||
_sparse_attn_decode_reduce_kernel[(num_queries, heads_blocks)](
|
||||
part_m,
|
||||
part_l,
|
||||
part_acc,
|
||||
attn_sink,
|
||||
out,
|
||||
out.stride(0),
|
||||
out.stride(1),
|
||||
part_m.stride(0),
|
||||
part_m.stride(1),
|
||||
part_acc.stride(0),
|
||||
part_acc.stride(1),
|
||||
part_acc.stride(2),
|
||||
num_heads,
|
||||
HAS_ATTN_SINK=has_attn_sink,
|
||||
COMB_DIM=comb_dim,
|
||||
BLOCK_H=block_h,
|
||||
NUM_SPLITS=num_splits,
|
||||
SPLITS_PAD=triton.next_power_of_2(num_splits),
|
||||
num_warps=4,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
Reference in New Issue
Block a user