[GLM-5.3 Flash] Restore and enable KPool metadata fusion (#38845)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2@liverpool.ac.uk> Co-authored-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Co-authored-by: zanes-ops <zanes@nvidia.com>
This commit is contained in:
co-authored by
Xinyuan Tong
zRzRzRzRzRzRzR
Shijin Zhang
zanes-ops
parent
288627e400
commit
a66451c058
@@ -0,0 +1 @@
|
|||||||
|
"""Opt-in KPool metadata kernels; ordinary DSA kernels remain unchanged."""
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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))
|
||||||
@@ -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)
|
||||||
@@ -1544,6 +1544,8 @@ class Envs:
|
|||||||
SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias(
|
SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias(
|
||||||
True, deprecated_name="SGLANG_NSA_FUSE_TOPK"
|
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_DETERMINISTIC = EnvBool(False)
|
||||||
SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None)
|
SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None)
|
||||||
SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvIntWithAlias(
|
SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvIntWithAlias(
|
||||||
|
|||||||
@@ -122,11 +122,9 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
"""Precompute metadata for normal decode mode."""
|
"""Precompute metadata for normal decode mode."""
|
||||||
max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1]
|
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:
|
if (
|
||||||
from sglang.kernels.ops.attention.dsa_metadata import (
|
(_is_cuda or _is_hip) and self.dsa_index_kpool <= 1
|
||||||
fused_dsa_decode_metadata,
|
) or self.experimental_kpool_metadata_fusion:
|
||||||
)
|
|
||||||
|
|
||||||
cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device)
|
cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device)
|
||||||
cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device)
|
cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device)
|
||||||
page_indices = torch.empty(
|
page_indices = torch.empty(
|
||||||
@@ -146,7 +144,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
real_page_table = None
|
real_page_table = None
|
||||||
real_page_table_arg = page_indices
|
real_page_table_arg = page_indices
|
||||||
|
|
||||||
fused_dsa_decode_metadata(
|
self._fused_decode_metadata(
|
||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
req_to_token=self.req_to_token,
|
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]
|
max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1]
|
||||||
seqlens_expanded_size = bs * self.speculative_num_draft_tokens
|
seqlens_expanded_size = bs * self.speculative_num_draft_tokens
|
||||||
|
|
||||||
if (_is_cuda or _is_hip) and self.dsa_index_kpool <= 1:
|
if (
|
||||||
from sglang.kernels.ops.attention.dsa_metadata import (
|
(_is_cuda or _is_hip) and self.dsa_index_kpool <= 1
|
||||||
fused_dsa_target_verify_metadata,
|
) or self.experimental_kpool_metadata_fusion:
|
||||||
)
|
|
||||||
|
|
||||||
cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device)
|
cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device)
|
||||||
cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device)
|
cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device)
|
||||||
page_indices = torch.empty(
|
page_indices = torch.empty(
|
||||||
@@ -282,7 +278,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
real_page_table = None
|
real_page_table = None
|
||||||
real_page_table_arg = page_indices
|
real_page_table_arg = page_indices
|
||||||
|
|
||||||
fused_dsa_target_verify_metadata(
|
self._fused_verify_metadata(
|
||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
req_to_token=self.req_to_token,
|
req_to_token=self.req_to_token,
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -35,11 +35,6 @@ from sglang.kernels.ops.attention.dsa.transform_index import (
|
|||||||
transform_index_page_table_decode,
|
transform_index_page_table_decode,
|
||||||
transform_index_page_table_prefill,
|
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 (
|
from sglang.kernels.ops.attention.utils import (
|
||||||
concat_mla_absorb_q_general,
|
concat_mla_absorb_q_general,
|
||||||
mla_quantize_and_rope_for_fp8,
|
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.kernels.ops.kvcache.cache_ops import concat_and_cast_q_fp8_pad
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
get_dsa_index_kpool,
|
get_dsa_index_kpool,
|
||||||
get_dsa_index_topk,
|
|
||||||
is_deepseek_dsa,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
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,
|
compute_cu_seqlens,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import DSAIndexerMetadata
|
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 (
|
from sglang.srt.layers.attention.dsa.dsa_topk_backend import (
|
||||||
DSATopKBackend,
|
DSATopKBackend,
|
||||||
TopkTransformMethod,
|
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.base import get_cp_strategy
|
||||||
from sglang.srt.layers.cp.utils import is_cp_active
|
from sglang.srt.layers.cp.utils import is_cp_active
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
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 (
|
from sglang.srt.utils import (
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
@@ -301,6 +296,7 @@ _DSA_IMPL_T: TypeAlias = Literal[
|
|||||||
|
|
||||||
|
|
||||||
class DeepseekSparseAttnBackend(
|
class DeepseekSparseAttnBackend(
|
||||||
|
DSAMetadataManagementMixin,
|
||||||
DeepseekSparseAttnBackendKPoolMixin,
|
DeepseekSparseAttnBackendKPoolMixin,
|
||||||
DeepseekSparseAttnBackendMTPPrecomputeMixin,
|
DeepseekSparseAttnBackendMTPPrecomputeMixin,
|
||||||
AttentionBackend,
|
AttentionBackend,
|
||||||
@@ -339,6 +335,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
self.dsa_index_topk = get_dsa_index_topk(hf_config)
|
self.dsa_index_topk = get_dsa_index_topk(hf_config)
|
||||||
self.dsa_index_kpool = get_dsa_index_kpool(hf_config)
|
self.dsa_index_kpool = get_dsa_index_kpool(hf_config)
|
||||||
self.needs_cpu_seq_lens = self.dsa_index_kpool > 1
|
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.max_context_len = model_runner.model_config.context_len
|
||||||
self.num_q_heads = (
|
self.num_q_heads = (
|
||||||
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
||||||
@@ -1523,8 +1520,10 @@ class DeepseekSparseAttnBackend(
|
|||||||
# Normal Decode
|
# Normal Decode
|
||||||
max_len = self._graph_page_table_width(metadata)
|
max_len = self._graph_page_table_width(metadata)
|
||||||
|
|
||||||
if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1:
|
if (
|
||||||
fused_dsa_decode_metadata(
|
(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,
|
seq_lens=seq_lens,
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
req_to_token=self.req_to_token,
|
req_to_token=self.req_to_token,
|
||||||
@@ -1563,7 +1562,9 @@ class DeepseekSparseAttnBackend(
|
|||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
max_seqlen_k = self._graph_page_table_width(metadata)
|
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
|
paged_mqa_ctx_lens_2d = None
|
||||||
if (
|
if (
|
||||||
self.speculative_num_draft_tokens >= 2
|
self.speculative_num_draft_tokens >= 2
|
||||||
@@ -1576,7 +1577,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
):
|
):
|
||||||
paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d
|
paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d
|
||||||
|
|
||||||
fused_dsa_target_verify_metadata(
|
self._fused_verify_metadata(
|
||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
req_to_token=self.req_to_token,
|
req_to_token=self.req_to_token,
|
||||||
@@ -1659,8 +1660,10 @@ class DeepseekSparseAttnBackend(
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
if (is_cuda() or _is_hip) and self.dsa_index_kpool <= 1:
|
if (
|
||||||
fused_dsa_draft_extend_metadata(
|
(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,
|
seq_lens=seq_lens,
|
||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
@@ -1811,125 +1814,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
|
|
||||||
metadata = self.decode_cuda_graph_metadata[bs]
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
|
|
||||||
# Track whether fused kernel succeeded
|
self._copy_base_replay_buffers(bs, metadata, precomputed, forward_mode)
|
||||||
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)
|
|
||||||
|
|
||||||
# Refresh the schedule because stale shape decomposition can deadlock
|
# Refresh the schedule because stale shape decomposition can deadlock
|
||||||
# DeepGEMM paged MQA.
|
# DeepGEMM paged MQA.
|
||||||
@@ -3788,6 +3673,20 @@ class DeepseekSparseAttnMultiStepBackend:
|
|||||||
forward_mode=ForwardMode.DECODE,
|
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
|
# Use multi-backend fused copy when we have 3 or more backends
|
||||||
# This is 3x faster than calling the single-backend copy 3 times
|
# This is 3x faster than calling the single-backend copy 3 times
|
||||||
if self.speculative_num_steps > 3:
|
if self.speculative_num_steps > 3:
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user