diff --git a/python/sglang/kernels/ops/attention/dsa_kpool_metadata/__init__.py b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/__init__.py new file mode 100644 index 000000000..fc84db50d --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/__init__.py @@ -0,0 +1 @@ +"""Opt-in KPool metadata kernels; ordinary DSA kernels remain unchanged.""" diff --git a/python/sglang/kernels/ops/attention/dsa_kpool_metadata/decode.py b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/decode.py new file mode 100644 index 000000000..0263dad48 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/decode.py @@ -0,0 +1,225 @@ +"""Pool-aware fused DSA decode metadata.""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.kernels.ops.attention.dsa_kpool_metadata.scan import bounded_scan_num_splits + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "max_len", + "num_splits", + ] +) +def _fused_dsa_decode_metadata_kernel( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + bs: tl.constexpr, + max_len, + num_splits, + dsa_index_topk: tl.constexpr, + index_kpool: tl.constexpr, + real_page_size: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + HAS_PAGE_TABLE_1: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + seq_i32 = seq.to(tl.int32) + if index_kpool <= 1: + dsa_seq = tl.minimum(seq_i32, dsa_index_topk) + else: + # Preserve the live partial pool after selecting pool-aligned history. + full_pool_tokens = (seq_i32 // index_kpool) * index_kpool + selected_history_tokens = tl.minimum(full_pool_tokens, dsa_index_topk) + tail_tokens = seq_i32 - full_pool_tokens + dsa_seq = selected_history_tokens + tail_tokens + + cu = tl.cumsum(seq_i32, 0) + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(cache_seqlens + offs_b, seq_i32, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + tl.store(dsa_cache_seqlens + offs_b, dsa_seq, mask=mask_b) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_b, dsa_cu, mask=mask_b) + return + + page_pid = pid - 1 + row = page_pid // num_splits + split_id = page_pid - row * num_splits + + req_idx = tl.load( + req_pool_indices + row * req_pool_indices_stride, + mask=row < bs, + other=0, + ) + kv_len = tl.load( + seq_lens + row * seq_lens_stride, + mask=row < bs, + other=0, + ).to(tl.int32) + # Page-table row offsets can overflow int32 at 1M context. + row_i64 = row.to(tl.int64) + num_live_blocks = tl.minimum(tl.cdiv(kv_len, BLOCK_N), tl.cdiv(max_len, BLOCK_N)) + # Three stages hide latency across strided copy iterations. + for col_block in tl.range(split_id, num_live_blocks, num_splits, num_stages=3): + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (row < bs) & (offs_n < max_len) + vals = tl.load( + req_to_token + + req_idx * req_to_token_stride_0 + + offs_n * req_to_token_stride_1, + mask=mask, + other=0, + ).to(tl.int32) + if HAS_PAGE_TABLE_1: + tl.store( + page_table_1 + + row_i64 * page_table_stride_0 + + offs_n * page_table_stride_1, + vals, + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + row_i64 * real_page_table_stride_0 + + real_cols * real_page_table_stride_1, + vals // real_page_size, + mask=real_mask, + ) + + +def fused_dsa_decode_metadata( + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: Optional[torch.Tensor], + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + max_len: int, + dsa_index_topk: int, + real_page_size: int, + index_kpool: int = 1, +) -> None: + """Fill decode-graph DSA metadata (seqlens + page tables) from req_to_token. + + ``page_table_1`` (the wide page_size=1 table) is optional: pass ``None`` to + skip materializing it and write only the compact ``real_page_table`` + (page_size=``real_page_size``). This is used by the fused decode CUDA graph, + 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. + + The column scan is bounded inside the kernel by each row's own kv length + (read at run time), so the cost scales with the live sequence lengths and + not with ``max_len`` (the table width); the grid itself stays + data-independent. See :func:`bounded_scan_num_splits`. + """ + assert seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + if bs == 0: + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + assert index_kpool > 0 + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + # page_size==1: real IS page_table_1, so page_table_1 must be present. + assert page_table_1 is not None + real_page_table = page_table_1 + + # page_table_1 (the wide page_size=1 table) may be dropped for the fused + # decode CUDA graph; the kernel then writes only real_page_table. + has_page_table_1 = page_table_1 is not None + if not has_page_table_1: + assert has_real_page_table + page_table_1 = real_page_table # dummy pointer for stride args + else: + assert page_table_1.is_cuda + + block_bs = triton.next_power_of_2(bs) + block_n = 128 + num_col_blocks = triton.cdiv(max_len, block_n) + num_splits = bounded_scan_num_splits(bs, num_col_blocks) + grid = (1 + bs * num_splits,) + + _fused_dsa_decode_metadata_kernel[grid]( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + bs, + max_len, + num_splits, + dsa_index_topk, + index_kpool, + real_page_size, + has_real_page_table, + has_page_table_1, + BLOCK_BS=block_bs, + BLOCK_N=block_n, + ) diff --git a/python/sglang/kernels/ops/attention/dsa_kpool_metadata/draft_extend.py b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/draft_extend.py new file mode 100644 index 000000000..8dc75674c --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/draft_extend.py @@ -0,0 +1,308 @@ +"""Pool-aware fused DSA draft extend metadata.""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.kernels.ops.attention.dsa_kpool_metadata.scan import bounded_scan_num_splits + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "total_len", + "max_seqlen_k", + "num_splits", + ] +) +def _fused_dsa_draft_extend_metadata_kernel( + seq_lens, + extend_seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens_stride: tl.constexpr, + extend_seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + bs: tl.constexpr, + total_len, + max_seqlen_k, + num_splits, + dsa_index_topk: tl.constexpr, + index_kpool: tl.constexpr, + real_page_size: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + HAS_PAGE_TABLE_1: tl.constexpr, + STATIC_EXTEND_LEN: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_EXPANDED: tl.constexpr, + BLOCK_ROWS: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + cache_seq = seq.to(tl.int32) + cu = tl.cumsum(cache_seq, 0) + + tl.store(cache_seqlens + offs_b, cache_seq, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + + offs_e = tl.arange(0, BLOCK_EXPANDED) + mask_e = offs_e < total_len + if STATIC_EXTEND_LEN: + static_qo_len = tl.load(extend_seq_lens).to(tl.int32) + req_row = offs_e // static_qo_len + local_off = offs_e - req_row * static_qo_len + qo_len_for_row = tl.zeros((BLOCK_EXPANDED,), tl.int32) + static_qo_len + else: + req_row = tl.full((BLOCK_EXPANDED,), 0, tl.int32) + local_off = tl.full((BLOCK_EXPANDED,), 0, tl.int32) + qo_len_for_row = tl.full((BLOCK_EXPANDED,), 1, tl.int32) + prefix = tl.full((), 0, tl.int32) + + for i in tl.range(0, bs): + qo_len = tl.load(extend_seq_lens + i * extend_seq_lens_stride).to( + tl.int32 + ) + in_row = (offs_e >= prefix) & (offs_e < prefix + qo_len) + req_row = tl.where(in_row, i, req_row) + local_off = tl.where(in_row, offs_e - prefix, local_off) + qo_len_for_row = tl.where(in_row, qo_len, qo_len_for_row) + prefix += qo_len + + base_seq = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=mask_e, + other=0, + ).to(tl.int32) + # Clamp to >= 0: DP-padded / idle-companion rows carry the CUDA-graph + # seq_len fill value (1), which is smaller than qo_len, so the raw + # per-row visible kv length goes negative. Consumers treat these + # lengths as unsigned (the top-k v2 kernel reads them as uint32), so a + # negative row becomes a ~4e9-token length and an illegal memory + # access. 0 keeps padded rows on the trivial all-(-1) output path. + expanded_seq = base_seq - qo_len_for_row + local_off + 1 + expanded_seq = tl.maximum(expanded_seq, 0) + expanded_seq = tl.where(mask_e, expanded_seq, 0) + if index_kpool <= 1: + dsa_seq = tl.minimum(expanded_seq, dsa_index_topk) + else: + # Preserve the live partial pool after selecting pool-aligned history. + full_pool_tokens = (expanded_seq // index_kpool) * index_kpool + selected_history_tokens = tl.minimum(full_pool_tokens, dsa_index_topk) + tail_tokens = expanded_seq - full_pool_tokens + dsa_seq = selected_history_tokens + tail_tokens + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(seqlens_expanded + offs_e, expanded_seq, mask=mask_e) + tl.store(dsa_cache_seqlens + offs_e, dsa_seq, mask=mask_e) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_e, dsa_cu, mask=mask_e) + return + + page_pid = pid - 1 + req_row = page_pid // num_splits + split_id = page_pid - req_row * num_splits + + qo_len = tl.load( + extend_seq_lens + req_row * extend_seq_lens_stride, + mask=req_row < bs, + other=0, + ).to(tl.int32) + kv_len = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=req_row < bs, + other=0, + ).to(tl.int32) + # Bound the scan by the live replay-time kv length. + num_live_blocks = tl.minimum( + tl.cdiv(kv_len, BLOCK_N), tl.cdiv(max_seqlen_k, BLOCK_N) + ) + if split_id >= num_live_blocks: + return + if STATIC_EXTEND_LEN: + prefix = req_row * qo_len + else: + prefix = tl.full((), 0, tl.int32) + for i in tl.range(0, bs): + prev_qo_len = tl.load(extend_seq_lens + i * extend_seq_lens_stride).to( + tl.int32 + ) + prefix += tl.where(i < req_row, prev_qo_len, 0) + offs_r = tl.arange(0, BLOCK_ROWS) + out_rows = prefix + offs_r + row_mask = (req_row < bs) & (offs_r < qo_len) & (out_rows < total_len) + has_rows = (req_row < bs) & (qo_len > 0) + + req_idx = tl.load( + req_pool_indices + req_row * req_pool_indices_stride, + mask=has_rows, + other=0, + ) + # Output-row offsets can overflow int32 at 1M context. + out_rows_i64 = out_rows.to(tl.int64) + for col_block in tl.range(split_id, num_live_blocks, num_splits, num_stages=3): + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + col_mask = offs_n < max_seqlen_k + mask = row_mask[:, None] & col_mask[None, :] + + vals = tl.load( + req_to_token + + req_idx * req_to_token_stride_0 + + offs_n * req_to_token_stride_1, + mask=col_mask & has_rows, + other=0, + ).to(tl.int32) + if HAS_PAGE_TABLE_1: + tl.store( + page_table_1 + + out_rows_i64[:, None] * page_table_stride_0 + + offs_n[None, :] * page_table_stride_1, + vals[None, :], + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n[None, :] % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + out_rows_i64[:, None] * real_page_table_stride_0 + + real_cols[None, :] * real_page_table_stride_1, + (vals // real_page_size)[None, :], + mask=real_mask, + ) + + +def fused_dsa_draft_extend_metadata( + seq_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: Optional[torch.Tensor], + seqlens_expanded: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + total_len: int, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + max_extend_len: int, + max_total_len: int, + static_extend_len: bool = False, + index_kpool: int = 1, +) -> None: + assert seq_lens.is_cuda + assert extend_seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert seqlens_expanded.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + if bs == 0: + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + if total_len == 0: + cache = seq_lens.to(torch.int32) + cache_seqlens.copy_(cache) + cu_seqlens_k[:1].zero_() + cu_seqlens_k[1 : bs + 1].copy_(torch.cumsum(cache, dim=0, dtype=torch.int32)) + dsa_cu_seqlens_k[:1].zero_() + return + assert total_len <= max_total_len + # Caller-owned graph metadata guarantees each request accepts at most + # max_extend_len tokens. Avoid checking extend_seq_lens.max() here because + # that would sync in the replay hot path. + assert max_extend_len > 0 + assert total_len <= bs * max_extend_len + assert index_kpool > 0 + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + assert page_table_1 is not None + real_page_table = page_table_1 + + # page_table_1 (the wide page_size=1 table) may be dropped for the fused + # decode CUDA graph; the kernel then writes only real_page_table. + has_page_table_1 = page_table_1 is not None + if not has_page_table_1: + assert has_real_page_table + page_table_1 = real_page_table # dummy pointer for stride args + else: + assert page_table_1.is_cuda + + block_bs = triton.next_power_of_2(bs) + block_expanded = triton.next_power_of_2(max_total_len) + block_rows = triton.next_power_of_2(max_extend_len) + block_n = 128 + num_col_blocks = triton.cdiv(max_seqlen_k, block_n) + num_splits = bounded_scan_num_splits(bs, num_col_blocks) + grid = (1 + bs * num_splits,) + + _fused_dsa_draft_extend_metadata_kernel[grid]( + seq_lens, + extend_seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens.stride(0), + extend_seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + bs, + total_len, + max_seqlen_k, + num_splits, + dsa_index_topk, + index_kpool, + real_page_size, + has_real_page_table, + has_page_table_1, + static_extend_len, + BLOCK_BS=block_bs, + BLOCK_EXPANDED=block_expanded, + BLOCK_ROWS=block_rows, + BLOCK_N=block_n, + ) diff --git a/python/sglang/kernels/ops/attention/dsa_kpool_metadata/scan.py b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/scan.py new file mode 100644 index 000000000..91f64c7f6 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/scan.py @@ -0,0 +1,9 @@ +"""Capture-safe launch bounds shared by KPool metadata kernels.""" + +_TILE_PROGRAM_TARGET = 8192 + + +def bounded_scan_num_splits(rows: int, num_col_blocks: int) -> int: + """Keep the grid capture-safe while bounding traversal by replay-time data.""" + assert rows > 0 + return max(1, min(num_col_blocks, _TILE_PROGRAM_TARGET // rows)) diff --git a/python/sglang/kernels/ops/attention/dsa_kpool_metadata/verify.py b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/verify.py new file mode 100644 index 000000000..0d49743ed --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsa_kpool_metadata/verify.py @@ -0,0 +1,313 @@ +"""Pool-aware fused DSA verify metadata.""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.kernels.ops.attention.dsa_kpool_metadata.scan import bounded_scan_num_splits + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "max_seqlen_k", + "num_splits", + ] +) +def _fused_dsa_target_verify_metadata_kernel( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + paged_mqa_ctx_lens_2d, + seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + paged_mqa_ctx_lens_stride_0: tl.constexpr, + paged_mqa_ctx_lens_stride_1: tl.constexpr, + bs: tl.constexpr, + max_seqlen_k, + num_splits, + dsa_index_topk: tl.constexpr, + index_kpool: tl.constexpr, + real_page_size: tl.constexpr, + next_n: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + HAS_PAGED_MQA_CTX_LENS: tl.constexpr, + HAS_PAGE_TABLE_1: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_EXPANDED: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + expanded_size: tl.constexpr = bs * next_n + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + cache_seq = seq.to(tl.int32) + next_n + cu = tl.cumsum(cache_seq, 0) + + tl.store(cache_seqlens + offs_b, cache_seq, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + + offs_e = tl.arange(0, BLOCK_EXPANDED) + mask_e = offs_e < expanded_size + req_row = offs_e // next_n + draft_off = offs_e - req_row * next_n + base_seq = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=mask_e, + other=0, + ).to(tl.int32) + expanded_seq = base_seq + draft_off + 1 + expanded_seq = tl.where(mask_e, expanded_seq, 0) + if index_kpool <= 1: + dsa_seq = tl.minimum(expanded_seq, dsa_index_topk) + else: + # Preserve the live partial pool after selecting pool-aligned history. + full_pool_tokens = (expanded_seq // index_kpool) * index_kpool + selected_history_tokens = tl.minimum(full_pool_tokens, dsa_index_topk) + tail_tokens = expanded_seq - full_pool_tokens + dsa_seq = selected_history_tokens + tail_tokens + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(seqlens_expanded + offs_e, expanded_seq, mask=mask_e) + tl.store(dsa_cache_seqlens + offs_e, dsa_seq, mask=mask_e) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_e, dsa_cu, mask=mask_e) + + if HAS_PAGED_MQA_CTX_LENS: + tl.store( + paged_mqa_ctx_lens_2d + + req_row * paged_mqa_ctx_lens_stride_0 + + draft_off * paged_mqa_ctx_lens_stride_1, + base_seq + next_n, + mask=mask_e, + ) + return + + page_pid = pid - 1 + out_row = page_pid // num_splits + split_id = page_pid - out_row * num_splits + + req_row = out_row // next_n + req_idx = tl.load( + req_pool_indices + req_row * req_pool_indices_stride, + mask=out_row < expanded_size, + other=0, + ) + kv_len = ( + tl.load( + seq_lens + req_row * seq_lens_stride, + mask=out_row < expanded_size, + other=0, + ).to(tl.int32) + + next_n + ) + # Output-row offsets can overflow int32 at 1M context. + out_row_i64 = out_row.to(tl.int64) + num_live_blocks = tl.minimum( + tl.cdiv(kv_len, BLOCK_N), tl.cdiv(max_seqlen_k, BLOCK_N) + ) + for col_block in tl.range(split_id, num_live_blocks, num_splits, num_stages=3): + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (out_row < expanded_size) & (offs_n < max_seqlen_k) + vals = tl.load( + req_to_token + + req_idx * req_to_token_stride_0 + + offs_n * req_to_token_stride_1, + mask=mask, + other=0, + ).to(tl.int32) + if HAS_PAGE_TABLE_1: + tl.store( + page_table_1 + + out_row_i64 * page_table_stride_0 + + offs_n * page_table_stride_1, + vals, + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + out_row_i64 * real_page_table_stride_0 + + real_cols * real_page_table_stride_1, + vals // real_page_size, + mask=real_mask, + ) + + +def _prep_fused_dsa_target_verify_metadata_launch( + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: Optional[torch.Tensor], + seqlens_expanded: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + next_n: int, + paged_mqa_ctx_lens_2d: torch.Tensor = None, + index_kpool: int = 1, +): + assert seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert seqlens_expanded.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + assert bs > 0 + assert next_n > 0 + assert index_kpool > 0 + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + assert page_table_1 is not None + real_page_table = page_table_1 + + # page_table_1 (the wide page_size=1 table) may be dropped for the fused + # decode CUDA graph; the kernel then writes only real_page_table. + has_page_table_1 = page_table_1 is not None + if not has_page_table_1: + assert has_real_page_table + page_table_1 = real_page_table # dummy pointer for stride args + else: + assert page_table_1.is_cuda + + has_paged_mqa_ctx_lens = paged_mqa_ctx_lens_2d is not None + if has_paged_mqa_ctx_lens: + assert paged_mqa_ctx_lens_2d.is_cuda + assert paged_mqa_ctx_lens_2d.dtype == torch.int32 + assert paged_mqa_ctx_lens_2d.dim() == 2 + assert paged_mqa_ctx_lens_2d.size(0) == bs + assert paged_mqa_ctx_lens_2d.size(1) == next_n + else: + paged_mqa_ctx_lens_2d = page_table_1 + + expanded_size = bs * next_n + block_bs = triton.next_power_of_2(bs) + block_expanded = triton.next_power_of_2(expanded_size) + block_n = 128 + num_col_blocks = triton.cdiv(max_seqlen_k, block_n) + num_splits = bounded_scan_num_splits(expanded_size, num_col_blocks) + grid = (1 + expanded_size * num_splits,) + + args = ( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + paged_mqa_ctx_lens_2d, + seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + paged_mqa_ctx_lens_2d.stride(0) if has_paged_mqa_ctx_lens else 0, + paged_mqa_ctx_lens_2d.stride(1) if has_paged_mqa_ctx_lens else 0, + bs, + max_seqlen_k, + num_splits, + dsa_index_topk, + index_kpool, + real_page_size, + next_n, + has_real_page_table, + has_paged_mqa_ctx_lens, + has_page_table_1, + ) + constexprs = dict( + BLOCK_BS=block_bs, + BLOCK_EXPANDED=block_expanded, + BLOCK_N=block_n, + ) + return grid, args, constexprs + + +def fused_dsa_target_verify_metadata( + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: Optional[torch.Tensor], + seqlens_expanded: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + next_n: int, + paged_mqa_ctx_lens_2d: torch.Tensor = None, + index_kpool: int = 1, +) -> None: + if bs == 0: + assert cu_seqlens_k.is_cuda + assert dsa_cu_seqlens_k.is_cuda + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + + grid, args, constexprs = _prep_fused_dsa_target_verify_metadata_launch( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + bs, + max_seqlen_k, + dsa_index_topk, + real_page_size, + next_n, + paged_mqa_ctx_lens_2d, + index_kpool, + ) + _fused_dsa_target_verify_metadata_kernel[grid](*args, **constexprs) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 8e870fdc8..db80e167d 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1544,6 +1544,8 @@ class Envs: SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias( True, deprecated_name="SGLANG_NSA_FUSE_TOPK" ) + # Enabled for supported CUDA KPool geometry; set to 0 to use ordinary metadata. + SGLANG_EXPERIMENTAL_DSA_KPOOL_METADATA_FUSION = EnvBool(True) SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC = EnvBool(False) SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None) SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvIntWithAlias( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py index b6f1291a2..538d9a544 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py @@ -122,11 +122,9 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: """Precompute metadata for normal decode mode.""" max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] - if (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1: - from sglang.kernels.ops.attention.dsa_metadata import ( - fused_dsa_decode_metadata, - ) - + if ( + (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1 + ) or self.experimental_kpool_metadata_fusion: cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) page_indices = torch.empty( @@ -146,7 +144,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: real_page_table = None real_page_table_arg = page_indices - fused_dsa_decode_metadata( + self._fused_decode_metadata( seq_lens=seq_lens, req_pool_indices=req_pool_indices, req_to_token=self.req_to_token, @@ -245,11 +243,9 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] seqlens_expanded_size = bs * self.speculative_num_draft_tokens - if (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1: - from sglang.kernels.ops.attention.dsa_metadata import ( - fused_dsa_target_verify_metadata, - ) - + if ( + (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1 + ) or self.experimental_kpool_metadata_fusion: cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) page_indices = torch.empty( @@ -282,7 +278,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: real_page_table = None real_page_table_arg = page_indices - fused_dsa_target_verify_metadata( + self._fused_verify_metadata( seq_lens=seq_lens, req_pool_indices=req_pool_indices, req_to_token=self.req_to_token, diff --git a/python/sglang/srt/layers/attention/dsa/dsa_metadata_manager.py b/python/sglang/srt/layers/attention/dsa/dsa_metadata_manager.py new file mode 100644 index 000000000..bc783b6c6 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/dsa_metadata_manager.py @@ -0,0 +1,320 @@ +"""DSA metadata fusion selection and MTP replay reuse.""" + +from __future__ import annotations + +import logging +from functools import partial +from typing import TYPE_CHECKING + +from sglang.kernels.ops.attention.dsa_metadata import ( + fused_dsa_decode_metadata, + fused_dsa_draft_extend_metadata, + fused_dsa_target_verify_metadata, +) +from sglang.srt.environ import envs +from sglang.srt.utils import is_cuda, is_hip + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( + PrecomputedMetadata, + ) + from sglang.srt.layers.attention.dsa_backend import ( + DeepseekSparseAttnBackend, + DSAMetadata, + ) + from sglang.srt.model_executor.forward_batch_info import ForwardMode + +_is_hip = is_hip() + +logger = logging.getLogger(__name__) + + +def kpool_metadata_fusion_supported(pool_size, page_size, topk): + return ( + pool_size > 1 + and page_size == 64 + and page_size % pool_size == 0 + and topk % pool_size == 0 + ) + + +class DSAMetadataManagementMixin: + experimental_kpool_metadata_fusion = False + + def _init_kpool_metadata_fusion(self): + requested = envs.SGLANG_EXPERIMENTAL_DSA_KPOOL_METADATA_FUSION.get() + supported = kpool_metadata_fusion_supported( + self.dsa_index_kpool, self.real_page_size, self.dsa_index_topk + ) + self.experimental_kpool_metadata_fusion = ( + requested and supported and is_cuda() and not is_hip() + ) + self._fused_decode_metadata = fused_dsa_decode_metadata + self._fused_verify_metadata = fused_dsa_target_verify_metadata + self._fused_draft_extend_metadata = fused_dsa_draft_extend_metadata + if self.experimental_kpool_metadata_fusion: + from sglang.kernels.ops.attention.dsa_kpool_metadata.decode import ( + fused_dsa_decode_metadata as decode, + ) + from sglang.kernels.ops.attention.dsa_kpool_metadata.draft_extend import ( + fused_dsa_draft_extend_metadata as draft_extend, + ) + from sglang.kernels.ops.attention.dsa_kpool_metadata.verify import ( + fused_dsa_target_verify_metadata as verify, + ) + + self._fused_decode_metadata = partial( + decode, index_kpool=self.dsa_index_kpool + ) + self._fused_verify_metadata = partial( + verify, index_kpool=self.dsa_index_kpool + ) + self._fused_draft_extend_metadata = partial( + draft_extend, index_kpool=self.dsa_index_kpool + ) + logger.info( + "DSA KPool metadata fusion enabled (pool=%d)", self.dsa_index_kpool + ) + elif requested and self.dsa_index_kpool > 1: + logger.warning( + "DSA KPool metadata fusion unsupported for this platform/geometry; retaining ordinary metadata" + ) + + def _copy_base_replay_buffers(self, bs, metadata, precomputed, forward_mode): + # Track whether fused kernel succeeded + fused_kernel_succeeded = False + + # Use fused CUDA kernel for all copy operations + if not _is_hip: + try: + from sglang.kernels.ops.attention.fused_metadata_copy import ( + fused_metadata_copy_cuda, + ) + + # Map forward_mode to integer enum + if forward_mode.is_decode_or_idle(): + mode_int = 0 # DECODE + elif forward_mode.is_target_verify(): + mode_int = 1 # TARGET_VERIFY + else: + raise ValueError(f"Unsupported forward_mode: {forward_mode}") + + # Prepare FlashMLA tensors if needed + flashmla_num_splits_src = None + flashmla_num_splits_dst = None + flashmla_metadata_src = None + flashmla_metadata_dst = None + if precomputed.flashmla_metadata is not None: + flashmla_num_splits_src = precomputed.flashmla_metadata.num_splits + flashmla_num_splits_dst = metadata.flashmla_metadata.num_splits + flashmla_metadata_src = ( + precomputed.flashmla_metadata.flashmla_metadata + ) + flashmla_metadata_dst = metadata.flashmla_metadata.flashmla_metadata + + # Call fused kernel + fused_metadata_copy_cuda( + # Source tensors + precomputed.cache_seqlens, + precomputed.cu_seqlens_k, + precomputed.page_indices, + precomputed.dsa_cache_seqlens, + precomputed.seqlens_expanded, + precomputed.dsa_cu_seqlens_k, + precomputed.real_page_table, + flashmla_num_splits_src, + flashmla_metadata_src, + # Destination tensors + metadata.cache_seqlens_int32, + metadata.cu_seqlens_k, + metadata.page_table_1, + metadata.dsa_cache_seqlens_int32, + metadata.dsa_seqlens_expanded, + metadata.dsa_cu_seqlens_k, + ( + metadata.real_page_table + if precomputed.real_page_table is not None + else None + ), + flashmla_num_splits_dst, + flashmla_metadata_dst, + # Parameters + mode_int, + bs, + precomputed.max_len, + precomputed.max_seqlen_k, + precomputed.seqlens_expanded_size, + ) + + # Successfully used fused kernel + fused_kernel_succeeded = True + + except ImportError: + print( + "Warning: Fused metadata copy kernel not available, falling back to individual copies." + ) + except Exception as e: + print( + f"Warning: Fused metadata copy kernel failed with error: {e}, falling back to individual copies." + ) + + # Fallback to individual copy operations if the fused kernel is unavailable + # or fails at runtime. + if not fused_kernel_succeeded: + # Copy basic seqlens + metadata.cache_seqlens_int32.copy_(precomputed.cache_seqlens) + metadata.cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:]) + + # Mode-specific copy logic + if forward_mode.is_decode_or_idle(): + # Decode mode + metadata.page_table_1[:, : precomputed.max_len].copy_( + precomputed.page_indices + ) + metadata.dsa_cache_seqlens_int32.copy_(precomputed.dsa_cache_seqlens) + # seqlens_expanded is same as cache_seqlens (already copied) + + elif forward_mode.is_target_verify(): + # Target verify mode + metadata.page_table_1[:, : precomputed.max_seqlen_k].copy_( + precomputed.page_indices + ) + metadata.dsa_seqlens_expanded.copy_(precomputed.seqlens_expanded) + metadata.dsa_cache_seqlens_int32.copy_(precomputed.dsa_cache_seqlens) + + # Copy DSA cu_seqlens + size = precomputed.seqlens_expanded_size + metadata.dsa_cu_seqlens_k[1 : 1 + size].copy_( + precomputed.dsa_cu_seqlens_k[1 : 1 + size] + ) + + # Copy real page table + if precomputed.real_page_table is not None: + rows, cols = precomputed.real_page_table.shape + metadata.real_page_table[:rows, :cols].copy_( + precomputed.real_page_table + ) + + # Copy FlashMLA metadata in fallback path + if precomputed.flashmla_metadata is not None: + size = precomputed.seqlens_expanded_size + flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1)) + flashmla_metadata.copy_(precomputed.flashmla_metadata) + + @staticmethod + def _sibling_replay_metadata_compatible(dst: DSAMetadata, src: DSAMetadata) -> bool: + """Check that both sides expose the same optional derived buffers.""" + + def _match(a, b) -> bool: + return (a is None) == (b is None) + + if not ( + _match(dst.paged_mqa_schedule_metadata, src.paged_mqa_schedule_metadata) + and _match(dst.topk_v2_plan, src.topk_v2_plan) + and _match(dst.pooled_cache_seqlens_int32, src.pooled_cache_seqlens_int32) + and _match(dst.pooled_real_page_table, src.pooled_real_page_table) + and _match( + dst.pooled_paged_mqa_schedule_metadata, + src.pooled_paged_mqa_schedule_metadata, + ) + and _match(dst.kpool_write_plan, src.kpool_write_plan) + ): + return False + dst_plan, src_plan = dst.kpool_write_plan, src.kpool_write_plan + if dst_plan is not None and not ( + _match(dst_plan.pool_seqlens_per_q, src_plan.pool_seqlens_per_q) + and _match(dst_plan.seqlens_per_q, src_plan.seqlens_per_q) + and _match(dst_plan.pool_schedule_metadata, src_plan.pool_schedule_metadata) + and _match(dst_plan.effective_n_per_batch, src_plan.effective_n_per_batch) + ): + return False + return True + + def _copy_replay_metadata_from_sibling( + self, + src_backend: DeepseekSparseAttnBackend, + bs: int, + precomputed: PrecomputedMetadata, + forward_mode: ForwardMode, + ) -> None: + """Copy replay metadata from a sibling using the same precomputed input.""" + metadata = self.decode_cuda_graph_metadata.get(bs) + src_metadata = src_backend.decode_cuda_graph_metadata.get(bs) + if ( + # The derived-copy body below is CUDA-only; any other platform + # must take the full recompute, not a partial copy that would + # leave the DeepGEMM schedule / top-k plan / kpool metadata + # stale. + not is_cuda() + or _is_hip + or not forward_mode.is_decode_or_idle() + or metadata is None + or src_metadata is None + # `src_backend` must have run the full recompute path for this bs + # in this replay, so its derived buffers are fresh. + or src_backend.forward_metadata is not src_metadata + or not self._sibling_replay_metadata_compatible(metadata, src_metadata) + ): + self.init_forward_metadata_replay_cuda_graph_from_precomputed( + bs=bs, precomputed=precomputed, forward_mode=forward_mode + ) + return + + self.set_dsa_prefill_impl(forward_batch=None) + self._copy_base_replay_buffers(bs, metadata, precomputed, forward_mode) + + if is_cuda(): + if metadata.paged_mqa_schedule_metadata is not None: + metadata.paged_mqa_schedule_metadata.copy_( + src_metadata.paged_mqa_schedule_metadata + ) + if metadata.topk_v2_plan is not None: + metadata.topk_v2_plan.copy_(src_metadata.topk_v2_plan) + # Decode: the 2D ctx lens are a (bs, 1) view of this backend's own + # cache_seqlens_int32 (just refreshed by the base copy above); keep + # the exact refresh the recompute path performs -- it is a single + # small view/copy, not part of the duplicated derived work. + seqlens_32_2d = metadata.cache_seqlens_int32.contiguous().view(bs, 1) + if metadata.paged_mqa_ctx_lens_2d is None: + object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) + else: + metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d) + + self._copy_kpool_metadata_from_sibling(metadata, src_metadata) + + self.forward_metadata = metadata + + def _copy_kpool_metadata_from_sibling( + self, metadata: DSAMetadata, src_metadata: DSAMetadata + ) -> None: + """Copy KPool metadata derived from identical inputs from a sibling.""" + if self.dsa_index_kpool <= 1 or not is_cuda(): + return + + if metadata.pooled_cache_seqlens_int32 is not None: + metadata.pooled_cache_seqlens_int32.copy_( + src_metadata.pooled_cache_seqlens_int32 + ) + if metadata.pooled_real_page_table is not None: + metadata.pooled_real_page_table.copy_(src_metadata.pooled_real_page_table) + if metadata.pooled_paged_mqa_schedule_metadata is not None: + metadata.pooled_paged_mqa_schedule_metadata.copy_( + src_metadata.pooled_paged_mqa_schedule_metadata + ) + + dst_plan = metadata.kpool_write_plan + src_plan = src_metadata.kpool_write_plan + if dst_plan is None: + return + dst_plan.req.copy_(src_plan.req) + dst_plan.write_start.copy_(src_plan.write_start) + dst_plan.tail_logical_start.copy_(src_plan.tail_logical_start) + dst_plan.write_loc.copy_(src_plan.write_loc) + if dst_plan.pool_seqlens_per_q is not None: + dst_plan.pool_seqlens_per_q.copy_(src_plan.pool_seqlens_per_q) + if dst_plan.seqlens_per_q is not None: + dst_plan.seqlens_per_q.copy_(src_plan.seqlens_per_q) + if dst_plan.pool_schedule_metadata is not None: + dst_plan.pool_schedule_metadata.copy_(src_plan.pool_schedule_metadata) + if dst_plan.effective_n_per_batch is not None: + dst_plan.effective_n_per_batch.copy_(src_plan.effective_n_per_batch) diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 67c0725a0..8fdd7ba12 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -35,11 +35,6 @@ from sglang.kernels.ops.attention.dsa.transform_index import ( transform_index_page_table_decode, transform_index_page_table_prefill, ) -from sglang.kernels.ops.attention.dsa_metadata import ( - fused_dsa_decode_metadata, - fused_dsa_draft_extend_metadata, - fused_dsa_target_verify_metadata, -) from sglang.kernels.ops.attention.utils import ( concat_mla_absorb_q_general, mla_quantize_and_rope_for_fp8, @@ -50,8 +45,6 @@ from sglang.kernels.ops.attention.utils import ( from sglang.kernels.ops.kvcache.cache_ops import concat_and_cast_q_fp8_pad from sglang.srt.configs.model_config import ( get_dsa_index_kpool, - get_dsa_index_topk, - is_deepseek_dsa, ) from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend @@ -65,6 +58,9 @@ from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( compute_cu_seqlens, ) from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import DSAIndexerMetadata +from sglang.srt.layers.attention.dsa.dsa_metadata_manager import ( + DSAMetadataManagementMixin, +) from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( DSATopKBackend, TopkTransformMethod, @@ -88,7 +84,6 @@ from sglang.srt.layers.attention.trtllm_mla_backend import ( from sglang.srt.layers.cp.base import get_cp_strategy from sglang.srt.layers.cp.utils import is_cp_active from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_buffer, get_exec, get_parallel, get_spec from sglang.srt.utils import ( is_cuda, is_gfx95_supported, @@ -301,6 +296,7 @@ _DSA_IMPL_T: TypeAlias = Literal[ class DeepseekSparseAttnBackend( + DSAMetadataManagementMixin, DeepseekSparseAttnBackendKPoolMixin, DeepseekSparseAttnBackendMTPPrecomputeMixin, AttentionBackend, @@ -339,6 +335,7 @@ class DeepseekSparseAttnBackend( self.dsa_index_topk = get_dsa_index_topk(hf_config) self.dsa_index_kpool = get_dsa_index_kpool(hf_config) self.needs_cpu_seq_lens = self.dsa_index_kpool > 1 + self._init_kpool_metadata_fusion() self.max_context_len = model_runner.model_config.context_len self.num_q_heads = ( model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size @@ -1523,8 +1520,10 @@ class DeepseekSparseAttnBackend( # Normal Decode max_len = self._graph_page_table_width(metadata) - if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1: - fused_dsa_decode_metadata( + if ( + (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1 + ) or self.experimental_kpool_metadata_fusion: + self._fused_decode_metadata( seq_lens=seq_lens, req_pool_indices=req_pool_indices, req_to_token=self.req_to_token, @@ -1563,7 +1562,9 @@ class DeepseekSparseAttnBackend( elif forward_mode.is_target_verify(): max_seqlen_k = self._graph_page_table_width(metadata) - if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1: + if ( + (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1 + ) or self.experimental_kpool_metadata_fusion: paged_mqa_ctx_lens_2d = None if ( self.speculative_num_draft_tokens >= 2 @@ -1576,7 +1577,7 @@ class DeepseekSparseAttnBackend( ): paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d - fused_dsa_target_verify_metadata( + self._fused_verify_metadata( seq_lens=seq_lens, req_pool_indices=req_pool_indices, req_to_token=self.req_to_token, @@ -1659,8 +1660,10 @@ class DeepseekSparseAttnBackend( device=self.device, ) - if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1: - fused_dsa_draft_extend_metadata( + if ( + (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1 + ) or self.experimental_kpool_metadata_fusion: + self._fused_draft_extend_metadata( seq_lens=seq_lens, extend_seq_lens=extend_seq_lens, req_pool_indices=req_pool_indices, @@ -1811,125 +1814,7 @@ class DeepseekSparseAttnBackend( metadata = self.decode_cuda_graph_metadata[bs] - # Track whether fused kernel succeeded - fused_kernel_succeeded = False - - # Use fused CUDA kernel for all copy operations - if not _is_hip: - try: - from sglang.kernels.ops.attention.fused_metadata_copy import ( - fused_metadata_copy_cuda, - ) - - # Map forward_mode to integer enum - if forward_mode.is_decode_or_idle(): - mode_int = 0 # DECODE - elif forward_mode.is_target_verify(): - mode_int = 1 # TARGET_VERIFY - else: - raise ValueError(f"Unsupported forward_mode: {forward_mode}") - - # Prepare FlashMLA tensors if needed - flashmla_num_splits_src = None - flashmla_num_splits_dst = None - flashmla_metadata_src = None - flashmla_metadata_dst = None - if precomputed.flashmla_metadata is not None: - flashmla_num_splits_src = precomputed.flashmla_metadata.num_splits - flashmla_num_splits_dst = metadata.flashmla_metadata.num_splits - flashmla_metadata_src = ( - precomputed.flashmla_metadata.flashmla_metadata - ) - flashmla_metadata_dst = metadata.flashmla_metadata.flashmla_metadata - - # Call fused kernel - fused_metadata_copy_cuda( - # Source tensors - precomputed.cache_seqlens, - precomputed.cu_seqlens_k, - precomputed.page_indices, - precomputed.dsa_cache_seqlens, - precomputed.seqlens_expanded, - precomputed.dsa_cu_seqlens_k, - precomputed.real_page_table, - flashmla_num_splits_src, - flashmla_metadata_src, - # Destination tensors - metadata.cache_seqlens_int32, - metadata.cu_seqlens_k, - metadata.page_table_1, - metadata.dsa_cache_seqlens_int32, - metadata.dsa_seqlens_expanded, - metadata.dsa_cu_seqlens_k, - ( - metadata.real_page_table - if precomputed.real_page_table is not None - else None - ), - flashmla_num_splits_dst, - flashmla_metadata_dst, - # Parameters - mode_int, - bs, - precomputed.max_len, - precomputed.max_seqlen_k, - precomputed.seqlens_expanded_size, - ) - - # Successfully used fused kernel - fused_kernel_succeeded = True - - except ImportError: - print( - "Warning: Fused metadata copy kernel not available, falling back to individual copies." - ) - except Exception as e: - print( - f"Warning: Fused metadata copy kernel failed with error: {e}, falling back to individual copies." - ) - - # Fallback to individual copy operations if the fused kernel is unavailable - # or fails at runtime. - if not fused_kernel_succeeded: - # Copy basic seqlens - metadata.cache_seqlens_int32.copy_(precomputed.cache_seqlens) - metadata.cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:]) - - # Mode-specific copy logic - if forward_mode.is_decode_or_idle(): - # Decode mode - metadata.page_table_1[:, : precomputed.max_len].copy_( - precomputed.page_indices - ) - metadata.dsa_cache_seqlens_int32.copy_(precomputed.dsa_cache_seqlens) - # seqlens_expanded is same as cache_seqlens (already copied) - - elif forward_mode.is_target_verify(): - # Target verify mode - metadata.page_table_1[:, : precomputed.max_seqlen_k].copy_( - precomputed.page_indices - ) - metadata.dsa_seqlens_expanded.copy_(precomputed.seqlens_expanded) - metadata.dsa_cache_seqlens_int32.copy_(precomputed.dsa_cache_seqlens) - - # Copy DSA cu_seqlens - size = precomputed.seqlens_expanded_size - metadata.dsa_cu_seqlens_k[1 : 1 + size].copy_( - precomputed.dsa_cu_seqlens_k[1 : 1 + size] - ) - - # Copy real page table - if precomputed.real_page_table is not None: - rows, cols = precomputed.real_page_table.shape - metadata.real_page_table[:rows, :cols].copy_( - precomputed.real_page_table - ) - - # Copy FlashMLA metadata in fallback path - if precomputed.flashmla_metadata is not None: - size = precomputed.seqlens_expanded_size - flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1)) - flashmla_metadata.copy_(precomputed.flashmla_metadata) + self._copy_base_replay_buffers(bs, metadata, precomputed, forward_mode) # Refresh the schedule because stale shape decomposition can deadlock # DeepGEMM paged MQA. @@ -3788,6 +3673,20 @@ class DeepseekSparseAttnMultiStepBackend: forward_mode=ForwardMode.DECODE, ) + if self.attn_backends[0].experimental_kpool_metadata_fusion: + first = self.attn_backends[0] + first.init_forward_metadata_replay_cuda_graph_from_precomputed( + bs=bs, precomputed=precomputed, forward_mode=ForwardMode.DECODE + ) + for backend in self.attn_backends[1 : self.speculative_num_steps - 1]: + backend._copy_replay_metadata_from_sibling( + src_backend=first, + bs=bs, + precomputed=precomputed, + forward_mode=ForwardMode.DECODE, + ) + return + # Use multi-backend fused copy when we have 3 or more backends # This is 3x faster than calling the single-backend copy 3 times if self.speculative_num_steps > 3: diff --git a/python/sglang/test/kits/dsa_metadata_kit.py b/python/sglang/test/kits/dsa_metadata_kit.py new file mode 100644 index 000000000..a3e5c3629 --- /dev/null +++ b/python/sglang/test/kits/dsa_metadata_kit.py @@ -0,0 +1,147 @@ +"""Small real CUDA metadata fixtures; no model runner or model weights required.""" + +from dataclasses import fields +from types import SimpleNamespace + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsa.dsa_topk_backend import DSATopKBackend +from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.runtime_context import get_parallel + +BS, NEXT_N, WIDTH, TOPK, POOL = 4, 6, 131072, 2048, 4 +ROUNDS = ( + ([64, 128, 2048, 65530], [0, 2, 4, 6]), + ([63, 129, 65539, 100001], [7, 5, 3, 1]), + ([128, 61, 2051, 1023], [2, 6, 0, 4]), +) + + +def inputs(lengths, requests): + return ( + torch.tensor(lengths, dtype=torch.int64, device="cuda"), + torch.tensor(requests, dtype=torch.int64, device="cuda"), + ) + + +def make_backend(mode, seq, req, *, fusion=True): + backend = object.__new__(DeepseekSparseAttnBackend) + backend.device = torch.device("cuda") + backend.device_sm_major = torch.cuda.get_device_capability()[0] + backend.num_q_heads = 64 + backend.real_page_size = 64 + backend.dsa_index_topk = TOPK + backend.dsa_index_kpool = POOL + backend.speculative_num_draft_tokens = NEXT_N + backend.dsa_drop_wide_page_table = False + backend.dsa_decode_impl = "fa3" + backend.dsa_prefill_impl = "fa3" + backend.enable_auto_select_prefill_impl = False + backend.token_to_kv_pool = SimpleNamespace(slots_per_page=64) + # Only attention-dispatch state is synthetic; every metadata kernel is real. + backend._is_in_breakable_cuda_graph = lambda: False + backend._is_in_tc_piecewise_cuda_graph = lambda: False + backend._get_device_sm = lambda: backend.device_sm_major * 10 + backend._is_blackwell = lambda: backend.device_sm_major == 10 + backend.dsa_topk_backend = DSATopKBackend.SGL_KERNEL + backend.req_to_token = torch.arange( + 8 * WIDTH, device="cuda", dtype=torch.int32 + ).view(8, WIDTH) + backend._arange_buf = torch.arange( + BS * NEXT_N + 1, device="cuda", dtype=torch.int32 + ) + backend.decode_cuda_graph_metadata = { + "page_table": torch.zeros(BS * NEXT_N, WIDTH, device="cuda", dtype=torch.int32), + "cu_seqlens_q": backend._arange_buf, + } + with envs.SGLANG_EXPERIMENTAL_DSA_KPOOL_METADATA_FUSION.override(fusion): + backend._init_kpool_metadata_fusion() + with envs.SGLANG_OPT_USE_TOPK_V2.override(True): + apply_metadata(backend, mode, seq, req) + apply_metadata(backend, mode, seq, req) + return backend + + +def apply_metadata(backend, mode, seq, req, spec_info=None): + backend._apply_cuda_graph_metadata( + bs=BS, + req_pool_indices=req, + seq_lens=seq, + seq_lens_cpu=seq.cpu(), + forward_mode=mode, + spec_info=spec_info, + ) + + +def tensor_buffers(metadata): + result = {} + for field in fields(metadata): + value = getattr(metadata, field.name) + if isinstance(value, torch.Tensor): + result[field.name] = value + if metadata.kpool_write_plan is not None: + for field in fields(metadata.kpool_write_plan): + value = getattr(metadata.kpool_write_plan, field.name) + if isinstance(value, torch.Tensor): + result["kpool." + field.name] = value + return result + + +def addresses(metadata): + return { + name: tensor.data_ptr() for name, tensor in tensor_buffers(metadata).items() + } + + +def assert_metadata_equal(test, actual, expected): + actual_buffers, expected_buffers = tensor_buffers(actual), tensor_buffers(expected) + test.assertEqual(actual_buffers.keys(), expected_buffers.keys()) + for name, value in actual_buffers.items(): + reference = expected_buffers[name] + if name == "topk_v2_plan": + # Unused plan rows are intentionally uninitialized. Active rows are + # compacted by atomicAdd, so compare them in request order. + torch.testing.assert_close(value[0], reference[0]) + count = int(reference[0, 1].item()) + lhs, rhs = value[1 : count + 1], reference[1 : count + 1] + torch.testing.assert_close( + lhs[lhs[:, 0].argsort()], rhs[rhs[:, 0].argsort()] + ) + elif name in ("page_table_1", "real_page_table", "pooled_real_page_table"): + lengths = expected.cache_seqlens_int32 + lengths = lengths.repeat_interleave(value.shape[0] // lengths.numel()) + step = 1 if name == "page_table_1" else 64 + if name == "pooled_real_page_table": + step *= POOL + live = ( + torch.arange(value.shape[1], device=value.device)[None, :] * step + < lengths[:, None] + ) + torch.testing.assert_close(value[live], reference[live], msg=name) + else: + torch.testing.assert_close(value, reference, msg=name) + + +def capture_verify_metadata(backend, seq, req, *, dg_out_of_graph=False): + backend.ingraph_verify_metadata_enabled = True + backend.ingraph_verify_metadata_dg_out_of_graph = dg_out_of_graph + batch = SimpleNamespace( + batch_size=BS, + forward_mode=ForwardMode.TARGET_VERIFY, + seq_lens=seq, + req_pool_indices=req, + ) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with get_parallel().override(dcp_enabled=False): + with torch.cuda.stream(stream): + # Compile kernels and prime the allocator before capture. + backend.init_forward_metadata_in_graph(batch) + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + backend.init_forward_metadata_in_graph(batch) + torch.cuda.current_stream().wait_stream(stream) + return graph diff --git a/test/registered/kernel/attention/test_dsa_kpool_metadata_fusion.py b/test/registered/kernel/attention/test_dsa_kpool_metadata_fusion.py new file mode 100644 index 000000000..021067027 --- /dev/null +++ b/test/registered/kernel/attention/test_dsa_kpool_metadata_fusion.py @@ -0,0 +1,98 @@ +"""KPool fused metadata must retain live tails and refresh captured buffers.""" + +import unittest + +import torch + +from sglang.kernels.ops.attention.dsa_kpool_metadata.verify import ( + fused_dsa_target_verify_metadata, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +class TestKPoolMetadataFusion(CustomTestCase): + def test_verify_replay_boundaries_and_request_remapping(self): + device = "cuda" + bs, next_n, width, topk, pool_size = 4, 6, 16384, 2048, 4 + seq = torch.tensor([1, 61, 2047, 8191], device=device, dtype=torch.int64) + req = torch.tensor([3, 1, 6, 0], device=device, dtype=torch.int64) + table = torch.arange(8 * width, device=device, dtype=torch.int32).view(8, width) + + def empty(*shape): + return torch.full(shape, -1, device=device, dtype=torch.int32) + + buffers = dict( + cache_seqlens=empty(bs), + cu_seqlens_k=empty(bs + 1), + page_table_1=empty(bs * next_n, width), + seqlens_expanded=empty(bs * next_n), + dsa_cache_seqlens=empty(bs * next_n), + dsa_cu_seqlens_k=empty(bs * next_n + 1), + real_page_table=empty(bs * next_n, width // 64), + paged_mqa_ctx_lens_2d=empty(bs, next_n), + ) + addresses = {key: value.data_ptr() for key, value in buffers.items()} + + def refresh(): + fused_dsa_target_verify_metadata( + seq_lens=seq, + req_pool_indices=req, + req_to_token=table, + bs=bs, + max_seqlen_k=width, + dsa_index_topk=topk, + real_page_size=64, + next_n=next_n, + index_kpool=pool_size, + **buffers, + ) + + refresh() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + refresh() + for lengths, requests in [ + ([2, 63, 2048, 8193], [0, 6, 1, 3]), + ([64, 128, 2051, 8190], [5, 2, 7, 4]), + ([0, 3, 2045, 9000], [7, 0, 3, 2]), + ]: + seq.copy_(torch.tensor(lengths, device=device)) + req.copy_(torch.tensor(requests, device=device)) + graph.replay() + expanded = ( + seq[:, None] + torch.arange(1, next_n + 1, device=device) + ).flatten() + expected = torch.minimum(expanded, topk + expanded % pool_size).int() + torch.testing.assert_close(buffers["seqlens_expanded"], expanded.int()) + torch.testing.assert_close(buffers["dsa_cache_seqlens"], expected) + torch.testing.assert_close( + buffers["dsa_cu_seqlens_k"][1:], expected.cumsum(0).int() + ) + torch.testing.assert_close(buffers["cache_seqlens"], (seq + next_n).int()) + torch.testing.assert_close( + buffers["paged_mqa_ctx_lens_2d"], + (seq + next_n).int()[:, None].expand(bs, next_n), + ) + expected_pages = table[req].repeat_interleave(next_n, dim=0) + row_lens = (seq + next_n).repeat_interleave(next_n) + live = torch.arange(width, device=device)[None, :] < row_lens[:, None] + torch.testing.assert_close( + buffers["page_table_1"][live], expected_pages[live] + ) + real_live = ( + torch.arange(0, width, 64, device=device)[None, :] < row_lens[:, None] + ) + torch.testing.assert_close( + buffers["real_page_table"][real_live], + (expected_pages[:, ::64] // 64)[real_live], + ) + self.assertEqual( + addresses, {key: value.data_ptr() for key, value in buffers.items()} + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernel/attention/test_dsa_metadata_replay.py b/test/registered/kernel/attention/test_dsa_metadata_replay.py new file mode 100644 index 000000000..e47b7aefa --- /dev/null +++ b/test/registered/kernel/attention/test_dsa_metadata_replay.py @@ -0,0 +1,117 @@ +"""Fused KPool replay and MTP sibling copies preserve captured buffer identity.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.dsa_metadata_kit import ( + BS, + NEXT_N, + ROUNDS, + addresses, + apply_metadata, + assert_metadata_equal, + inputs, + make_backend, +) +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +class TestDSAMetadataReplay(CustomTestCase): + def test_fusion_matches_ordinary_metadata(self): + for mode in ( + ForwardMode.DECODE, + ForwardMode.TARGET_VERIFY, + ForwardMode.DRAFT_EXTEND_V2, + ): + with self.subTest(mode=mode): + seq, req = inputs(*ROUNDS[0]) + fused = make_backend(mode, seq, req) + ordinary = make_backend(mode, seq, req, fusion=False) + pointers = addresses(fused.forward_metadata) + for lengths, requests in ROUNDS: + seq.copy_(torch.tensor(lengths, device="cuda")) + req.copy_(torch.tensor(requests, device="cuda")) + spec = None + if mode.is_draft_extend_v2(): + spec = SimpleNamespace( + num_accept_tokens=torch.tensor( + [1, 2, 5, NEXT_N], device="cuda", dtype=torch.int32 + ) + ) + apply_metadata(fused, mode, seq, req, spec) + apply_metadata(ordinary, mode, seq, req, spec) + assert_metadata_equal( + self, fused.forward_metadata, ordinary.forward_metadata + ) + self.assertEqual(pointers, addresses(fused.forward_metadata)) + + def test_precomputed_verify_retains_live_tail(self): + mode = ForwardMode.TARGET_VERIFY + seq, req = inputs(*ROUNDS[0]) + fused = make_backend(mode, seq, req) + ordinary = make_backend(mode, seq, req, fusion=False) + pointers = addresses(fused.forward_metadata) + for lengths, requests in ROUNDS[1:]: + seq.copy_(torch.tensor(lengths, device="cuda")) + req.copy_(torch.tensor(requests, device="cuda")) + precomputed = fused._precompute_replay_metadata( + BS, req, seq, seq.cpu(), mode + ) + fused.init_forward_metadata_replay_cuda_graph_from_precomputed( + BS, precomputed, mode + ) + apply_metadata(ordinary, mode, seq, req) + assert_metadata_equal( + self, fused.forward_metadata, ordinary.forward_metadata + ) + self.assertEqual(pointers, addresses(fused.forward_metadata)) + + def test_precomputed_and_sibling_copy_refresh_derived_metadata(self): + mode = ForwardMode.DECODE + seq, req = inputs(*ROUNDS[0]) + source = make_backend(mode, seq, req) + sibling = make_backend(mode, seq, req) + ordinary = make_backend(mode, seq, req, fusion=False) + pointers = addresses(sibling.forward_metadata) + for lengths, requests in ROUNDS[1:]: + seq.copy_(torch.tensor(lengths, device="cuda")) + req.copy_(torch.tensor(requests, device="cuda")) + precomputed = source._precompute_replay_metadata( + BS, req, seq, seq.cpu(), mode + ) + source.init_forward_metadata_replay_cuda_graph_from_precomputed( + BS, precomputed, mode + ) + # An eligible sibling must reuse the derived results, not silently + # fall through to the full recomputation path. + with patch.object( + sibling, + "init_forward_metadata_replay_cuda_graph_from_precomputed", + side_effect=AssertionError("unexpected sibling fallback"), + ): + sibling._copy_replay_metadata_from_sibling( + source, BS, precomputed, mode + ) + apply_metadata(ordinary, mode, seq, req) + assert_metadata_equal( + self, source.forward_metadata, ordinary.forward_metadata + ) + assert_metadata_equal( + self, sibling.forward_metadata, ordinary.forward_metadata + ) + self.assertEqual(pointers, addresses(sibling.forward_metadata)) + self.assertIsNot( + sibling.forward_metadata.kpool_write_plan, + source.forward_metadata.kpool_write_plan, + ) + + +if __name__ == "__main__": + unittest.main()