diff --git a/python/sglang/kernels/ops/attention/utils.py b/python/sglang/kernels/ops/attention/utils.py index baa7dd27a..3833847b5 100644 --- a/python/sglang/kernels/ops/attention/utils.py +++ b/python/sglang/kernels/ops/attention/utils.py @@ -48,6 +48,9 @@ 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, ) +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 ( fused_qk_rope_reshape_and_cache as fused_qk_rope_reshape_and_cache, ) diff --git a/python/sglang/kernels/ops/kvcache/kv_indices.py b/python/sglang/kernels/ops/kvcache/kv_indices.py index fc6d6879c..e11c1b7cd 100644 --- a/python/sglang/kernels/ops/kvcache/kv_indices.py +++ b/python/sglang/kernels/ops/kvcache/kv_indices.py @@ -4,6 +4,19 @@ import triton.language as tl _FLASHMLA_CREATE_KV_BLOCK_SIZE = 4096 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 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). req_to_token_ptr_stride, ENTRY_PAGE_SIZE: tl.constexpr = 1, + TOKEN_BLOCK_PARALLEL: tl.constexpr = False, ): """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 as ``token = entry * ps + pos % ps``, exact because converting an id keeps 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 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 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) 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 offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE mask = offset < kv_end - kv_start diff --git a/python/sglang/kernels/ops/speculative/cache_locs.py b/python/sglang/kernels/ops/speculative/cache_locs.py index ec1953d66..c8aad1ab0 100644 --- a/python/sglang/kernels/ops/speculative/cache_locs.py +++ b/python/sglang/kernels/ops/speculative/cache_locs.py @@ -67,13 +67,30 @@ def generate_draft_decode_kv_indices( iter_upper: tl.constexpr, num_tokens_upper: tl.constexpr, page_size: tl.constexpr, + NUM_STEPS: tl.constexpr = 0, ): - BLOCK_SIZE: tl.constexpr = 128 - iters = tl.program_id(axis=0) + # Optional token-block parallelism (NUM_STEPS > 0): the first grid axis + # 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) topk_id = tl.program_id(axis=2) - num_steps = tl.num_programs(axis=0) + if NUM_STEPS == 0: + iters = pid0 + 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) topk = tl.num_programs(axis=2) @@ -81,56 +98,90 @@ def generate_draft_decode_kv_indices( kv_indptr += kv_indptr_stride * iters iters += 1 - load_offset = tl.arange(0, bs_upper) - seq_lens = tl.load(paged_kernel_lens + load_offset, mask=load_offset < bid, other=0) - seq_len = tl.load(paged_kernel_lens + bid) - cum_seq_len = tl.sum(seq_lens) + if NUM_STEPS == 0: + load_offset = tl.arange(0, bs_upper) + seq_lens = tl.load( + paged_kernel_lens + load_offset, mask=load_offset < bid, other=0 + ) + seq_len = tl.load(paged_kernel_lens + bid) + 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 kv_offset = cum_seq_len * topk + bid * iters * topk + topk_id * (seq_len + iters) kv_ptr = kv_indices + kv_offset token_pool_ptr = req_to_token + tl.load(req_pool_indices + bid) * pool_len - kv_offset = tl.arange(0, BLOCK_SIZE) - num_loop = tl.cdiv(seq_len, BLOCK_SIZE) - for _ in range(num_loop): - mask = kv_offset < seq_len - data = tl.load(token_pool_ptr + kv_offset, mask=mask) - tl.store(kv_ptr + kv_offset, data, mask=mask) - kv_offset += BLOCK_SIZE - - extend_offset = tl.arange(0, iter_upper) - if page_size == 1 or topk == 1: - extend_data = tl.load( - token_pool_ptr + seq_len + topk_id * num_steps + tl.arange(0, iter_upper), - mask=extend_offset < iters, - ) + if NUM_STEPS == 0: + kv_offset = tl.arange(0, BLOCK_SIZE) + num_loop = tl.cdiv(seq_len, BLOCK_SIZE) + for _ in range(num_loop): + mask = kv_offset < seq_len + data = tl.load(token_pool_ptr + kv_offset, mask=mask) + tl.store(kv_ptr + kv_offset, data, mask=mask) + kv_offset += BLOCK_SIZE else: - prefix_len = seq_len - last_page_len = prefix_len % page_size - num_new_pages_per_topk = ( - last_page_len + num_steps + page_size - 1 - ) // page_size - prefix_base = seq_len // page_size * page_size - start = ( - prefix_base + topk_id * num_new_pages_per_topk * page_size + last_page_len - ) - extend_data = tl.load( - token_pool_ptr + start + extend_offset, + 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) + if page_size == 1 or topk == 1: + extend_data = tl.load( + token_pool_ptr + + seq_len + + topk_id * num_steps + + tl.arange(0, iter_upper), + mask=extend_offset < iters, + ) + else: + prefix_len = seq_len + last_page_len = prefix_len % page_size + num_new_pages_per_topk = ( + last_page_len + num_steps + page_size - 1 + ) // page_size + prefix_base = seq_len // page_size * page_size + start = ( + prefix_base + + topk_id * num_new_pages_per_topk * page_size + + last_page_len + ) + extend_data = tl.load( + token_pool_ptr + start + extend_offset, + 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 + bs_offset = tl.arange(0, num_tokens_upper) - # Update kv_indptr - bs_offset = tl.arange(0, num_tokens_upper) - - zid = bid * topk + topk_id - if zid == 0: - zid = num_seqs * topk - positions = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0) - base = tl.sum(positions) - tl.store(kv_indptr + zid, base + zid * iters) + zid = bid * topk + topk_id + if zid == 0: + zid = num_seqs * topk + pos_vals = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0) + base = tl.sum(pos_vals) + tl.store(kv_indptr + zid, base + zid * iters) @triton.jit diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 7aa557185..93188aceb 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -23,6 +23,7 @@ from sglang.kernels.ops.attention.utils import ( create_flashinfer_kv_indices_triton, create_flashmla_kv_indices_triton, get_num_kv_index_blocks_flashmla, + kv_indices_num_token_blocks, ) from sglang.kernels.ops.kvcache.aiter_unified_attention import ( scatter_ragged_to_page_table_kernel, @@ -124,6 +125,13 @@ fast_mode = 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): SLIDING_WINDOW = auto() CROSS_ATTENTION = auto() @@ -1170,6 +1178,11 @@ class AiterAttnBackend(AttentionBackend): ) 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( self, forward_batch: ForwardBatch, @@ -1374,7 +1387,8 @@ class AiterAttnBackend(AttentionBackend): 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, forward_batch.req_pool_indices, forward_batch.seq_lens, @@ -1382,6 +1396,7 @@ class AiterAttnBackend(AttentionBackend): None, kv_indices, self.req_to_token.stride(0), + TOKEN_BLOCK_PARALLEL=num_token_blocks > 1, ) if _use_mla_ps_kernel: @@ -1467,7 +1482,8 @@ class AiterAttnBackend(AttentionBackend): kv_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, forward_batch.req_pool_indices, kv_lens, @@ -1475,6 +1491,7 @@ class AiterAttnBackend(AttentionBackend): None, kv_indices, self.req_to_token.stride(0), + TOKEN_BLOCK_PARALLEL=num_token_blocks > 1, ) # if self.kv_cache_dtype == fp8_dtype: @@ -1564,7 +1581,8 @@ class AiterAttnBackend(AttentionBackend): kv_indices = torch.empty( 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, forward_batch.req_pool_indices, forward_batch.seq_lens, @@ -1572,6 +1590,7 @@ class AiterAttnBackend(AttentionBackend): None, kv_indices, self.req_to_token.stride(0), + TOKEN_BLOCK_PARALLEL=num_token_blocks > 1, ) custom_mask = spec_info.custom_mask @@ -2018,7 +2037,8 @@ class AiterAttnBackend(AttentionBackend): bs=bs, 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, req_pool_indices, kv_lens, @@ -2026,6 +2046,7 @@ class AiterAttnBackend(AttentionBackend): None, kv_indices, 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] @@ -2144,7 +2165,8 @@ class AiterAttnBackend(AttentionBackend): kv_indptr = self.kv_indptr[: bs + 1] kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) 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, req_pool_indices, seq_lens, @@ -2152,6 +2174,7 @@ class AiterAttnBackend(AttentionBackend): None, kv_indices, 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] @@ -3544,8 +3567,15 @@ class AiterMultiStepDraftBackend: bs = self.topk * num_seqs 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.speculative_num_steps, num_seqs, self.topk) + (self.speculative_num_steps * num_token_blocks, num_seqs, self.topk) ]( forward_batch.req_pool_indices, 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(bs), 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): diff --git a/test/registered/kernel/speculative/test_spec_kv_indices_grid.py b/test/registered/kernel/speculative/test_spec_kv_indices_grid.py new file mode 100644 index 000000000..d93ad83de --- /dev/null +++ b/test/registered/kernel/speculative/test_spec_kv_indices_grid.py @@ -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()