Build the unified read stream directly, without the page-table rectangle (#37512)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-09-02 16:55:16 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 18d5ffb42a
commit d9848b9ecd
11 changed files with 576 additions and 425 deletions
+50 -57
View File
@@ -193,9 +193,7 @@ def _fused_metadata_kernel_general(
use_swa: tl.constexpr,
SHIFT: tl.constexpr,
BLOCK_COLS: tl.constexpr,
# 1: the two table pointers carry PAGE-granular, already kernel-facing
# read tables; emit verbatim -- no >>SHIFT, no v2p, no mapping gather.
SRC_IS_KERNEL_PAGE_TABLE: tl.constexpr = 0,
SKIP_PAGE_TABLE: tl.constexpr = 0,
):
pid_b = tl.program_id(0) # batch index
pid_c = tl.program_id(1) # column chunk index
@@ -212,6 +210,8 @@ def _fused_metadata_kernel_general(
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
# 2. Gather for this batch and column chunk
if SKIP_PAGE_TABLE:
return
if max_seq_pages == 0:
return
@@ -236,11 +236,8 @@ def _fused_metadata_kernel_general(
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
mask = col_offsets < num_live_pages
# Compute column indices in the source tensor (token offset; page offset
# when the source is already the page-granular canonical)
if SRC_IS_KERNEL_PAGE_TABLE:
col_idx = col_offsets
elif page_size == 1:
# Compute column indices in the source tensor
if page_size == 1:
col_idx = col_offsets
else:
col_idx = col_offsets << SHIFT # faster than multiplication for power-of-two
@@ -252,9 +249,7 @@ def _fused_metadata_kernel_general(
)
# Compute page_table
if SRC_IS_KERNEL_PAGE_TABLE:
page_table_val = page_index # read-table entries are the page ids
elif page_size == 1:
if page_size == 1:
page_table_val = page_index
else:
page_table_val = page_index >> SHIFT
@@ -264,26 +259,16 @@ def _fused_metadata_kernel_general(
tl.store(page_table + pt_offsets, page_table_val, mask=mask, cache_modifier=".cg")
if use_swa:
if SRC_IS_KERNEL_PAGE_TABLE:
# The swa canonical shares the full canonical's shape and strides,
# so the SAME rt_offsets address the matching swa entry.
swa_val = tl.load(
full_to_swa_mapping + rt_offsets,
mask=mask,
other=0,
cache_modifier=".cg",
)
swa_slot = tl.load(
full_to_swa_mapping + page_index * full_to_swa_mapping_stride_0,
mask=mask,
other=0,
cache_modifier=".cg",
)
if page_size == 1:
swa_val = swa_slot
else:
swa_slot = tl.load(
full_to_swa_mapping + page_index * full_to_swa_mapping_stride_0,
mask=mask,
other=0,
cache_modifier=".cg",
)
if page_size == 1:
swa_val = swa_slot
else:
swa_val = swa_slot >> SHIFT
swa_val = swa_slot >> SHIFT
swa_offsets = (
i * swa_page_table_stride_0 + col_offsets * swa_page_table_stride_1
)
@@ -313,6 +298,7 @@ def _fused_metadata_kernel_ps1_no_swa(
max_seq_pages,
seq_len_delta: tl.constexpr,
BLOCK_COLS: tl.constexpr,
SKIP_PAGE_TABLE: tl.constexpr = 0,
):
pid_b = tl.program_id(0) # batch index
pid_c = tl.program_id(1) # column chunk index
@@ -329,6 +315,8 @@ def _fused_metadata_kernel_ps1_no_swa(
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
# 2. Gather for this batch and column chunk
if SKIP_PAGE_TABLE:
return
if max_seq_pages == 0:
return
@@ -581,8 +569,7 @@ def normal_decode_set_metadata(
page_size: int,
swa_page_table: Optional[torch.Tensor] = None,
token_to_kv_pool: Optional["SWAKVPool"] = None,
src_is_read_table: bool = False,
swa_src_table: Optional[torch.Tensor] = None,
skip_page_table: bool = False,
):
"""
Fused Triton implementation that replaces 4-5 sequential CUDA kernels with 1-2 kernels:
@@ -592,13 +579,8 @@ def normal_decode_set_metadata(
4. page_table = page_indices // page_size (floor-divide)
5. (optional) swa_page_table for sliding window attention
Unified pool (``src_is_read_table=True``): ``req_to_token`` /
``req_pool_indices`` carry the translator's PAGE-granular read table and its
row indices instead (entries already kernel-facing; ``swa_src_table`` is
the swa canonical, same shape and strides); steps 3-5 become verbatim
copies of the read table's rows' live prefixes, folded into the same launch
so the capture-stable page_table is written translated with no separate
pass a caller could forget.
Unified pool (``skip_page_table=True``): the translator has already filled
the page tables in place, so only steps 1-2 run.
Achieves ~5.2x speedup on H200 hardware for typical decode workloads.
@@ -628,9 +610,34 @@ def normal_decode_set_metadata(
page_table_stride_0 = page_table.stride(0)
page_table_stride_1 = page_table.stride(1)
use_swa = swa_page_table is not None and (
token_to_kv_pool is not None or swa_src_table is not None
)
if skip_page_table:
# One block does the prefix sum, so one block is the whole grid.
_fused_metadata_kernel_ps1_no_swa[(1, 1)](
seq_lens,
seq_lens_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
req_pool_indices,
req_pool_indices_stride_0,
cache_seqlens_int32,
cache_seqlens_int32_stride_0,
cu_seqlens_k,
cu_seqlens_k_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
batch_size,
0,
seq_len_delta,
BLOCK_COLS=256,
SKIP_PAGE_TABLE=1,
num_warps=8,
num_stages=3,
)
return
use_swa = swa_page_table is not None and token_to_kv_pool is not None
# Unified SWA uses an independent SWA v2p table.
swa_v2p_page_table = None
@@ -678,20 +685,7 @@ def normal_decode_set_metadata(
else:
# General kernel for page_size > 1 or SWA cases
# SWA parameters
if use_swa and src_is_read_table:
# Unified pool: the swa canonical rides in the mapping slot; the
# kernel addresses it with the SAME row/col offsets as the full
# canonical, so their layouts must match exactly.
assert swa_src_table is not None
assert (
swa_src_table.stride() == req_to_token.stride()
), "swa canonical must share the full canonical's strides"
swa_page_table = swa_page_table.contiguous()
swa_page_table_stride_0 = swa_page_table.stride(0)
swa_page_table_stride_1 = swa_page_table.stride(1)
full_to_swa_mapping = swa_src_table
full_to_swa_mapping_stride_0 = 0 # unused under the canonical source
elif use_swa:
if use_swa:
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
swa_page_table = swa_page_table.contiguous()
@@ -751,7 +745,6 @@ def normal_decode_set_metadata(
use_swa,
shift,
BLOCK_COLS=BLOCK_COLS,
SRC_IS_KERNEL_PAGE_TABLE=1 if src_is_read_table else 0,
num_warps=4,
num_stages=3,
)
@@ -63,6 +63,7 @@ _TRITON_KERNELS = [
("cache_ops", "launch_reshape_and_cache_flash"),
("pd_dcp_gather", "copy_mla_rows_into_pack"),
("kv_read_table", "build_kv_read_table"),
("kv_read_table", "build_kv_read_table_packed"),
("kv_indices", "create_flashinfer_kv_indices_triton"),
("kv_indices", "create_flashmla_kv_indices_triton"),
("kv_indices", "create_chunked_prefix_cache_kv_indices"),
@@ -11,21 +11,26 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Builds the per-batch read table for the unified memory pool.
"""Builds the unified memory pool's read indices.
One fused gather-and-translate. For each request row it reads the virtual ids
out of `req_to_token`, converts each to the id the kernels can use, and writes
the result into `out`:
One gather-and-translate: for each request row, read the virtual ids out of
`req_to_token` and convert each to the id the kernels can use.
out[b, c] = clamp(v2p[req_to_token[req[b], c * ps] // ps] * multiplier, 0)
for c < ceil(seq_lens[b] / ps) -- the row's LIVE prefix
page(b, c) = req_to_token[req[b], c * ps] // ps -- the VIRTUAL page
entry(b, c) = clamp(v2p[page(b, c)] * multiplier, 0) -- kernel-facing
`v2p` is the pool's virtual->physical page table and `multiplier` scales a
physical page into the id space the per-layer views use (1 when one page maps
to one row-block). Since only the page number is rewritten, a token-level consumer can
rebuild flat ids as `entry * ps + offset`.
Two delivery forms over that one formula:
PREFIX-ONLY per row: columns past the live prefix are never written, so a
PAGE TABLE `out[b, c] = entry(b, c)`, rows a uniform stride apart, for a
consumer whose kernel reads a page table directly.
TOKEN STREAM `out[row_starts[b] + p] = entry(b, p // ps) * ps + p % ps`, the
indptr-addressed form a paged wrapper plans over. Converting an
id keeps its offset inside the page, so the token id is exact.
Its length is `sum(seq_lens)` -- one id per resident token,
which the pool bounds, where a page table's width is bounded
only by `max_context_len`.
PREFIX-ONLY per row: nothing past the row's live prefix is written, so a
caller-owned buffer keeps what it had there -- which is what lets a captured
cuda-graph buffer be refreshed in place. Readers bound themselves by
`cache_seqlens` and never look past the prefix.
@@ -33,58 +38,140 @@ cuda-graph buffer be refreshed in place. Readers bound themselves by
A `-1` in `req_to_token` and a freed (`-1`) v2p row both clamp to entry 0, the
reserved padding slot, so a kernel dereferences padding, not a wild address.
The grid is sized from `bs` alone and each program strides over the columns it
owns, bounded by the device-side `seq_lens`. A cuda-graph capture bakes the
grid, so a grid spanning `max_pages` would replay `max_context_len`/BLOCK column
blocks every step no matter how short the sequences actually are.
The grid is sized from `bs` alone and each program strides over the items it
owns, bounded by the device-side lengths. A cuda-graph capture bakes the grid,
so a grid spanning the full width would replay `max_context_len`/BLOCK blocks
every step no matter how short the sequences actually are.
"""
from __future__ import annotations
from typing import Optional
import torch
import triton
import triton.language as tl
_BLOCK_COLS = 512
_BLOCK_ITEMS = 512
_NUM_WARPS = 8
# Enough blocks to fill the device without oversubscribing the column loop;
# Enough blocks to fill the device without oversubscribing the item loop;
# measured on H100 over bs 1..256 x seq 1k..128k, flat within ~10% either side.
_TARGET_BLOCKS = 1024
@triton.jit
def build_kv_read_table_kernel(
def build_kv_read_indices_kernel(
req_to_token_ptr, # in: [max_reqs, max_context] -- VIRTUAL token ids
req_pool_indices_ptr, # in: [bs] -- row per batch lane
seq_lens_ptr, # in: [bs]
req_pool_indices_ptr, # in: [bs] -- req_to_token row per batch lane
seq_lens_ptr, # in: [bs] -- live TOKENS per row
v2p_ptr, # in: [num_pages + 1] int64 -- virtual->physical page table
out_ptr, # out: [>=bs, >=max_pages] int32 -- the read table
row_starts_ptr, # in: [bs + 1] or null -- CSR row starts; null = uniform
kv_start_idx_ptr, # in: [bs] or null -- first token of the row's window
out_ptr, # out: int32
req_stride, # runtime: req_to_token row stride (elements)
out_stride, # runtime: out row stride (elements)
out_stride, # runtime: uniform row stride, used when row_starts is null
mult, # runtime: kernel_page_multiplier of the target sub-pool
col_stride, # runtime: columns one program advances per loop trip
item_stride, # runtime: items one program advances per loop trip
PAGE_SIZE: tl.constexpr,
EMIT_PER_TOKEN: tl.constexpr,
OUT_INT64: tl.constexpr,
BLOCK: tl.constexpr,
):
bid = tl.program_id(0)
req = tl.load(req_pool_indices_ptr + bid).to(tl.int64)
seqlen = tl.load(seq_lens_ptr + bid)
n_pages = (seqlen + PAGE_SIZE - 1) // PAGE_SIZE
# Derived here, not on the host: one elementwise op there costs a whole
# launch, which a captured graph then replays every step.
if EMIT_PER_TOKEN:
n_items = seqlen
else:
n_items = (seqlen + PAGE_SIZE - 1) // PAGE_SIZE
kv_start = 0
if kv_start_idx_ptr:
kv_start = tl.load(kv_start_idx_ptr + bid).to(tl.int32)
row_in = req_to_token_ptr + req * req_stride
row_out = out_ptr + bid.to(tl.int64) * out_stride
if row_starts_ptr:
row_out = out_ptr + tl.load(row_starts_ptr + bid).to(tl.int64)
else:
row_out = out_ptr + bid.to(tl.int64) * out_stride
for start in range(tl.program_id(1) * BLOCK, n_pages, col_stride):
cols = start + tl.arange(0, BLOCK)
mask = cols < n_pages
tok = tl.load(row_in + cols.to(tl.int64) * PAGE_SIZE, mask=mask, other=0).to(
for start in range(tl.program_id(1) * BLOCK, n_items, item_stride):
item = start + tl.arange(0, BLOCK)
mask = item < n_items
pos = kv_start + item
if EMIT_PER_TOKEN:
page = pos // PAGE_SIZE
else:
page = pos
tok = tl.load(row_in + page.to(tl.int64) * PAGE_SIZE, mask=mask, other=0).to(
tl.int64
)
# Triton's `//` truncates toward zero, so `-1 // ps` is 0 for ps > 1 but
# -1 at ps == 1, which would read one element BEFORE `v2p`.
page = tl.where(tok < 0, 0, tok // PAGE_SIZE)
phys = tl.load(v2p_ptr + page, mask=mask, other=0)
entry = tl.maximum(phys * mult, 0).to(tl.int32)
tl.store(row_out + cols, entry, mask=mask)
vpage = tl.where(tok < 0, 0, tok // PAGE_SIZE)
entry = tl.maximum(tl.load(v2p_ptr + vpage, mask=mask, other=0) * mult, 0)
if EMIT_PER_TOKEN:
value = entry * PAGE_SIZE + pos % PAGE_SIZE
else:
value = entry
if OUT_INT64:
tl.store(row_out + item, value, mask=mask)
else:
tl.store(row_out + item, value.to(tl.int32), mask=mask)
def _launch(
*,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
v2p: torch.Tensor,
multiplier: int,
page_size: int,
max_items: int,
out: torch.Tensor,
out_stride: int,
row_starts: Optional[torch.Tensor],
kv_start_idx: Optional[torch.Tensor],
emit_per_token: bool,
) -> None:
bs = int(req_pool_indices.numel())
item_programs = min(
triton.cdiv(_TARGET_BLOCKS, bs), triton.cdiv(max_items, _BLOCK_ITEMS)
)
build_kv_read_indices_kernel[(bs, item_programs)](
req_to_token,
req_pool_indices,
seq_lens,
v2p,
row_starts,
kv_start_idx,
out,
req_to_token.stride(0),
out_stride,
multiplier,
item_programs * _BLOCK_ITEMS,
PAGE_SIZE=page_size,
EMIT_PER_TOKEN=emit_per_token,
OUT_INT64=out.dtype == torch.int64,
BLOCK=_BLOCK_ITEMS,
num_warps=_NUM_WARPS,
)
def _entries(
*,
req_to_token: torch.Tensor,
req: int,
page_cols: torch.Tensor,
v2p: torch.Tensor,
multiplier: int,
page_size: int,
) -> torch.Tensor:
"""The formula above, in torch. The allocator's unit tests run on CPU, so
without this the Triton kernel would have no coverage there."""
tok = req_to_token[req, page_cols * page_size].to(torch.int64)
return (v2p[torch.where(tok < 0, 0, tok // page_size)] * multiplier).clamp(min=0)
def build_kv_read_table(
@@ -98,7 +185,7 @@ def build_kv_read_table(
max_pages: int,
out: torch.Tensor,
) -> torch.Tensor:
"""Fill ``out``'s live prefix with read-table entries.
"""Fill ``out``'s live prefix with PAGE TABLE entries.
``out`` is caller-owned (fresh zeros for the eager path, the module's
capture-stable buffer for replay) and only its ``[:bs, :max_pages]``
@@ -122,33 +209,102 @@ def build_kv_read_table(
if not req_to_token.is_cuda:
cols = torch.arange(max_pages, device=req_to_token.device)
live = cols[None, :] < (
(seq_lens[:bs, None].to(torch.int64) + page_size - 1) // page_size
)
tok = req_to_token[
req_pool_indices[:bs, None].to(torch.int64), (cols * page_size)[None, :]
].to(torch.int64)
pages = torch.where(tok < 0, 0, tok // page_size)
entry = (v2p[pages] * multiplier).clamp(min=0).to(torch.int32)
dst = out[:bs, :max_pages]
dst.copy_(torch.where(live, entry, dst))
for b in range(bs):
n_pages = (int(seq_lens[b]) + page_size - 1) // page_size
live = min(n_pages, max_pages)
out[b, :live] = _entries(
req_to_token=req_to_token,
req=int(req_pool_indices[b]),
page_cols=cols[:live],
v2p=v2p,
multiplier=multiplier,
page_size=page_size,
).to(torch.int32)
return out
col_programs = min(
triton.cdiv(_TARGET_BLOCKS, bs), triton.cdiv(max_pages, _BLOCK_COLS)
_launch(
req_to_token=req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
v2p=v2p,
multiplier=multiplier,
page_size=page_size,
max_items=max_pages,
out=out,
out_stride=out.stride(0),
row_starts=None,
kv_start_idx=None,
emit_per_token=False,
)
build_kv_read_table_kernel[(bs, col_programs)](
req_to_token,
req_pool_indices,
seq_lens,
v2p,
out,
req_to_token.stride(0),
out.stride(0),
multiplier,
col_programs * _BLOCK_COLS,
PAGE_SIZE=page_size,
BLOCK=_BLOCK_COLS,
num_warps=_NUM_WARPS,
return out
def build_kv_read_table_packed(
*,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
v2p: torch.Tensor,
indptr: torch.Tensor,
multiplier: int,
page_size: int,
max_tokens: int,
out: torch.Tensor,
kv_start_idx: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Fill ``out``'s CSR rows with TOKEN STREAM ids.
``seq_lens`` counts tokens per row and ``indptr`` gives each row's start, so
the live stream is ``sum(seq_lens)`` long; ``max_tokens`` is the capacity
``out`` must have for that, and callers holding a capture-stable buffer pass
its size. ``kv_start_idx`` shifts a row's window start without moving where
it lands.
"""
bs = int(req_pool_indices.numel())
assert out.dtype in (torch.int32, torch.int64), (
f"build_kv_read_table_packed: out must be int32 or int64, got " f"{out.dtype}"
)
assert out.dim() == 1 and out.numel() >= max_tokens, (
f"build_kv_read_table_packed: out {tuple(out.shape)} cannot hold "
f"max_tokens={max_tokens}"
)
assert indptr.numel() > bs, (
f"build_kv_read_table_packed: indptr holds {indptr.numel()} entries, "
f"need {bs + 1}"
)
if bs == 0 or max_tokens == 0:
return out
if not req_to_token.is_cuda:
for b in range(bs):
n = int(seq_lens[b])
pos = torch.arange(n, device=req_to_token.device) + (
0 if kv_start_idx is None else int(kv_start_idx[b])
)
entry = _entries(
req_to_token=req_to_token,
req=int(req_pool_indices[b]),
page_cols=pos // page_size,
v2p=v2p,
multiplier=multiplier,
page_size=page_size,
)
start = int(indptr[b])
out[start : start + n] = (entry * page_size + pos % page_size).to(out.dtype)
return out
_launch(
req_to_token=req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
v2p=v2p,
multiplier=multiplier,
page_size=page_size,
max_items=max_tokens,
out=out,
out_stride=0,
row_starts=indptr,
kv_start_idx=kv_start_idx,
emit_per_token=True,
)
return out
@@ -192,7 +192,6 @@ class FlashAttentionBackend(AttentionBackend):
self.needs_cpu_seq_lens = False
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.kv_index_translator = model_runner.kv_index_translator
self.kv_read_tables = None
self.skip_prefill = skip_prefill
self.attn_cp_size = model_runner.ps.attn_cp_size
self._verify_mask = None
@@ -2166,12 +2165,6 @@ class FlashAttentionBackend(AttentionBackend):
"""
max_num_pages = (self.max_context_len + self.page_size - 1) // self.page_size
if self.kv_index_translator.is_translating:
# Zero-filled: slot 0 is the reserved sink in every id space.
self.kv_read_tables = self.kv_index_translator.make_capture_tables(
max_bs=max_bs, max_context_len=self.max_context_len
)
# This is being used by normal decode and draft decode when topk == 1
self.decode_cuda_graph_metadata = {
"cache_seqlens": torch.zeros(max_bs, dtype=torch.int32, device=self.device),
@@ -2713,6 +2706,44 @@ class FlashAttentionBackend(AttentionBackend):
src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu()
return src.max().item()
def _set_decode_page_metadata(
self,
metadata,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_len_delta: int,
) -> None:
"""Fill `cache_seqlens_int32`, `cu_seqlens_k` and the page table(s).
Under the unified pool the translator writes the page tables in place,
so the fused kernel is left with the prefix sum alone -- one pass over
the rows instead of a translated build plus a verbatim copy of it.
"""
translated = self.kv_index_translator.reads_are_translated
normal_decode_set_metadata(
metadata.cache_seqlens_int32,
metadata.cu_seqlens_k,
metadata.page_table,
self.req_to_token,
req_pool_indices,
self.max_num_pages,
seq_lens,
seq_len_delta,
self.page_size,
metadata.swa_page_table,
self.token_to_kv_pool if self.use_sliding_window_kv_pool else None,
skip_page_table=translated,
)
if translated:
# Fill to `cache_seqlens_int32`, which the kernels bound their reads
# by: a draft decode reads `seq_len_delta` past `seq_lens`.
self.kv_index_translator.fill_read_table(
out=metadata.page_table,
sliding_window_out=metadata.swa_page_table,
req_pool_indices=req_pool_indices,
seq_lens=metadata.cache_seqlens_int32,
)
def _apply_cuda_graph_metadata(
self,
bs: int,
@@ -2766,29 +2797,11 @@ class FlashAttentionBackend(AttentionBackend):
# is normal-decode-only).
# Spec is asserted off under the unified pool, so this
# captured view is always the passthrough (req_to_token).
kv_view = self.kv_index_translator.build_index_table(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
into=self.kv_read_tables,
)
normal_decode_set_metadata(
metadata.cache_seqlens_int32,
metadata.cu_seqlens_k,
metadata.page_table,
kv_view.ids,
kv_view.row_ids,
self.max_num_pages,
self._set_decode_page_metadata(
metadata,
req_pool_indices,
seq_lens,
self.speculative_step_id + 1,
self.page_size,
metadata.swa_page_table,
(
self.token_to_kv_pool
if self.use_sliding_window_kv_pool
else None
),
src_is_read_table=kv_view.is_translated,
swa_src_table=kv_view.sliding_window_ids,
)
else:
@@ -2888,29 +2901,8 @@ class FlashAttentionBackend(AttentionBackend):
if seq_lens_cpu is not None
else self.max_context_len
)
kv_view = self.kv_index_translator.build_index_table(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
into=self.kv_read_tables,
)
normal_decode_set_metadata(
metadata.cache_seqlens_int32,
metadata.cu_seqlens_k,
metadata.page_table,
kv_view.ids,
kv_view.row_ids,
self.max_num_pages,
seq_lens,
0,
self.page_size,
metadata.swa_page_table,
(
self.token_to_kv_pool
if self.use_sliding_window_kv_pool
else None
),
src_is_read_table=kv_view.is_translated,
swa_src_table=kv_view.sliding_window_ids,
self._set_decode_page_metadata(
metadata, req_pool_indices, seq_lens, 0
)
self._maybe_update_local_attn_metadata_for_replay(
@@ -26,7 +26,6 @@ import torch
from sglang.kernels.kernel_api_logging import debug_kernel_api
from sglang.kernels.ops.attention.utils import (
assert_buffer_fits,
create_flashinfer_kv_indices_triton,
)
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.environ import envs
@@ -36,7 +35,6 @@ from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
)
from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.kv_index_translator import KVIndexTable
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
@@ -312,7 +310,6 @@ class FlashInferAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.kv_index_translator = model_runner.kv_index_translator
self.kv_read_tables = None
self._swa_kv_pool: Optional[BaseSWAKVPool] = self._resolve_swa_kv_pool(
model_runner
)
@@ -724,14 +721,6 @@ class FlashInferAttnBackend(AttentionBackend):
num_tokens = forward_batch.positions.numel()
self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info)
# All flashinfer gathers run OUT-of-graph (plan time), so the
# capture-stable read table is buffer reuse, not pointer stability.
kv_view = self.kv_index_translator.build_index_table(
req_pool_indices=req_pool_indices[:bs],
seq_lens=seq_lens[:bs],
into=self.kv_read_tables,
)
if forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
seq_lens[:bs],
@@ -742,7 +731,7 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=spec_info,
fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split,
kv_view=kv_view,
req_pool_indices=req_pool_indices,
)
elif forward_mode.is_target_verify():
self.indices_updater_prefill.update(
@@ -755,7 +744,6 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
kv_view=kv_view,
)
elif forward_mode.is_dllm_extend():
self.indices_updater_prefill.update(
@@ -768,7 +756,6 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged=not self.use_paged,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=None,
kv_view=kv_view,
)
elif forward_mode.is_draft_extend_v2():
self.indices_updater_prefill.update(
@@ -781,7 +768,6 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
kv_view=kv_view,
)
elif forward_mode.is_extend():
# Plain EXTEND under full prefill CUDA graph. plan() runs
@@ -800,7 +786,6 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=None,
kv_view=kv_view,
)
else:
raise ValueError("Invalid forward mode")
@@ -954,7 +939,6 @@ class FlashInferAttnBackend(AttentionBackend):
return layer.k_scale, layer.v_scale
def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_view = self.kv_index_translator.index_table_for_batch(forward_batch)
swa_out_cache_loc = None
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
swa_out_cache_loc = self.kv_index_translator.sliding_window_write_loc_for(
@@ -971,7 +955,7 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=forward_batch.spec_info,
fixed_split_size=self.decode_split_tile_size,
disable_split_kv=False,
kv_view=kv_view,
req_pool_indices=forward_batch.req_pool_indices,
)
self.forward_metadata = DecodeMetadata(
self.decode_wrappers, swa_out_cache_loc=swa_out_cache_loc
@@ -987,7 +971,6 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged=False,
encoder_lens=forward_batch.encoder_lens,
spec_info=forward_batch.spec_info,
kv_view=kv_view,
)
self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_verify,
@@ -1041,7 +1024,6 @@ class FlashInferAttnBackend(AttentionBackend):
cross_attention_custom_mask=forward_batch.cross_attention_custom_mask,
extend_prefix_lens_cpu=forward_batch.extend_prefix_lens_cpu,
custom_kv_indices=self.dq_page_table,
kv_view=kv_view,
)
self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_paged,
@@ -1057,9 +1039,6 @@ class FlashInferAttnBackend(AttentionBackend):
max_num_tokens: int,
kv_indices_buf: Optional[torch.Tensor] = None,
):
self.kv_read_tables = self.kv_index_translator.make_capture_tables(
max_bs=max_bs, max_context_len=self.max_context_len
)
if kv_indices_buf is None:
cuda_graph_kv_indices = torch.zeros(
(max_num_tokens * self.max_context_len,),
@@ -1586,7 +1565,7 @@ class FlashInferIndicesUpdaterDecode:
fixed_split_size: Optional[int] = None,
disable_split_kv: Optional[bool] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
):
# Keep the signature for type checking. It will be assigned during runtime.
raise NotImplementedError()
@@ -1602,7 +1581,7 @@ class FlashInferIndicesUpdaterDecode:
fixed_split_size: Optional[int] = None,
disable_split_kv: Optional[bool] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
):
decode_wrappers = decode_wrappers or self.decode_wrappers
self.call_begin_forward(
@@ -1615,7 +1594,7 @@ class FlashInferIndicesUpdaterDecode:
seq_lens_cpu,
fixed_split_size=fixed_split_size,
disable_split_kv=disable_split_kv,
kv_view=kv_view,
req_pool_indices=req_pool_indices,
)
def update_sliding_window(
@@ -1629,7 +1608,7 @@ class FlashInferIndicesUpdaterDecode:
fixed_split_size: Optional[int] = None,
disable_split_kv: Optional[bool] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
):
assert self.sliding_window_size is not None
for wrapper_id in range(2):
@@ -1668,7 +1647,7 @@ class FlashInferIndicesUpdaterDecode:
use_sliding_window_kv_pool=use_sliding_window_kv_pool,
fixed_split_size=fixed_split_size,
disable_split_kv=disable_split_kv,
kv_view=kv_view,
req_pool_indices=req_pool_indices,
)
def update_cross_attention(
@@ -1682,7 +1661,7 @@ class FlashInferIndicesUpdaterDecode:
fixed_split_size: Optional[int] = None,
disable_split_kv: Optional[bool] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
):
# Cache encoder_lens on CPU to avoid GPU→CPU transfer per call
encoder_lens_cpu = encoder_lens.cpu() if encoder_lens is not None else None
@@ -1708,7 +1687,7 @@ class FlashInferIndicesUpdaterDecode:
seq_lens_cpu=kv_lens_cpu,
fixed_split_size=fixed_split_size,
disable_split_kv=disable_split_kv,
kv_view=kv_view,
req_pool_indices=req_pool_indices,
)
def call_begin_forward(
@@ -1724,12 +1703,13 @@ class FlashInferIndicesUpdaterDecode:
fixed_split_size: Optional[int] = None,
disable_split_kv: Optional[bool] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
):
# Unified SWA wrapper-0: gather from the swa canonical directly -- its
# entries are already swa-side kernel-facing ids, so the in-place
# full->swa translate below must not run on top of them.
use_swa_source = use_sliding_window_kv_pool and kv_view.is_translated
translator = self.attn_backend.kv_index_translator
use_swa_source = use_sliding_window_kv_pool and translator.reads_are_translated
if spec_info is None or getattr(spec_info, "kv_indptr", None) is None:
bs = len(paged_kernel_lens)
kv_indptr[1 : bs + 1] = torch.cumsum(paged_kernel_lens, dim=0)
@@ -1743,20 +1723,14 @@ class FlashInferIndicesUpdaterDecode:
paged_kernel_lens_sum, dtype=torch.int32, device="cuda"
)
if use_swa_source:
assert kv_view.sliding_window_ids is not None
src_table = kv_view.sliding_window_ids
else:
src_table = kv_view.ids
create_flashinfer_kv_indices_triton[(bs,)](
src_table,
kv_view.row_ids,
paged_kernel_lens,
kv_indptr,
kv_start_idx,
kv_indices,
kv_view.row_stride,
ENTRY_PAGE_SIZE=kv_view.entry_page_size,
translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices,
seq_lens=paged_kernel_lens,
indptr=kv_indptr,
total_tokens=paged_kernel_lens_sum,
out=kv_indices,
kv_start_idx=kv_start_idx,
sliding_window=use_swa_source,
)
else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
@@ -1882,8 +1856,6 @@ class FlashInferIndicesUpdaterPrefill:
cross_attention_custom_mask: Optional[torch.Tensor] = None,
extend_prefix_lens_cpu: Optional[List[int]] = None,
custom_kv_indices: Optional[torch.Tensor] = None,
*,
kv_view: KVIndexTable,
):
# Keep the signature for type checking. It will be assigned during runtime.
raise NotImplementedError()
@@ -1904,8 +1876,6 @@ class FlashInferIndicesUpdaterPrefill:
cross_attention_custom_mask: Optional[torch.Tensor] = None,
extend_prefix_lens_cpu: Optional[List[int]] = None,
custom_kv_indices: Optional[torch.Tensor] = None,
*,
kv_view: KVIndexTable,
):
if use_ragged:
assert prefix_lens is not None
@@ -1936,7 +1906,6 @@ class FlashInferIndicesUpdaterPrefill:
multi_item_params=multi_item_params,
seq_lens_cpu=seq_lens_cpu,
custom_kv_indices=custom_kv_indices,
kv_view=kv_view,
)
def update_sliding_window(
@@ -1955,8 +1924,6 @@ class FlashInferIndicesUpdaterPrefill:
cross_attention_custom_mask: Optional[torch.Tensor] = None,
extend_prefix_lens_cpu: Optional[List[int]] = None,
custom_kv_indices: Optional[torch.Tensor] = None,
*,
kv_view: KVIndexTable,
):
if custom_kv_indices is not None:
raise RuntimeError(
@@ -2044,7 +2011,6 @@ class FlashInferIndicesUpdaterPrefill:
if (wrapper_id == 0 and not use_ragged and spec_info is None)
else -1
),
kv_view=kv_view,
)
def _build_swa_prefix_custom_mask(
@@ -2104,8 +2070,6 @@ class FlashInferIndicesUpdaterPrefill:
cross_attention_custom_mask: Optional[torch.Tensor] = None,
extend_prefix_lens_cpu: Optional[List[int]] = None,
custom_kv_indices: Optional[torch.Tensor] = None,
*,
kv_view: KVIndexTable,
):
if custom_kv_indices is not None:
raise RuntimeError(
@@ -2141,7 +2105,6 @@ class FlashInferIndicesUpdaterPrefill:
cross_attention_custom_mask=(
cross_attention_custom_mask if wrapper_id == 1 else None
),
kv_view=kv_view,
)
def call_begin_forward(
@@ -2165,14 +2128,13 @@ class FlashInferIndicesUpdaterPrefill:
seq_lens_cpu: Optional[torch.Tensor] = None,
custom_kv_indices: Optional[torch.Tensor] = None,
window_left: int = -1,
*,
kv_view: KVIndexTable,
):
bs = len(seq_lens)
# Unified SWA wrapper-0: gather from the swa canonical directly -- its
# entries are already swa-side kernel-facing ids, so the in-place
# full->swa translate below must not run on top of them.
use_swa_source = use_sliding_window_kv_pool and kv_view.is_translated
translator = self.attn_backend.kv_index_translator
use_swa_source = use_sliding_window_kv_pool and translator.reads_are_translated
if spec_info is None:
assert prefix_lens is not None
assert len(seq_lens) == len(req_pool_indices)
@@ -2198,20 +2160,14 @@ class FlashInferIndicesUpdaterPrefill:
dtype=torch.int32,
device=req_pool_indices.device,
)
if use_swa_source:
assert kv_view.sliding_window_ids is not None
src_table = kv_view.sliding_window_ids
else:
src_table = kv_view.ids
create_flashinfer_kv_indices_triton[(bs,)](
src_table,
kv_view.row_ids,
paged_kernel_lens,
kv_indptr,
kv_start_idx,
kv_indices,
kv_view.row_stride,
ENTRY_PAGE_SIZE=kv_view.entry_page_size,
translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices,
seq_lens=paged_kernel_lens,
indptr=kv_indptr,
total_tokens=paged_kernel_lens_sum,
out=kv_indices,
kv_start_idx=kv_start_idx,
sliding_window=use_swa_source,
)
qo_indptr[1 : bs + 1] = torch.cumsum(seq_lens - prefix_lens, dim=0)
qo_indptr = qo_indptr[: bs + 1]
@@ -27,15 +27,11 @@ import torch
from sglang.kernels.ops.attention.utils import assert_buffer_fits
from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.flashinfer_backend import (
create_flashinfer_kv_indices_triton,
)
from sglang.srt.layers.dcp import (
DecodeContextParallelMetadata,
update_local_kv_lens_for_dcp,
)
from sglang.srt.layers.dcp.planner import plan_dcp_decode_metadata
from sglang.srt.mem_cache.kv_index_translator import KVIndexTable
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
is_in_breakable_cuda_graph,
@@ -239,7 +235,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.kv_index_translator = model_runner.kv_index_translator
self.kv_read_tables = None
self.enable_chunk_kv = (
not skip_prefill
and get_disagg().disaggregation_mode != "decode"
@@ -344,14 +339,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
forward_mode = forward_batch.forward_mode
spec_info = forward_batch.spec_info
# All flashinfer gathers run OUT-of-graph (plan time), so the
# capture-stable table is buffer reuse, not pointer stability.
kv_view = self.kv_index_translator.build_index_table(
req_pool_indices=req_pool_indices[:bs],
seq_lens=seq_lens[:bs],
into=self.kv_read_tables,
)
if in_capture:
num_tokens = forward_batch.positions.numel()
seq_lens_sum = seq_lens.sum().item()
@@ -373,7 +360,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
decode_wrapper=decode_wrapper,
init_metadata_replay=False,
spec_info=spec_info,
kv_view=kv_view,
req_pool_indices=req_pool_indices[:bs],
)
self.decode_cuda_graph_metadata[bs] = decode_wrapper
self.forward_metadata = DecodeMetadata(decode_wrapper)
@@ -406,7 +393,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
in_capture=True,
kv_view=kv_view,
)
if forward_mode.is_target_verify() and (
spec_info is None
@@ -423,18 +409,16 @@ class FlashInferMLAAttnBackend(AttentionBackend):
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
kv_view=kv_view,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_view = self.kv_index_translator.index_table_for_batch(forward_batch)
if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
decode_wrapper=self.decode_wrapper,
init_metadata_replay=False,
kv_view=kv_view,
req_pool_indices=forward_batch.req_pool_indices,
)
self.forward_metadata = DecodeMetadata(self.decode_wrapper)
elif forward_batch.forward_mode.is_target_verify():
@@ -446,7 +430,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
prefill_wrapper_paged=self.prefill_wrapper_verify,
use_ragged=False,
spec_info=forward_batch.spec_info,
kv_view=kv_view,
)
self.forward_metadata = PrefillMetadata(self.prefill_wrapper_verify, False)
else:
@@ -495,7 +478,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
qo_indptr_cpu=qo_indptr_cpu,
kv_indptr_cpu=kv_indptr_cpu,
kv_len_arr_cpu=kv_len_arr_cpu,
kv_view=kv_view,
)
self.forward_metadata = PrefillMetadata(
self.prefill_wrapper_paged, use_ragged
@@ -507,9 +489,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
max_num_tokens: int,
kv_indices_buf: Optional[torch.Tensor] = None,
):
self.kv_read_tables = self.kv_index_translator.make_capture_tables(
max_bs=max_bs, max_context_len=self.max_context_len
)
if kv_indices_buf is None:
cuda_graph_kv_indices = torch.zeros(
(max_bs * self.max_context_len,),
@@ -553,7 +532,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
kv_view: KVIndexTable,
in_capture: bool = False,
):
"""Shared capture+replay body for the cuda-graph init path.
@@ -581,7 +559,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
decode_wrapper=self.decode_cuda_graph_metadata[bs],
init_metadata_replay=True,
spec_info=spec_info,
kv_view=kv_view,
req_pool_indices=req_pool_indices[:bs],
**self.fast_decode_kwargs,
)
elif forward_mode.is_target_verify():
@@ -632,7 +610,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
if use_generic_fast_plan
else None
),
kv_view=kv_view,
)
else:
raise ValueError(f"Invalid forward mode: {forward_mode=}")
@@ -875,7 +852,7 @@ class FlashInferMLAIndicesUpdaterDecode:
init_metadata_replay: bool = False,
spec_info: Optional[SpecInput] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
**fast_decode_kwargs,
):
decode_wrapper = decode_wrapper or self.decode_wrapper
@@ -887,7 +864,7 @@ class FlashInferMLAIndicesUpdaterDecode:
self.kv_indptr,
init_metadata_replay,
spec_info,
kv_view=kv_view,
req_pool_indices=req_pool_indices,
**fast_decode_kwargs,
)
@@ -901,7 +878,7 @@ class FlashInferMLAIndicesUpdaterDecode:
init_metadata_replay: bool = False,
spec_info: Optional[SpecInput] = None,
*,
kv_view: KVIndexTable,
req_pool_indices: torch.Tensor,
**fast_decode_kwargs,
):
bs = len(paged_kernel_lens)
@@ -919,18 +896,16 @@ class FlashInferMLAIndicesUpdaterDecode:
if not init_metadata_replay
else fast_decode_kwargs["kv_indices"]
)
create_flashinfer_kv_indices_triton[(bs,)](
kv_view.ids,
kv_view.row_ids,
paged_kernel_lens,
kv_indptr,
None,
kv_indices,
kv_view.row_stride,
ENTRY_PAGE_SIZE=kv_view.entry_page_size,
translator = self.attn_backend.kv_index_translator
is_translated = translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices,
seq_lens=paged_kernel_lens,
indptr=kv_indptr,
total_tokens=paged_kernel_lens_sum,
out=kv_indices,
)
# The table above is deliberately VIRTUAL under DCP.
# The stream above is deliberately VIRTUAL under DCP.
n_kernel_ids = paged_kernel_lens_sum
if get_parallel().dcp_enabled:
n_kernel_ids = plan_dcp_decode_metadata(
@@ -945,9 +920,8 @@ class FlashInferMLAIndicesUpdaterDecode:
# capture-stable buffer the captured wrapper reads, so rebinding the
# local name would leave the graph on virtual ids. Only the prefix
# just filled is translated; the stale tail never indexes v2p.
translator = self.attn_backend.kv_index_translator
if (
not kv_view.is_translated
not is_translated
and n_kernel_ids > 0
and translator.needs_read_translate
):
@@ -1008,7 +982,7 @@ class FlashInferMLAIndicesUpdaterPrefill:
self.qo_indptr = attn_backend.qo_indptr
# Kept ONLY for the spec-info branch (generate_attn_arg_prefill), which
# is static-pool-only: unified memory asserts spec off. The normal
# builder sources from the per-batch KVIndexTable.
# builder reads req_to_token through the translator.
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.prefill_wrapper_ragged = attn_backend.prefill_wrapper_ragged
@@ -1024,7 +998,6 @@ class FlashInferMLAIndicesUpdaterPrefill:
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None,
fast_verify_plan_kwargs: Optional[dict] = None,
*,
kv_view: KVIndexTable,
qo_indptr_cpu: Optional[torch.Tensor] = None,
kv_indptr_cpu: Optional[torch.Tensor] = None,
kv_len_arr_cpu: Optional[torch.Tensor] = None,
@@ -1053,7 +1026,6 @@ class FlashInferMLAIndicesUpdaterPrefill:
qo_indptr_cpu=qo_indptr_cpu,
kv_indptr_cpu=kv_indptr_cpu,
kv_len_arr_cpu=kv_len_arr_cpu,
kv_view=kv_view,
)
def call_begin_forward(
@@ -1072,7 +1044,6 @@ class FlashInferMLAIndicesUpdaterPrefill:
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None,
fast_verify_plan_kwargs: Optional[dict] = None,
*,
kv_view: KVIndexTable,
qo_indptr_cpu: Optional[torch.Tensor] = None,
kv_indptr_cpu: Optional[torch.Tensor] = None,
kv_len_arr_cpu: Optional[torch.Tensor] = None,
@@ -1089,15 +1060,12 @@ class FlashInferMLAIndicesUpdaterPrefill:
dtype=torch.int32,
device=req_pool_indices.device,
)
create_flashinfer_kv_indices_triton[(bs,)](
kv_view.ids,
kv_view.row_ids,
paged_kernel_lens,
kv_indptr,
None,
kv_indices,
kv_view.row_stride,
ENTRY_PAGE_SIZE=kv_view.entry_page_size,
self.attn_backend.kv_index_translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices,
seq_lens=paged_kernel_lens,
indptr=kv_indptr,
total_tokens=paged_kernel_lens_sum,
out=kv_indices,
)
qo_indptr[1 : bs + 1] = torch.cumsum(seq_lens - prefix_lens, dim=0)
qo_indptr = qo_indptr[: bs + 1]
@@ -7,9 +7,6 @@ import torch
import triton
from sglang.kernels.ops.attention.metadata import get_num_kv_splits_triton
from sglang.kernels.ops.kvcache.kv_indices import (
create_flashinfer_kv_indices_triton,
)
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.configs.model_config import (
AttentionArch,
@@ -210,7 +207,6 @@ class TritonAttnBackend(AttentionBackend):
# byte-identical to the slot-based envelope.
self.page_size = getattr(model_runner, "page_size", 1) or 1
self.kv_index_translator = model_runner.kv_index_translator
self.kv_read_tables = None
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = get_spec().speculative_eagle_topk or 0
@@ -461,20 +457,17 @@ class TritonAttnBackend(AttentionBackend):
self,
bs: int,
seq_lens: torch.Tensor,
index_table,
req_pool_indices: torch.Tensor,
kv_indices: torch.Tensor,
) -> torch.Tensor:
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1:] = torch.cumsum(seq_lens, dim=0)
create_flashinfer_kv_indices_triton[(bs,)](
index_table.ids,
index_table.row_ids,
seq_lens,
kv_indptr,
None,
kv_indices,
index_table.row_stride,
ENTRY_PAGE_SIZE=index_table.entry_page_size,
self.kv_index_translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices[:bs],
seq_lens=seq_lens[:bs],
indptr=kv_indptr,
total_tokens=kv_indices.numel(),
out=kv_indices,
)
return kv_indptr
@@ -483,7 +476,6 @@ class TritonAttnBackend(AttentionBackend):
bs: int,
seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
index_table,
):
"""Fill KV (and SWA) cuda-graph buffers for decode/idle mode.
@@ -492,9 +484,6 @@ class TritonAttnBackend(AttentionBackend):
``num_kv_splits_lens`` is the per-request length used to size kv splits
(per-DCP-rank length clamped to >=1 when DCP is enabled, full seq_lens
otherwise).
``index_table`` is the captured read-index view: under the unified pool the
gathers below read the converted tables.
"""
seq_lens = seq_lens[:bs]
req_pool_indices = req_pool_indices[:bs]
@@ -512,7 +501,7 @@ class TritonAttnBackend(AttentionBackend):
num_kv_splits_lens = dcp_seq_lens.clamp_min(1)
else:
kv_indptr = self._fill_kv_indptr_and_indices(
bs, seq_lens, index_table, self.cuda_graph_kv_indices
bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices
)
num_kv_splits_lens = seq_lens
window_kv_indptr = self.window_kv_indptr
@@ -520,7 +509,8 @@ class TritonAttnBackend(AttentionBackend):
if self.sliding_window_size is not None and self.sliding_window_size > 0:
window_kv_indptr, _, window_kv_lens, _ = update_sliding_window_buffer(
self.window_kv_indptr,
index_table,
self.kv_index_translator,
req_pool_indices,
self.sliding_window_size,
seq_lens,
bs,
@@ -534,7 +524,7 @@ class TritonAttnBackend(AttentionBackend):
bs: int,
seq_lens: torch.Tensor,
spec_info,
index_table,
req_pool_indices: torch.Tensor,
):
"""Fill all cuda-graph buffers for target_verify mode."""
# Prefer the spec_info's per-request query length (DSpark draft propose
@@ -554,7 +544,7 @@ class TritonAttnBackend(AttentionBackend):
device=self.device,
)
kv_indptr = self._fill_kv_indptr_and_indices(
bs, seq_lens, index_table, self.cuda_graph_kv_indices
bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices
)
window_kv_indptr = self.window_kv_indptr
window_kv_indices = None
@@ -567,7 +557,8 @@ class TritonAttnBackend(AttentionBackend):
window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = (
update_sliding_window_buffer(
self.window_kv_indptr,
index_table,
self.kv_index_translator,
req_pool_indices,
self.sliding_window_size,
seq_lens[:bs],
bs,
@@ -605,7 +596,7 @@ class TritonAttnBackend(AttentionBackend):
seq_lens: torch.Tensor,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
index_table,
req_pool_indices: torch.Tensor,
):
"""Fill QO + KV cuda-graph buffers for draft_extend mode."""
seq_lens = seq_lens[:bs]
@@ -636,7 +627,7 @@ class TritonAttnBackend(AttentionBackend):
extend_seq_lens = torch.zeros(bs, dtype=torch.int32, device=seq_lens.device)
kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32)
kv_indptr = self._fill_kv_indptr_and_indices(
bs, kv_lens, index_table, self.cuda_graph_kv_indices
bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices
)
return qo_indptr, kv_indptr, num_tokens_per_req
@@ -763,9 +754,6 @@ class TritonAttnBackend(AttentionBackend):
if forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None or spec_info.kv_indptr is None:
index_table = self.kv_index_translator.index_table_for_batch(
forward_batch
)
# kv_indptr is None for draft-extend's idle batch; build from seq_lens.
if self.dcp_size > 1:
# DCP: per-rank sharded KV indices, else each rank reads the
@@ -786,7 +774,7 @@ class TritonAttnBackend(AttentionBackend):
kv_indptr = self._fill_kv_indptr_and_indices(
bs,
forward_batch.seq_lens,
index_table,
forward_batch.req_pool_indices,
kv_indices,
)
if (
@@ -796,7 +784,8 @@ class TritonAttnBackend(AttentionBackend):
window_kv_indptr, window_kv_indices, window_kv_lens, _ = (
update_sliding_window_buffer(
self.window_kv_indptr,
index_table,
self.kv_index_translator,
forward_batch.req_pool_indices,
self.sliding_window_size,
forward_batch.seq_lens,
bs,
@@ -888,11 +877,10 @@ class TritonAttnBackend(AttentionBackend):
kv_indices = torch.empty(
seq_lens_sum, dtype=torch.int64, device=self.device
)
index_table = self.kv_index_translator.index_table_for_batch(forward_batch)
kv_indptr = self._fill_kv_indptr_and_indices(
bs,
forward_batch.seq_lens,
index_table,
forward_batch.req_pool_indices,
kv_indices,
)
@@ -905,7 +893,8 @@ class TritonAttnBackend(AttentionBackend):
window_kv_offsets,
) = update_sliding_window_buffer(
self.window_kv_indptr,
index_table,
self.kv_index_translator,
forward_batch.req_pool_indices,
self.sliding_window_size,
forward_batch.seq_lens,
bs,
@@ -926,7 +915,6 @@ class TritonAttnBackend(AttentionBackend):
attn_lse = None
else:
index_table = self.kv_index_translator.index_table_for_batch(forward_batch)
if self.dcp_size > 1:
kv_indptr, kv_indices, _ = self._dcp_kv_indices(
forward_batch.req_pool_indices,
@@ -947,7 +935,7 @@ class TritonAttnBackend(AttentionBackend):
kv_indptr = self._fill_kv_indptr_and_indices(
bs,
forward_batch.extend_prefix_lens,
index_table,
forward_batch.req_pool_indices,
kv_indices,
)
if self.sliding_window_size is not None and self.sliding_window_size > 0:
@@ -958,7 +946,8 @@ class TritonAttnBackend(AttentionBackend):
window_kv_offsets,
) = update_sliding_window_buffer(
self.window_kv_indptr,
index_table,
self.kv_index_translator,
forward_batch.req_pool_indices,
self.sliding_window_size,
forward_batch.extend_prefix_lens,
bs,
@@ -1135,9 +1124,6 @@ class TritonAttnBackend(AttentionBackend):
dtype=torch.int64,
device=self.device,
)
self.kv_read_tables = self.kv_index_translator.make_capture_tables(
max_bs=max_bs, max_context_len=self.max_context_len
)
def _build_cuda_graph_forward_metadata(
self,
@@ -1257,15 +1243,10 @@ class TritonAttnBackend(AttentionBackend):
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
# NOTE: encoder_lens expected to be zeros or None
index_table = self.kv_index_translator.build_index_table(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
into=self.kv_read_tables,
)
if forward_mode.is_decode_or_idle():
assert spec_info is None, "Multi-step cuda graph init is not done here."
_, _, window_kv_lens, num_kv_splits_lens = self._update_decode_kv_buffers(
bs, seq_lens, req_pool_indices, index_table
bs, seq_lens, req_pool_indices
)
self.get_num_kv_splits(
self.cuda_graph_num_kv_splits[:bs], num_kv_splits_lens[:bs]
@@ -1276,10 +1257,12 @@ class TritonAttnBackend(AttentionBackend):
)
elif forward_mode.is_target_verify():
bs = len(req_pool_indices)
self._update_target_verify_buffers(bs, seq_lens, spec_info, index_table)
self._update_target_verify_buffers(
bs, seq_lens, spec_info, req_pool_indices
)
elif forward_mode.is_draft_extend_v2():
self._update_draft_extend_buffers(
bs, seq_lens, forward_mode, spec_info, index_table
bs, seq_lens, forward_mode, spec_info, req_pool_indices
)
else:
raise ValueError(
@@ -2198,7 +2181,8 @@ class TritonMultiStepDraftBackend:
def update_sliding_window_buffer(
window_kv_indptr,
index_table,
translator,
req_pool_indices,
sliding_window_size,
seq_lens,
bs,
@@ -2212,12 +2196,11 @@ def update_sliding_window_buffer(
path); omit it (or pass ``None``) to allocate a fresh tensor (eager path,
requires ``device``).
``index_table`` is the batch's read-index source view. Unified pool: the
gather reads the parallel SWA array (built directly from virtual ids
through the swa side's own v2p), so the window indices come out
already swa-side ids -- no translate here, eager or captured. Static SWA
pools gather full-token ids from req_to_token and keep the legacy
full->swa translate below.
Unified pool: the gather reads the swa sub-pool's own id space (built
directly from virtual ids through the swa side's own v2p), so the window
indices come out already swa-side ids -- no translate here, eager or
captured. Static SWA pools gather full-token ids from req_to_token and keep
the legacy full->swa translate below.
"""
window_kv_lens = torch.minimum(
seq_lens,
@@ -2230,18 +2213,16 @@ def update_sliding_window_buffer(
window_kv_indptr[-1], dtype=torch.int64, device=device
)
window_kv_start_idx = seq_lens - window_kv_lens
source_ids = index_table.sliding_window_read_ids()
create_flashinfer_kv_indices_triton[(bs,)](
source_ids,
index_table.row_ids,
window_kv_lens,
window_kv_indptr,
window_kv_start_idx,
window_kv_indices,
source_ids.stride(0),
ENTRY_PAGE_SIZE=index_table.entry_page_size,
translated = translator.fill_packed_read_stream(
req_pool_indices=req_pool_indices[:bs],
seq_lens=window_kv_lens,
indptr=window_kv_indptr,
total_tokens=window_kv_indices.numel(),
out=window_kv_indices,
kv_start_idx=window_kv_start_idx,
sliding_window=translator.reads_are_translated,
)
if not index_table.is_translated and isinstance(token_to_kv_pool, BaseSWAKVPool):
if not translated and isinstance(token_to_kv_pool, BaseSWAKVPool):
kv_last_index = window_kv_indptr[-1]
window_kv_indices[:kv_last_index] = (
token_to_kv_pool.translate_loc_from_full_to_swa(
@@ -469,9 +469,6 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
kv_indices_buf: Optional[torch.Tensor] = None,
):
"""Initialize CUDA graph state for TRTLLM MHA."""
self.kv_read_tables = self.kv_index_translator.make_capture_tables(
max_bs=max_bs, max_context_len=self.max_context_len
)
max_num_pages = self.max_num_pages
self.decode_cuda_graph_metadata = {
"cache_seqlens": torch.zeros(max_bs, dtype=torch.int32, device=self.device),
@@ -898,21 +895,21 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
)
if self.kv_index_translator.is_translating:
# Unified pool: refresh the capture-stable read table (this runs
# Unified pool: refill this mode's own page table (this runs
# out-of-graph on BOTH capture and every replay-prep; the recorded
# fused kernel skips its page-table writes so the graph reads the
# refreshed content through pointers baked at capture).
kv_view = self.kv_index_translator.build_index_table(
req_pool_indices=forward_batch.req_pool_indices[:bs],
seq_lens=forward_batch.seq_lens[:bs],
into=self.kv_read_tables,
)
metadata = self.forward_metadata
if in_capture:
# Bind ONCE: the attention kernels bake these pointers at capture.
metadata.page_table = kv_view.ids[:bs]
if kv_view.sliding_window_ids is not None:
metadata.swa_page_table = kv_view.sliding_window_ids[:bs]
# `cache_seqlens_int32` is what the attention kernels bound their
# page-table reads by, and the fused metadata call above wrote it.
# A target verify reads `draft_token_num` further than `seq_lens`
# goes, so filling to `seq_lens` leaves those columns untranslated.
self.kv_index_translator.fill_read_table(
out=metadata.page_table,
req_pool_indices=forward_batch.req_pool_indices[:bs],
seq_lens=metadata.cache_seqlens_int32,
sliding_window_out=metadata.swa_page_table,
)
# A capture batch carries no prepared write loc; zeros are the
# page-0 sink.
if (
@@ -57,7 +57,13 @@ from typing import Optional, Tuple
import msgspec
import torch
from sglang.kernels.ops.kvcache.kv_read_table import build_kv_read_table
from sglang.kernels.ops.kvcache.kv_indices import (
create_flashinfer_kv_indices_triton,
)
from sglang.kernels.ops.kvcache.kv_read_table import (
build_kv_read_table,
build_kv_read_table_packed,
)
from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator,
UnifiedSWATokenToKVPoolAllocator,
@@ -161,32 +167,85 @@ class KVIndexTranslator:
)
self._index_table_memo: Optional[Tuple[weakref.ref, KVIndexTable]] = None
def make_capture_tables(
self, *, max_bs: int, max_context_len: int
) -> Optional[KVReadTables]:
"""Capture-stable destinations for a backend to own, or None when this
pool needs no translation and the backend will never fill any.
Zero-filled: entry 0 is the reserved padding slot in every id space, so
a captured graph replaying before its first refresh reads padding, not
garbage.
"""
if not self.is_translating:
return None
max_pages = -(-max_context_len // self.page_size)
def _zeros():
return torch.zeros(
(max_bs, max_pages), dtype=torch.int32, device=self.device
)
return KVReadTables(
full=_zeros(),
sliding_window=_zeros() if self._swa_v2p_table is not None else None,
)
# -- per-batch view --------------------------------------------------------
@property
def reads_are_translated(self) -> bool:
"""Whether a read this translator fills comes out kernel-facing. False
on a non-unified pool, and under DCP, where the ids stay VIRTUAL for
``translate_dcp_read_ids`` to finish."""
return self.is_translating and not self.defer_read_translate
def fill_packed_read_stream(
self,
*,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
indptr: torch.Tensor,
total_tokens: int,
out: torch.Tensor,
kv_start_idx: Optional[torch.Tensor] = None,
sliding_window: bool = False,
) -> bool:
"""Fill ``out``'s CSR rows with the ids a paged wrapper plans over, and
report whether they came out translated.
Non-unified: the historical gather straight from ``req_to_token``.
Unified: one fused gather-and-translate, so no caller needs a
``[bs, max_pages]`` rectangle to repack from -- ``out`` holds one id per
resident token, a length the pool bounds.
``sliding_window`` selects the swa sub-pool's own id space, built from
VIRTUAL ids and never chained through full-physical. A ``False`` return
means the ids are still VIRTUAL: the DCP path defers translation to
``translate_dcp_read_ids``, and a static SWA pool maps the full ids
through its own full->swa table.
"""
# `seq_lens` sizes the batch: a caller may hold a wider req_pool_indices
# (the padded graph buffer), and the extra lanes have no length to bound.
bs = int(seq_lens.numel())
assert req_pool_indices.numel() >= bs, (
f"fill_packed_read_stream: {req_pool_indices.numel()} req rows for "
f"{bs} lengths"
)
req_pool_indices = req_pool_indices[:bs]
if not self.reads_are_translated:
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
indptr,
kv_start_idx,
out,
self.req_to_token.stride(0),
ENTRY_PAGE_SIZE=1,
)
return False
if sliding_window:
assert self._swa_v2p_table is not None, (
"fill_packed_read_stream: sliding_window on a pool with no swa "
"sub-pool"
)
build_kv_read_table_packed(
req_to_token=self.req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
v2p=self._swa_v2p_table if sliding_window else self._full_v2p_table,
indptr=indptr,
multiplier=(
self._swa_page_multiplier
if sliding_window
else self._full_page_multiplier
),
page_size=self.page_size,
max_tokens=total_tokens,
out=out,
kv_start_idx=kv_start_idx,
)
return True
def build_index_table(
self,
*,
@@ -273,18 +332,31 @@ class KVIndexTranslator:
out: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
sliding_window_out: Optional[torch.Tensor] = None,
) -> None:
"""`build_index_table(into=...)` for a caller that owns a bare block
table rather than a KVReadTables: trtllm_mla / flashmla consume that
table rather than a KVReadTables: the page-table consumers read that
table directly, its rows already being the index table's rows.
`sliding_window_out` fills the swa twin in the same pass, for a hybrid
model whose kernels take two block tables.
"""
assert (
self.is_translating
), "KVIndexTranslator.fill_read_table on a pool that needs no translation"
# `reads_are_translated`, not `is_translating`: under DCP the builder
# returns the passthrough view and writes nothing, so `is_translating`
# would let a caller keep a stale table and never hear about it.
assert self.reads_are_translated, (
"KVIndexTranslator.fill_read_table cannot fill a page table when "
"reads stay virtual (a non-unified pool, or DCP, where the caller "
"must select this rank's share itself)"
)
assert sliding_window_out is None or self._swa_v2p_table is not None, (
"KVIndexTranslator.fill_read_table: asked for a sliding-window "
"table on a pool with no swa sub-pool"
)
self.build_index_table(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
into=KVReadTables(full=out, sliding_window=None),
into=KVReadTables(full=out, sliding_window=sliding_window_out),
)
def index_table_for_batch(self, forward_batch) -> KVIndexTable: