From 08cb081ed0bbb75d61233d343b52fb306801db74 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Tue, 21 Jul 2026 16:20:14 -0700 Subject: [PATCH] [Perf] Skip page-table columns past kv length in DSA draft-extend metadata kernel (#31981) --- .../kernels/ops/attention/dsa_metadata.py | 9 +++++++ test/registered/kernels/test_dsa_metadata.py | 24 +++++++++++++++---- 2 files changed, 29 insertions(+), 4 deletions(-) diff --git a/python/sglang/kernels/ops/attention/dsa_metadata.py b/python/sglang/kernels/ops/attention/dsa_metadata.py index caa8ed027..593c64bb4 100644 --- a/python/sglang/kernels/ops/attention/dsa_metadata.py +++ b/python/sglang/kernels/ops/attention/dsa_metadata.py @@ -524,6 +524,15 @@ def _fused_dsa_draft_extend_metadata_kernel( mask=req_row < bs, other=0, ).to(tl.int32) + # Skip column blocks past the request's kv length: no consumer reads there + # (attention and the indexer both stay within cache_seqlens). + kv_len = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=req_row < bs, + other=0, + ).to(tl.int32) + if col_block * BLOCK_N >= kv_len: + return if STATIC_EXTEND_LEN: prefix = req_row * qo_len else: diff --git a/test/registered/kernels/test_dsa_metadata.py b/test/registered/kernels/test_dsa_metadata.py index 9e62ede04..d78562360 100644 --- a/test/registered/kernels/test_dsa_metadata.py +++ b/test/registered/kernels/test_dsa_metadata.py @@ -304,10 +304,19 @@ 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. + 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) + _assert_equal(cache_seqlens, expected_cache, "draft cache_seqlens") _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "draft cu_seqlens_k") _assert_equal( - page_table_1[:total_len], expected_page_table, "draft page_table_1" + page_table_1[:total_len][live_mask], + expected_page_table[live_mask], + "draft page_table_1 (live [:kv_len] prefix)", ) _assert_equal( seqlens_expanded[:total_len], expected_expanded, "draft seqlens_expanded" @@ -321,10 +330,17 @@ class TestDSAMetadataKernels(CustomTestCase): "draft dsa_cu_seqlens_k", ) if real_page_size > 1: + # Real-page column real_col maps to source column real_col*real_page_size. + 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[:total_len], - _real_page_table(expected_page_table, real_page_size), - "draft real_page_table", + real_page_table[:total_len][real_live_mask], + expected_real[real_live_mask], + "draft real_page_table (live [:kv_len] prefix)", ) def test_decode_matches_eager_reference(self):