From 1b63155efead7494843f6db8681851ba94611f73 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 22 Jul 2026 22:00:25 -0700 Subject: [PATCH] [Perf] Skip blocks past per-request live length in full-width Triton kernels (#32109) --- .../ops/attention/dsa/index_buf_accessor.py | 3 + .../ops/attention/dsa/transform_index.py | 3 + .../kernels/ops/attention/dsa_metadata.py | 28 +++++++++ .../kernels/ops/attention/extend_attention.py | 8 +++ .../sglang/kernels/ops/attention/metadata.py | 32 ++++++++-- .../ops/kvcache/trtllm_mha_graph_metadata.py | 14 ++++- .../test_normal_decode_set_metadata.py | 45 +++++++++++++- .../test_trtllm_mha_graph_metadata.py | 24 +++++--- .../ops/attention/test_dsa_metadata.py | 58 +++++++++++++++---- 9 files changed, 187 insertions(+), 28 deletions(-) diff --git a/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py index 20d58409a..494ab948f 100644 --- a/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py +++ b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py @@ -639,6 +639,9 @@ def _get_k_and_s_triton_kernel( k_offsets = thread_idx * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K) seq_len = tl.load(seq_len_ptr + batch_id) + # Grid axis 1 spans the batch-max seq len; fully-masked blocks store nothing. + if block_token_start >= seq_len: + return token_valid_mask = token_ids < seq_len pre_batch_idx = tl.arange(0, seq_len_num_pow) diff --git a/python/sglang/kernels/ops/attention/dsa/transform_index.py b/python/sglang/kernels/ops/attention/dsa/transform_index.py index 9ae724580..2d7dd105e 100644 --- a/python/sglang/kernels/ops/attention/dsa/transform_index.py +++ b/python/sglang/kernels/ops/attention/dsa/transform_index.py @@ -87,6 +87,9 @@ def transform_index_page_table_prefill_kernel( query_start = tl.load(cu_seqlens_q_ptr + request_id) query_end = tl.load(cu_seqlens_q_ptr + request_id + 1) + # Grid axis 1 spans the batch-max extend len; fully-masked blocks store nothing. + if query_start + tl.program_id(1) * BLOCK_Q >= query_end: + return token_indices = query_start + query_offsets mask = (token_indices[:, None] < query_end) & (topk_offsets[None, :] < TOPK) diff --git a/python/sglang/kernels/ops/attention/dsa_metadata.py b/python/sglang/kernels/ops/attention/dsa_metadata.py index 593c64bb4..21a1fedb4 100644 --- a/python/sglang/kernels/ops/attention/dsa_metadata.py +++ b/python/sglang/kernels/ops/attention/dsa_metadata.py @@ -71,6 +71,16 @@ def _fused_dsa_decode_metadata_kernel( mask=row < bs, other=0, ) + # Skip column blocks past the request's kv length: no consumer reads there + # (attention and the indexer both stay within cache_seqlens). Loaded after + # req_idx so the two scalar loads pipeline (no added latency when live). + kv_len = tl.load( + seq_lens + row * seq_lens_stride, + mask=row < bs, + other=0, + ).to(tl.int32) + if col_block * BLOCK_N >= kv_len: + return vals = tl.load( req_to_token + req_idx * req_to_token_stride_0 + offs_n * req_to_token_stride_1, mask=mask, @@ -120,6 +130,10 @@ def fused_dsa_decode_metadata( where the wide table is never read (attention uses topk_indices, the indexer uses real_page_table); ``real_page_size`` must be >1 in that case. When a tensor is passed, behavior is unchanged (both tables are written). + + Contract: each page-table row is written only over its live prefix + ([:cache_seqlens]); the tail keeps stale values across CUDA-graph replays, so + consumers must bound reads by cache_seqlens. """ assert seq_lens.is_cuda assert req_pool_indices.is_cuda @@ -283,6 +297,20 @@ def _fused_dsa_target_verify_metadata_kernel( mask=out_row < expanded_size, other=0, ) + # Skip column blocks past the request's kv length (seq_len + next_n): no + # consumer reads there (attention and the indexer stay within cache_seqlens). + # Loaded after req_idx so the two scalar loads pipeline (no added latency + # when live). + kv_len = ( + tl.load( + seq_lens + req_row * seq_lens_stride, + mask=out_row < expanded_size, + other=0, + ).to(tl.int32) + + next_n + ) + if col_block * BLOCK_N >= kv_len: + return vals = tl.load( req_to_token + req_idx * req_to_token_stride_0 + offs_n * req_to_token_stride_1, mask=mask, diff --git a/python/sglang/kernels/ops/attention/extend_attention.py b/python/sglang/kernels/ops/attention/extend_attention.py index 4c456b8c9..c41da1d56 100644 --- a/python/sglang/kernels/ops/attention/extend_attention.py +++ b/python/sglang/kernels/ops/attention/extend_attention.py @@ -313,6 +313,10 @@ def _fwd_kernel( cur_seq_len_prefix = tl.load(kv_indptr + cur_seq + 1) - cur_seq_kv_start_idx cur_seq_len = cur_seq_len_prefix + cur_seq_len_extend + # Grid axis 2 spans the batch-max extend length; all stores are masked by mask_m. + if cur_block_m * BLOCK_M >= cur_seq_len_extend: + return + if USE_CUSTOM_MASK: cur_seq_mask_start_idx = tl.load(mask_indptr + cur_seq) @@ -911,6 +915,10 @@ def _fwd_kernel_unified( cur_seq_kv_len = tl.load(kv_indptr + cur_seq + 1) - cur_seq_kv_start_idx cur_seq_prefix_len = tl.load(prefix_lens + cur_seq) + # Grid axis 2 spans the batch-max extend length; the store is masked by mask_m. + if cur_block_m * BLOCK_M >= cur_seq_q_len: + return + # Load window start position for sliding window attention # This is the absolute position of the first key in the window (0 if no sliding window) cur_window_start = 0 diff --git a/python/sglang/kernels/ops/attention/metadata.py b/python/sglang/kernels/ops/attention/metadata.py index 461767ccf..858042733 100644 --- a/python/sglang/kernels/ops/attention/metadata.py +++ b/python/sglang/kernels/ops/attention/metadata.py @@ -213,13 +213,25 @@ def _fused_metadata_kernel_general( return i = pid_b + # Self-guard on the device-side seq_len: skip column chunks past the + # request's live pages (tails keep stale values the attention kernels + # never read past cache_seqlens). + seq_len = tl.load(seq_lens + i * seq_lens_stride_0).to(tl.int32) + if page_size == 1: + num_live_pages = seq_len + seq_len_delta + else: + num_live_pages = (seq_len + seq_len_delta + (1 << SHIFT) - 1) >> SHIFT + num_live_pages = tl.minimum(num_live_pages, max_seq_pages) + col_start = pid_c * BLOCK_COLS + if col_start >= num_live_pages: + return + # Load row index for this batch (all threads in block have same i) row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0) row_offset = row_idx * req_to_token_stride_0 - col_start = pid_c * BLOCK_COLS col_offsets = col_start + tl.arange(0, BLOCK_COLS) - mask = col_offsets < max_seq_pages + mask = col_offsets < num_live_pages # Compute column indices in the source tensor (token offset) if page_size == 1: @@ -303,13 +315,21 @@ def _fused_metadata_kernel_ps1_no_swa( return i = pid_b + # Self-guard on the device-side seq_len: skip column chunks past the + # request's live pages (tails keep stale values the attention kernels + # never read past cache_seqlens). + seq_len = tl.load(seq_lens + i * seq_lens_stride_0).to(tl.int32) + num_live_pages = tl.minimum(seq_len + seq_len_delta, max_seq_pages) + col_start = pid_c * BLOCK_COLS + if col_start >= num_live_pages: + return + # Load row index for this batch (all threads in block have same i) row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0) row_offset = row_idx * req_to_token_stride_0 - col_start = pid_c * BLOCK_COLS col_offsets = col_start + tl.arange(0, BLOCK_COLS) - mask = col_offsets < max_seq_pages + mask = col_offsets < num_live_pages # page_size = 1: col_idx = col_offsets rt_offsets = row_offset + col_offsets * req_to_token_stride_1 @@ -555,6 +575,10 @@ def normal_decode_set_metadata( 5. (optional) swa_page_table for sliding window attention Achieves ~5.2x speedup on H200 hardware for typical decode workloads. + + Contract: only the live prefix (cdiv(cache_seqlens, page_size) pages) of each + page_table / swa_page_table row is (re)written; the tail keeps stale values + across CUDA-graph replays, so consumers must bound reads by cache_seqlens. """ assert ( page_size > 0 and (page_size & (page_size - 1)) == 0 diff --git a/python/sglang/kernels/ops/kvcache/trtllm_mha_graph_metadata.py b/python/sglang/kernels/ops/kvcache/trtllm_mha_graph_metadata.py index 73a828658..b1a062fc2 100644 --- a/python/sglang/kernels/ops/kvcache/trtllm_mha_graph_metadata.py +++ b/python/sglang/kernels/ops/kvcache/trtllm_mha_graph_metadata.py @@ -75,9 +75,12 @@ def update_trtllm_mha_graph_metadata_kernel( row_out = page_table_ptr + pid.to(tl.int64) * page_table_stride if HAS_SWA: swa_row_out = swa_page_table_ptr + pid.to(tl.int64) * swa_page_table_stride - for i in range(tl.cdiv(max_seq_pages, PAGE_BLOCK)): + # Self-guard on the device-side seqlen: pages past cdiv(cache_seqlen, + # PAGE_SIZE) keep stale values the attention kernels never read. + num_live_pages = tl.minimum(tl.cdiv(seqlen, PAGE_SIZE), max_seq_pages) + for i in range(tl.cdiv(num_live_pages, PAGE_BLOCK)): page_idx = i * PAGE_BLOCK + tl.arange(0, PAGE_BLOCK) - mask = page_idx < max_seq_pages + mask = page_idx < num_live_pages token = tl.load( row_in + page_idx.to(tl.int64) * PAGE_SIZE, mask=mask, other=0 ) @@ -143,7 +146,12 @@ def update_trtllm_mha_graph_metadata( q_stride: int = 0, q_mode: int = Q_MODE_NONE, ): - """Launch the fused metadata update (one kernel for the whole replay init).""" + """Launch the fused metadata update (one kernel for the whole replay init). + + Contract: only the live prefix (cdiv(cache_seqlens, page_size) pages) of each + page_table / swa_page_table row is (re)written; the tail keeps stale values + across replays, so consumers must bound reads by cache_seqlens. + """ if bs == 0: return diff --git a/test/registered/attention/test_normal_decode_set_metadata.py b/test/registered/attention/test_normal_decode_set_metadata.py index 99d5eeef6..b5b28fa1b 100644 --- a/test/registered/attention/test_normal_decode_set_metadata.py +++ b/test/registered/attention/test_normal_decode_set_metadata.py @@ -56,6 +56,22 @@ def reference_normal_decode_set_metadata( swa_page_table[:, :max_seq_pages].copy_(swa_page_indices // page_size) +def page_table_live_mask( + seq_lens: torch.Tensor, + seq_len_delta: int, + page_size: int, + max_seq_pages: int, + width: int, +) -> torch.Tensor: + """Per-row live region of the page table: the fused kernel self-guards on + seq_len and leaves columns past pages(seq_len + delta) untouched.""" + live_pages = torch.clamp( + (seq_lens + seq_len_delta + page_size - 1) // page_size, max=max_seq_pages + ) + cols = torch.arange(width, device=seq_lens.device) + return cols.view(1, -1) < live_pages.view(-1, 1) + + @unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA") class TestNormalDecodeSetMetadata(CustomTestCase): """Test fused Triton kernel in normal_decode_set_metadata.""" @@ -235,14 +251,26 @@ class TestNormalDecodeSetMetadata(CustomTestCase): f"cu_seqlens_k mismatch. Expected:\n{ref_data['cu_seqlens_k']}\nGot:\n{test_data['cu_seqlens_k']}", ) + live_mask = page_table_live_mask( + test_data["seq_lens"], + test_data["seq_len_delta"], + page_size, + test_data["max_seq_pages"], + test_data["page_table"].shape[1], + ) self.assertTrue( - torch.equal(test_data["page_table"], ref_data["page_table"]), + torch.equal( + test_data["page_table"][live_mask], ref_data["page_table"][live_mask] + ), f"page_table mismatch at bs={batch_size}, page_size={page_size}", ) if has_swa: self.assertTrue( - torch.equal(test_data["swa_page_table"], ref_data["swa_page_table"]), + torch.equal( + test_data["swa_page_table"][live_mask], + ref_data["swa_page_table"][live_mask], + ), f"swa_page_table mismatch at bs={batch_size}, page_size={page_size}", ) @@ -412,7 +440,18 @@ class TestNormalDecodeSetMetadata(CustomTestCase): self.assertTrue( torch.equal(test_data["cu_seqlens_k"], ref_data["cu_seqlens_k"]) ) - self.assertTrue(torch.equal(test_data["page_table"], ref_data["page_table"])) + live_mask = page_table_live_mask( + test_data["seq_lens"], + 0, + page_size, + test_data["max_seq_pages"], + test_data["page_table"].shape[1], + ) + self.assertTrue( + torch.equal( + test_data["page_table"][live_mask], ref_data["page_table"][live_mask] + ) + ) if __name__ == "__main__": diff --git a/test/registered/attention/test_trtllm_mha_graph_metadata.py b/test/registered/attention/test_trtllm_mha_graph_metadata.py index 780162d33..fa28fb22b 100644 --- a/test/registered/attention/test_trtllm_mha_graph_metadata.py +++ b/test/registered/attention/test_trtllm_mha_graph_metadata.py @@ -266,10 +266,8 @@ def _make_swa_mapping(pool_token_cap, seed): @pytest.mark.parametrize("q_mode", [Q_MODE_NONE, Q_MODE_CUMSUM, Q_MODE_STRIDED]) @pytest.mark.parametrize("with_swa", [False, True]) # static_width=True exercises the production path: the backend passes the STATIC -# max_num_pages (full upper bound), not a per-batch dynamic width, so the kernel -# rewrites the whole page-table width every replay (tail pages beyond a request's -# seq are gathered but ignored by the attention kernel via cache_seqlens). The -# aten reference gathers the same full width, so this stays a bit-exact check. +# max_num_pages upper bound, not a per-batch dynamic width. The kernel self-guards +# on the device-side seqlen, so the checks compare only the live prefix per row. @pytest.mark.parametrize("static_width", [False, True]) def test_metadata_correctness(bs, seqlen_offset, q_mode, with_swa, static_width): if not torch.cuda.is_available(): @@ -372,9 +370,18 @@ def test_metadata_correctness(bs, seqlen_offset, q_mode, with_swa, static_width) cu_k_ref[1:] = torch.cumsum(cache_seqlens_ref, dim=0, dtype=torch.int32) torch.testing.assert_close(cu_seqlens_k, cu_k_ref, rtol=0, atol=0) - # ---- page_table ---- + # ---- page_table (live [:pages(cache_seqlen)] prefix per row) ---- pt_ref, gathered = _ref_page_table(req_to_token, req_pool_indices, max_seq_pages) - torch.testing.assert_close(page_table[:, :max_seq_pages], pt_ref, rtol=0, atol=0) + live_pages = torch.clamp( + (cache_seqlens_ref.to(torch.int64) + PAGE_SIZE - 1) // PAGE_SIZE, + max=max_seq_pages, + ) + live_mask = torch.arange(max_seq_pages, device=DEVICE).view( + 1, -1 + ) < live_pages.view(-1, 1) + torch.testing.assert_close( + page_table[:, :max_seq_pages][live_mask], pt_ref[live_mask], rtol=0, atol=0 + ) # ---- cu_seqlens_q ---- if q_mode == Q_MODE_CUMSUM: @@ -392,7 +399,10 @@ def test_metadata_correctness(bs, seqlen_offset, q_mode, with_swa, static_width) if with_swa: swa_pt_ref = _ref_swa_page_table(gathered, swa_mapping) torch.testing.assert_close( - swa_page_table[:, :max_seq_pages], swa_pt_ref, rtol=0, atol=0 + swa_page_table[:, :max_seq_pages][live_mask], + swa_pt_ref[live_mask], + rtol=0, + atol=0, ) # swa_out_cache_loc reference: translate real prefix, zero-fill tail. diff --git a/test/registered/kernels/ops/attention/test_dsa_metadata.py b/test/registered/kernels/ops/attention/test_dsa_metadata.py index d78562360..3b7a985d5 100644 --- a/test/registered/kernels/ops/attention/test_dsa_metadata.py +++ b/test/registered/kernels/ops/attention/test_dsa_metadata.py @@ -102,18 +102,35 @@ class TestDSAMetadataKernels(CustomTestCase): expected_page_table = req_to_token[req_pool_indices, :max_len].contiguous() expected_dsa = _dsa_seqlens(expected_cache, dsa_index_topk) + # Compare only the live prefix [:seq_len]: whole blocks starting past + # the kv length are skipped (keep stale values), while the last + # partially live block still writes lanes past it -- the tail is + # unspecified either way, and consumers never read past cache_seqlens. + cols = torch.arange(max_len, dtype=torch.int32, device=self.device) + live_mask = cols.view(1, -1) < expected_cache.view(-1, 1) + _assert_equal(cache_seqlens, expected_cache, "decode cache_seqlens") _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "decode cu_seqlens_k") - _assert_equal(page_table_1, expected_page_table, "decode page_table_1") + _assert_equal( + page_table_1[live_mask], + expected_page_table[live_mask], + "decode page_table_1 (live [:seq_len] prefix)", + ) _assert_equal(dsa_cache_seqlens, expected_dsa, "decode dsa_cache_seqlens") _assert_equal( dsa_cu_seqlens_k, _cu_seqlens(expected_dsa), "decode dsa_cu_seqlens_k" ) if real_page_size > 1: + real_width = real_page_table.shape[1] + real_cols = torch.arange(real_width, dtype=torch.int32, device=self.device) + real_live_mask = ( + real_cols.view(1, -1) * real_page_size + ) < expected_cache.view(-1, 1) + expected_real = _real_page_table(expected_page_table, real_page_size) _assert_equal( - real_page_table, - _real_page_table(expected_page_table, real_page_size), - "decode real_page_table", + real_page_table[real_live_mask], + expected_real[real_live_mask], + "decode real_page_table (live [:seq_len] prefix)", ) def _check_target_verify( @@ -194,19 +211,37 @@ class TestDSAMetadataKernels(CustomTestCase): expected_expanded = expected_expanded.reshape(-1).contiguous() expected_dsa = _dsa_seqlens(expected_expanded, dsa_index_topk) + # Compare only the live prefix [:seq_len + next_n] per expanded row: + # whole blocks starting past the kv length are skipped, the last + # partially live block may still write past it -- the tail is + # unspecified, and consumers never read past cache_seqlens. + row_kv_lens = torch.repeat_interleave(expected_cache, next_n) + cols = torch.arange(max_seqlen_k, dtype=torch.int32, device=self.device) + live_mask = cols.view(1, -1) < row_kv_lens.view(-1, 1) + _assert_equal(cache_seqlens, expected_cache, "target cache_seqlens") _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "target cu_seqlens_k") - _assert_equal(page_table_1, expected_page_table, "target page_table_1") + _assert_equal( + page_table_1[live_mask], + expected_page_table[live_mask], + "target page_table_1 (live [:seq_len + next_n] prefix)", + ) _assert_equal(seqlens_expanded, expected_expanded, "target seqlens_expanded") _assert_equal(dsa_cache_seqlens, expected_dsa, "target dsa_cache_seqlens") _assert_equal( dsa_cu_seqlens_k, _cu_seqlens(expected_dsa), "target dsa_cu_seqlens_k" ) if real_page_size > 1: + real_width = real_page_table.shape[1] + real_cols = torch.arange(real_width, dtype=torch.int32, device=self.device) + real_live_mask = ( + real_cols.view(1, -1) * real_page_size + ) < row_kv_lens.view(-1, 1) + expected_real = _real_page_table(expected_page_table, real_page_size) _assert_equal( - real_page_table, - _real_page_table(expected_page_table, real_page_size), - "target real_page_table", + real_page_table[real_live_mask], + expected_real[real_live_mask], + "target real_page_table (live [:seq_len + next_n] prefix)", ) if fill_ctx_lens: expected_ctx = expected_cache.view(bs, 1).expand(bs, next_n).contiguous() @@ -304,9 +339,10 @@ class TestDSAMetadataKernels(CustomTestCase): ) expected_dsa = _dsa_seqlens(expected_expanded, dsa_index_topk) - # Only the live prefix [:kv_len] per request is defined; the kernel - # leaves columns past kv_len untouched. All expanded rows of a request - # share its kv length. + # Compare only the live prefix [:kv_len]: whole blocks starting past + # kv_len are skipped, the last partially live block may still write + # past it -- the tail is unspecified, and consumers never read past + # cache_seqlens. All expanded rows of a request share its kv length. row_kv_lens = torch.repeat_interleave(seq_lens.to(torch.int32), extend_seq_lens) cols = torch.arange(max_seqlen_k, dtype=torch.int32, device=self.device) live_mask = cols.view(1, -1) < row_kv_lens.view(-1, 1)