[AMD] Parallelize aiter spec-decode KV index building over token blocks (#37659)
Co-authored-by: Zijie Chen <300606707+zijiecode@users.noreply.github.com> Co-authored-by: jacky.cheng <yichiche@amd.com>
This commit is contained in:
co-authored by
Zijie Chen
jacky.cheng
parent
880d6fa64d
commit
92d831d3d7
@@ -48,6 +48,9 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
|||||||
from sglang.kernels.ops.kvcache.kv_indices import (
|
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||||
get_num_page_per_block_flashmla as get_num_page_per_block_flashmla,
|
get_num_page_per_block_flashmla as get_num_page_per_block_flashmla,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||||
|
kv_indices_num_token_blocks as kv_indices_num_token_blocks,
|
||||||
|
)
|
||||||
from sglang.kernels.ops.kvcache.rope_cache import (
|
from sglang.kernels.ops.kvcache.rope_cache import (
|
||||||
fused_qk_rope_reshape_and_cache as fused_qk_rope_reshape_and_cache,
|
fused_qk_rope_reshape_and_cache as fused_qk_rope_reshape_and_cache,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4,6 +4,19 @@ import triton.language as tl
|
|||||||
_FLASHMLA_CREATE_KV_BLOCK_SIZE = 4096
|
_FLASHMLA_CREATE_KV_BLOCK_SIZE = 4096
|
||||||
FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON = tl.constexpr(_FLASHMLA_CREATE_KV_BLOCK_SIZE)
|
FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON = tl.constexpr(_FLASHMLA_CREATE_KV_BLOCK_SIZE)
|
||||||
|
|
||||||
|
# Token-block parallelism for the index-copy kernels below: aim for about
|
||||||
|
# _TARGET_PROGRAMS programs in total, one extra block per
|
||||||
|
# _MIN_TOKENS_PER_BLOCK of table width at most, and fall back to the
|
||||||
|
# historical single-block grid when the base grid is already wide.
|
||||||
|
_MIN_TOKENS_PER_BLOCK = 8192
|
||||||
|
_TARGET_PROGRAMS = 512
|
||||||
|
|
||||||
|
|
||||||
|
def kv_indices_num_token_blocks(table_width: int, base_programs: int) -> int:
|
||||||
|
cap = (table_width + _MIN_TOKENS_PER_BLOCK - 1) // _MIN_TOKENS_PER_BLOCK
|
||||||
|
want = _TARGET_PROGRAMS // max(1, base_programs)
|
||||||
|
return max(1, min(cap, want))
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def create_flashinfer_kv_indices_triton(
|
def create_flashinfer_kv_indices_triton(
|
||||||
@@ -19,6 +32,7 @@ def create_flashinfer_kv_indices_triton(
|
|||||||
# (a recompile every few decode steps at small page sizes).
|
# (a recompile every few decode steps at small page sizes).
|
||||||
req_to_token_ptr_stride,
|
req_to_token_ptr_stride,
|
||||||
ENTRY_PAGE_SIZE: tl.constexpr = 1,
|
ENTRY_PAGE_SIZE: tl.constexpr = 1,
|
||||||
|
TOKEN_BLOCK_PARALLEL: tl.constexpr = False,
|
||||||
):
|
):
|
||||||
"""Gather per-request token ids into a flat CSR kv_indices stream.
|
"""Gather per-request token ids into a flat CSR kv_indices stream.
|
||||||
|
|
||||||
@@ -28,9 +42,21 @@ def create_flashinfer_kv_indices_triton(
|
|||||||
read table (entries already kernel-facing page ids); token ids are rebuilt
|
read table (entries already kernel-facing page ids); token ids are rebuilt
|
||||||
as ``token = entry * ps + pos % ps``, exact because converting an id keeps
|
as ``token = entry * ps + pos % ps``, exact because converting an id keeps
|
||||||
its offset inside the page.
|
its offset inside the page.
|
||||||
|
``TOKEN_BLOCK_PARALLEL`` (default False): launched on a 2D grid
|
||||||
|
``(batch, num_blocks)``, the programs of a request stride over its copy
|
||||||
|
loop together instead of one program crawling the whole context serially
|
||||||
|
(which bottlenecks long-context spec decode, where this kernel runs every
|
||||||
|
iteration). With the default, the kernel is the historical
|
||||||
|
one-program-per-request loop and 1D launch sites are unaffected.
|
||||||
"""
|
"""
|
||||||
BLOCK_SIZE: tl.constexpr = 512
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
pid = tl.program_id(axis=0)
|
pid = tl.program_id(axis=0)
|
||||||
|
if TOKEN_BLOCK_PARALLEL:
|
||||||
|
blk = tl.program_id(axis=1)
|
||||||
|
num_blk = tl.num_programs(axis=1)
|
||||||
|
else:
|
||||||
|
blk = 0
|
||||||
|
num_blk = 1
|
||||||
|
|
||||||
# find the req pool idx, this is for batch to token
|
# find the req pool idx, this is for batch to token
|
||||||
req_pool_index = tl.load(req_pool_indices_ptr + pid).to(tl.int64)
|
req_pool_index = tl.load(req_pool_indices_ptr + pid).to(tl.int64)
|
||||||
@@ -44,7 +70,11 @@ def create_flashinfer_kv_indices_triton(
|
|||||||
kv_end += tl.load(page_kernel_lens_ptr + pid).to(tl.int32)
|
kv_end += tl.load(page_kernel_lens_ptr + pid).to(tl.int32)
|
||||||
|
|
||||||
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
||||||
for i in range(num_loop):
|
if TOKEN_BLOCK_PARALLEL:
|
||||||
|
# Blocks with no copy work exit early.
|
||||||
|
if blk >= num_loop:
|
||||||
|
return
|
||||||
|
for i in range(blk, num_loop, num_blk):
|
||||||
# index into req_to_token_ptr needs to be int64
|
# index into req_to_token_ptr needs to be int64
|
||||||
offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE
|
offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE
|
||||||
mask = offset < kv_end - kv_start
|
mask = offset < kv_end - kv_start
|
||||||
|
|||||||
@@ -67,13 +67,30 @@ def generate_draft_decode_kv_indices(
|
|||||||
iter_upper: tl.constexpr,
|
iter_upper: tl.constexpr,
|
||||||
num_tokens_upper: tl.constexpr,
|
num_tokens_upper: tl.constexpr,
|
||||||
page_size: tl.constexpr,
|
page_size: tl.constexpr,
|
||||||
|
NUM_STEPS: tl.constexpr = 0,
|
||||||
):
|
):
|
||||||
BLOCK_SIZE: tl.constexpr = 128
|
# Optional token-block parallelism (NUM_STEPS > 0): the first grid axis
|
||||||
iters = tl.program_id(axis=0)
|
# packs (draft step, token block) as ``step + NUM_STEPS * block``,
|
||||||
|
# spreading the per-request index copy below over many programs instead
|
||||||
|
# of one program crawling the whole context serially (which bottlenecks
|
||||||
|
# long-context spec decode, where this kernel runs every iteration).
|
||||||
|
# NUM_STEPS == 0 (default) is the historical one-program-per-step kernel:
|
||||||
|
# the same 128-wide copy loop, in the same order, with the token-block
|
||||||
|
# branches folded away at compile time.
|
||||||
|
BLOCK_SIZE: tl.constexpr = 128 if NUM_STEPS == 0 else 512
|
||||||
|
pid0 = tl.program_id(axis=0)
|
||||||
bid = tl.program_id(axis=1)
|
bid = tl.program_id(axis=1)
|
||||||
topk_id = tl.program_id(axis=2)
|
topk_id = tl.program_id(axis=2)
|
||||||
|
|
||||||
|
if NUM_STEPS == 0:
|
||||||
|
iters = pid0
|
||||||
num_steps = tl.num_programs(axis=0)
|
num_steps = tl.num_programs(axis=0)
|
||||||
|
blk = 0
|
||||||
|
else:
|
||||||
|
iters = pid0 % NUM_STEPS
|
||||||
|
blk = pid0 // NUM_STEPS
|
||||||
|
num_steps = NUM_STEPS
|
||||||
|
num_blk = tl.num_programs(axis=0) // NUM_STEPS
|
||||||
num_seqs = tl.num_programs(axis=1)
|
num_seqs = tl.num_programs(axis=1)
|
||||||
topk = tl.num_programs(axis=2)
|
topk = tl.num_programs(axis=2)
|
||||||
|
|
||||||
@@ -81,16 +98,32 @@ def generate_draft_decode_kv_indices(
|
|||||||
kv_indptr += kv_indptr_stride * iters
|
kv_indptr += kv_indptr_stride * iters
|
||||||
iters += 1
|
iters += 1
|
||||||
|
|
||||||
|
if NUM_STEPS == 0:
|
||||||
load_offset = tl.arange(0, bs_upper)
|
load_offset = tl.arange(0, bs_upper)
|
||||||
seq_lens = tl.load(paged_kernel_lens + load_offset, mask=load_offset < bid, other=0)
|
seq_lens = tl.load(
|
||||||
|
paged_kernel_lens + load_offset, mask=load_offset < bid, other=0
|
||||||
|
)
|
||||||
seq_len = tl.load(paged_kernel_lens + bid)
|
seq_len = tl.load(paged_kernel_lens + bid)
|
||||||
cum_seq_len = tl.sum(seq_lens)
|
cum_seq_len = tl.sum(seq_lens)
|
||||||
|
else:
|
||||||
|
seq_len = tl.load(paged_kernel_lens + bid)
|
||||||
|
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
||||||
|
# Blocks with no copy work exit before the O(bs) prefix-sum below;
|
||||||
|
# block 0 always continues (it owns the extension and kv_indptr).
|
||||||
|
if blk >= num_loop and blk > 0:
|
||||||
|
return
|
||||||
|
load_offset = tl.arange(0, bs_upper)
|
||||||
|
seq_lens = tl.load(
|
||||||
|
paged_kernel_lens + load_offset, mask=load_offset < bid, other=0
|
||||||
|
)
|
||||||
|
cum_seq_len = tl.sum(seq_lens)
|
||||||
|
|
||||||
# Update kv_indices
|
# Update kv_indices
|
||||||
kv_offset = cum_seq_len * topk + bid * iters * topk + topk_id * (seq_len + iters)
|
kv_offset = cum_seq_len * topk + bid * iters * topk + topk_id * (seq_len + iters)
|
||||||
kv_ptr = kv_indices + kv_offset
|
kv_ptr = kv_indices + kv_offset
|
||||||
token_pool_ptr = req_to_token + tl.load(req_pool_indices + bid) * pool_len
|
token_pool_ptr = req_to_token + tl.load(req_pool_indices + bid) * pool_len
|
||||||
|
|
||||||
|
if NUM_STEPS == 0:
|
||||||
kv_offset = tl.arange(0, BLOCK_SIZE)
|
kv_offset = tl.arange(0, BLOCK_SIZE)
|
||||||
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
||||||
for _ in range(num_loop):
|
for _ in range(num_loop):
|
||||||
@@ -98,11 +131,23 @@ def generate_draft_decode_kv_indices(
|
|||||||
data = tl.load(token_pool_ptr + kv_offset, mask=mask)
|
data = tl.load(token_pool_ptr + kv_offset, mask=mask)
|
||||||
tl.store(kv_ptr + kv_offset, data, mask=mask)
|
tl.store(kv_ptr + kv_offset, data, mask=mask)
|
||||||
kv_offset += BLOCK_SIZE
|
kv_offset += BLOCK_SIZE
|
||||||
|
else:
|
||||||
|
for i in range(blk, num_loop, num_blk):
|
||||||
|
tok_off = i * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = tok_off < seq_len
|
||||||
|
data = tl.load(token_pool_ptr + tok_off, mask=mask)
|
||||||
|
tl.store(kv_ptr + tok_off, data, mask=mask)
|
||||||
|
|
||||||
|
# Extension entries and kv_indptr belong to token block 0 alone; other
|
||||||
|
# blocks neither compute nor store them.
|
||||||
|
if blk == 0:
|
||||||
extend_offset = tl.arange(0, iter_upper)
|
extend_offset = tl.arange(0, iter_upper)
|
||||||
if page_size == 1 or topk == 1:
|
if page_size == 1 or topk == 1:
|
||||||
extend_data = tl.load(
|
extend_data = tl.load(
|
||||||
token_pool_ptr + seq_len + topk_id * num_steps + tl.arange(0, iter_upper),
|
token_pool_ptr
|
||||||
|
+ seq_len
|
||||||
|
+ topk_id * num_steps
|
||||||
|
+ tl.arange(0, iter_upper),
|
||||||
mask=extend_offset < iters,
|
mask=extend_offset < iters,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -113,14 +158,20 @@ def generate_draft_decode_kv_indices(
|
|||||||
) // page_size
|
) // page_size
|
||||||
prefix_base = seq_len // page_size * page_size
|
prefix_base = seq_len // page_size * page_size
|
||||||
start = (
|
start = (
|
||||||
prefix_base + topk_id * num_new_pages_per_topk * page_size + last_page_len
|
prefix_base
|
||||||
|
+ topk_id * num_new_pages_per_topk * page_size
|
||||||
|
+ last_page_len
|
||||||
)
|
)
|
||||||
extend_data = tl.load(
|
extend_data = tl.load(
|
||||||
token_pool_ptr + start + extend_offset,
|
token_pool_ptr + start + extend_offset,
|
||||||
mask=extend_offset < iters,
|
mask=extend_offset < iters,
|
||||||
)
|
)
|
||||||
|
|
||||||
tl.store(kv_ptr + seq_len + extend_offset, extend_data, mask=extend_offset < iters)
|
tl.store(
|
||||||
|
kv_ptr + seq_len + extend_offset,
|
||||||
|
extend_data,
|
||||||
|
mask=extend_offset < iters,
|
||||||
|
)
|
||||||
|
|
||||||
# Update kv_indptr
|
# Update kv_indptr
|
||||||
bs_offset = tl.arange(0, num_tokens_upper)
|
bs_offset = tl.arange(0, num_tokens_upper)
|
||||||
@@ -128,8 +179,8 @@ def generate_draft_decode_kv_indices(
|
|||||||
zid = bid * topk + topk_id
|
zid = bid * topk + topk_id
|
||||||
if zid == 0:
|
if zid == 0:
|
||||||
zid = num_seqs * topk
|
zid = num_seqs * topk
|
||||||
positions = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0)
|
pos_vals = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0)
|
||||||
base = tl.sum(positions)
|
base = tl.sum(pos_vals)
|
||||||
tl.store(kv_indptr + zid, base + zid * iters)
|
tl.store(kv_indptr + zid, base + zid * iters)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.kernels.ops.attention.utils import (
|
|||||||
create_flashinfer_kv_indices_triton,
|
create_flashinfer_kv_indices_triton,
|
||||||
create_flashmla_kv_indices_triton,
|
create_flashmla_kv_indices_triton,
|
||||||
get_num_kv_index_blocks_flashmla,
|
get_num_kv_index_blocks_flashmla,
|
||||||
|
kv_indices_num_token_blocks,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
||||||
scatter_ragged_to_page_table_kernel,
|
scatter_ragged_to_page_table_kernel,
|
||||||
@@ -124,6 +125,13 @@ fast_mode = False
|
|||||||
intra_batch_mode = True if _use_mla_ps_kernel else False
|
intra_batch_mode = True if _use_mla_ps_kernel else False
|
||||||
|
|
||||||
|
|
||||||
|
# Token-block parallel KV-index building is enabled only where it pays:
|
||||||
|
# the speculative-decoding paths (target_verify / draft_extend / draft
|
||||||
|
# decode) of long-context servers. Everything else keeps the historical
|
||||||
|
# one-program-per-request launch.
|
||||||
|
_KV_INDEX_BLOCKS_MIN_CONTEXT = 32768
|
||||||
|
|
||||||
|
|
||||||
class WrapperDispatch(Enum):
|
class WrapperDispatch(Enum):
|
||||||
SLIDING_WINDOW = auto()
|
SLIDING_WINDOW = auto()
|
||||||
CROSS_ATTENTION = auto()
|
CROSS_ATTENTION = auto()
|
||||||
@@ -1170,6 +1178,11 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
return output[:, : layer.tp_q_head_num, :] if head_pad else output
|
return output[:, : layer.tp_q_head_num, :] if head_pad else output
|
||||||
|
|
||||||
|
def _kv_index_blocks(self, bs: int) -> int:
|
||||||
|
if self.max_context_len < _KV_INDEX_BLOCKS_MIN_CONTEXT:
|
||||||
|
return 1
|
||||||
|
return kv_indices_num_token_blocks(self.req_to_token.shape[1], bs)
|
||||||
|
|
||||||
def init_forward_metadata_out_graph(
|
def init_forward_metadata_out_graph(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
@@ -1374,7 +1387,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
forward_batch.seq_lens_sum, device
|
forward_batch.seq_lens_sum, device
|
||||||
)
|
)
|
||||||
|
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
num_token_blocks = self._kv_index_blocks(bs)
|
||||||
|
create_flashinfer_kv_indices_triton[(bs, num_token_blocks)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
@@ -1382,6 +1396,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
if _use_mla_ps_kernel:
|
if _use_mla_ps_kernel:
|
||||||
@@ -1467,7 +1482,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_lens_sum,
|
kv_lens_sum,
|
||||||
device,
|
device,
|
||||||
)
|
)
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
num_token_blocks = self._kv_index_blocks(bs)
|
||||||
|
create_flashinfer_kv_indices_triton[(bs, num_token_blocks)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
kv_lens,
|
kv_lens,
|
||||||
@@ -1475,6 +1491,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# if self.kv_cache_dtype == fp8_dtype:
|
# if self.kv_cache_dtype == fp8_dtype:
|
||||||
@@ -1564,7 +1581,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices = torch.empty(
|
kv_indices = torch.empty(
|
||||||
kv_indptr[-1], dtype=torch.int64, device=self.device
|
kv_indptr[-1], dtype=torch.int64, device=self.device
|
||||||
)
|
)
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
num_token_blocks = self._kv_index_blocks(bs)
|
||||||
|
create_flashinfer_kv_indices_triton[(bs, num_token_blocks)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
@@ -1572,6 +1590,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
custom_mask = spec_info.custom_mask
|
custom_mask = spec_info.custom_mask
|
||||||
@@ -2018,7 +2037,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
bs=bs,
|
bs=bs,
|
||||||
seq_lens_sum=seq_lens_sum,
|
seq_lens_sum=seq_lens_sum,
|
||||||
)
|
)
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
num_token_blocks = self._kv_index_blocks(bs)
|
||||||
|
create_flashinfer_kv_indices_triton[(bs, num_token_blocks)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
kv_lens,
|
kv_lens,
|
||||||
@@ -2026,6 +2046,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
|
|
||||||
@@ -2144,7 +2165,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indptr = self.kv_indptr[: bs + 1]
|
kv_indptr = self.kv_indptr[: bs + 1]
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
||||||
kv_indices = self.cuda_graph_kv_indices
|
kv_indices = self.cuda_graph_kv_indices
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
num_token_blocks = self._kv_index_blocks(bs)
|
||||||
|
create_flashinfer_kv_indices_triton[(bs, num_token_blocks)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
@@ -2152,6 +2174,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
@@ -3544,8 +3567,15 @@ class AiterMultiStepDraftBackend:
|
|||||||
bs = self.topk * num_seqs
|
bs = self.topk * num_seqs
|
||||||
seq_lens_sum = forward_batch.seq_lens_sum
|
seq_lens_sum = forward_batch.seq_lens_sum
|
||||||
|
|
||||||
|
num_token_blocks = (
|
||||||
|
kv_indices_num_token_blocks(
|
||||||
|
self.pool_len, self.speculative_num_steps * num_seqs * self.topk
|
||||||
|
)
|
||||||
|
if self.max_context_len >= _KV_INDEX_BLOCKS_MIN_CONTEXT
|
||||||
|
else 1
|
||||||
|
)
|
||||||
self.generate_draft_decode_kv_indices[
|
self.generate_draft_decode_kv_indices[
|
||||||
(self.speculative_num_steps, num_seqs, self.topk)
|
(self.speculative_num_steps * num_token_blocks, num_seqs, self.topk)
|
||||||
](
|
](
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
self.req_to_token_pool.req_to_token,
|
self.req_to_token_pool.req_to_token,
|
||||||
@@ -3560,6 +3590,9 @@ class AiterMultiStepDraftBackend:
|
|||||||
triton.next_power_of_2(self.speculative_num_steps),
|
triton.next_power_of_2(self.speculative_num_steps),
|
||||||
triton.next_power_of_2(bs),
|
triton.next_power_of_2(bs),
|
||||||
self.page_size,
|
self.page_size,
|
||||||
|
# A single token block is the historical launch; NUM_STEPS=0 keeps
|
||||||
|
# its 128-wide program instead of the token-block specialization.
|
||||||
|
NUM_STEPS=self.speculative_num_steps if num_token_blocks > 1 else 0,
|
||||||
)
|
)
|
||||||
|
|
||||||
for i in range(self.speculative_num_steps - 1):
|
for i in range(self.speculative_num_steps - 1):
|
||||||
|
|||||||
@@ -0,0 +1,195 @@
|
|||||||
|
"""2D (token-block) launches of the KV index-copy kernels must be
|
||||||
|
bit-identical to the historical 1D launches, with no unwritten or
|
||||||
|
overwritten bytes (sentinel-checked over the full buffers), and the draft
|
||||||
|
kernel's output must match a Python reference."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.utils import (
|
||||||
|
create_flashinfer_kv_indices_triton,
|
||||||
|
kv_indices_num_token_blocks,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
|
generate_draft_decode_kv_indices,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=25, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||||
|
register_amd_ci(est_time=25, stage="jit-kernel-unit", runner_config="amd")
|
||||||
|
|
||||||
|
SENTINEL = 0x7EADBEEF
|
||||||
|
POOL_LEN = 262_144
|
||||||
|
LENSETS = [
|
||||||
|
[0, 1, 511, 512, 513, 8191, 8192, 8193],
|
||||||
|
[100_000, 33, 4096],
|
||||||
|
[1000] * 7 + [100_000],
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _npo2(x: int) -> int:
|
||||||
|
return max(1, 1 << (max(1, x) - 1).bit_length())
|
||||||
|
|
||||||
|
|
||||||
|
def _draft_inputs(seqs, topk, steps, device, idx_dtype=torch.int64):
|
||||||
|
bs = len(seqs)
|
||||||
|
req_pool = torch.arange(bs, dtype=idx_dtype, device=device)
|
||||||
|
r2t = torch.randint(
|
||||||
|
0, 6_000_000, (bs + 1, POOL_LEN), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
lens = torch.tensor(seqs, dtype=idx_dtype, device=device)
|
||||||
|
width = topk * (max(seqs) + steps) + 64
|
||||||
|
return bs, req_pool, r2t, lens, width, bs * topk
|
||||||
|
|
||||||
|
|
||||||
|
def _run_draft(kern_inputs, topk, steps, page_size, nb, kw):
|
||||||
|
bs, req_pool, r2t, lens, width, tot = kern_inputs
|
||||||
|
dev = lens.device
|
||||||
|
kv_i = torch.full((steps, bs * width), SENTINEL, dtype=torch.int32, device=dev)
|
||||||
|
kv_p = torch.full((steps, tot + 1), SENTINEL, dtype=torch.int32, device=dev)
|
||||||
|
# positions is per draft token in production (bs * topk entries); the
|
||||||
|
# kernel reads positions[:bs * topk] for the kv_indptr prefix sums.
|
||||||
|
positions = torch.repeat_interleave(lens, topk)
|
||||||
|
generate_draft_decode_kv_indices[(steps * nb, bs, topk)](
|
||||||
|
req_pool,
|
||||||
|
r2t,
|
||||||
|
lens,
|
||||||
|
kv_i,
|
||||||
|
kv_p,
|
||||||
|
positions,
|
||||||
|
POOL_LEN,
|
||||||
|
kv_i.shape[1],
|
||||||
|
kv_p.shape[1],
|
||||||
|
_npo2(bs),
|
||||||
|
_npo2(steps),
|
||||||
|
_npo2(tot),
|
||||||
|
page_size,
|
||||||
|
**kw,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return kv_i, kv_p
|
||||||
|
|
||||||
|
|
||||||
|
class TestSpecKvIndicesGrid(CustomTestCase):
|
||||||
|
def test_draft_grid_equivalence(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
for seqs in LENSETS:
|
||||||
|
for topk, page_size in [(1, 1), (4, 1), (4, 16)]:
|
||||||
|
for steps, idx_dtype in [
|
||||||
|
(2, torch.int64),
|
||||||
|
(3, torch.int32),
|
||||||
|
(4, torch.int64),
|
||||||
|
]:
|
||||||
|
inputs = _draft_inputs(seqs, topk, steps, "cuda", idx_dtype)
|
||||||
|
ref = _run_draft(inputs, topk, steps, page_size, 1, {})
|
||||||
|
for nb in [
|
||||||
|
1,
|
||||||
|
kv_indices_num_token_blocks(POOL_LEN, steps * len(seqs) * topk),
|
||||||
|
triton.cdiv(POOL_LEN, 8192),
|
||||||
|
]:
|
||||||
|
out = _run_draft(
|
||||||
|
inputs, topk, steps, page_size, nb, {"NUM_STEPS": steps}
|
||||||
|
)
|
||||||
|
self.assertTrue(torch.equal(ref[0], out[0]), (seqs, topk, nb))
|
||||||
|
self.assertTrue(torch.equal(ref[1], out[1]), (seqs, topk, nb))
|
||||||
|
|
||||||
|
def test_draft_reference(self):
|
||||||
|
torch.manual_seed(1)
|
||||||
|
seqs, steps = [100_000, 33, 4096, 16], 3
|
||||||
|
for topk, page_size in [(1, 1), (4, 1), (4, 16)]:
|
||||||
|
inputs = _draft_inputs(seqs, topk, steps, "cuda")
|
||||||
|
bs, _, r2t, _, _, tot = inputs
|
||||||
|
nb = kv_indices_num_token_blocks(POOL_LEN, steps * bs * topk)
|
||||||
|
kv_i, kv_p = _run_draft(
|
||||||
|
inputs, topk, steps, page_size, nb, {"NUM_STEPS": steps}
|
||||||
|
)
|
||||||
|
for it in range(steps):
|
||||||
|
iters = it + 1
|
||||||
|
for s in range(bs):
|
||||||
|
ln = seqs[s]
|
||||||
|
for k in range(topk):
|
||||||
|
off = sum(seqs[:s]) * topk + s * iters * topk + k * (ln + iters)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(r2t[s, :ln], kv_i[it, off : off + ln]),
|
||||||
|
(it, s, k, topk, page_size),
|
||||||
|
)
|
||||||
|
if page_size == 1 or topk == 1:
|
||||||
|
src = ln + k * steps
|
||||||
|
else:
|
||||||
|
last = ln % page_size
|
||||||
|
pages = -(-(last + steps) // page_size)
|
||||||
|
src = (
|
||||||
|
ln // page_size * page_size
|
||||||
|
+ k * pages * page_size
|
||||||
|
+ last
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
r2t[s, src : src + iters],
|
||||||
|
kv_i[it, off + ln : off + ln + iters],
|
||||||
|
),
|
||||||
|
(it, s, k, topk, page_size),
|
||||||
|
)
|
||||||
|
positions = [ln for ln in seqs for _ in range(topk)]
|
||||||
|
for z in range(1, tot + 1):
|
||||||
|
self.assertEqual(
|
||||||
|
int(kv_p[it, z]),
|
||||||
|
sum(positions[:z]) + z * iters,
|
||||||
|
(it, z, topk, page_size),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_flat_grid_equivalence(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
dev = "cuda"
|
||||||
|
for seqs in LENSETS:
|
||||||
|
for use_start, entry_page_size in [(False, 1), (True, 1), (False, 16)]:
|
||||||
|
bs = len(seqs)
|
||||||
|
req_pool = torch.arange(bs, dtype=torch.int64, device=dev)
|
||||||
|
r2t = torch.randint(
|
||||||
|
0, 6_000_000, (bs + 1, POOL_LEN), dtype=torch.int32, device=dev
|
||||||
|
)
|
||||||
|
lens = torch.tensor(
|
||||||
|
seqs,
|
||||||
|
dtype=torch.int32 if use_start else torch.int64,
|
||||||
|
device=dev,
|
||||||
|
)
|
||||||
|
indptr = torch.zeros(bs + 1, dtype=torch.int32, device=dev)
|
||||||
|
indptr[1:] = torch.cumsum(lens, 0).to(torch.int32)
|
||||||
|
start = (
|
||||||
|
torch.full((bs,), 7, dtype=torch.int32, device=dev)
|
||||||
|
if use_start
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
n = int(indptr[-1]) + (7 * bs if use_start else 0) + 8
|
||||||
|
outs = []
|
||||||
|
for grid, parallel in [
|
||||||
|
((bs,), False),
|
||||||
|
((bs, 1), True),
|
||||||
|
((bs, kv_indices_num_token_blocks(POOL_LEN, bs)), True),
|
||||||
|
((bs, triton.cdiv(POOL_LEN, 8192)), True),
|
||||||
|
]:
|
||||||
|
kv_i = torch.full((n,), SENTINEL, dtype=torch.int32, device=dev)
|
||||||
|
create_flashinfer_kv_indices_triton[grid](
|
||||||
|
r2t,
|
||||||
|
req_pool,
|
||||||
|
lens,
|
||||||
|
indptr,
|
||||||
|
start,
|
||||||
|
kv_i,
|
||||||
|
r2t.shape[1],
|
||||||
|
ENTRY_PAGE_SIZE=entry_page_size,
|
||||||
|
TOKEN_BLOCK_PARALLEL=parallel,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
outs.append(kv_i)
|
||||||
|
for kv_i in outs[1:]:
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(outs[0], kv_i), (seqs, use_start, entry_page_size)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user