[Perf] Unified memory: close the DCP decode gap on Blackwell (#37926)

This commit is contained in:
Cheng Wan
2026-09-07 01:10:44 -07:00
committed by GitHub
parent a8edafff7c
commit b5766336d4
21 changed files with 609 additions and 75 deletions
@@ -14,6 +14,8 @@ MTP inter-phase seam:
This module fuses that chain into a single launch. This module fuses that chain into a single launch.
""" """
from typing import Optional
import torch import torch
import triton import triton
import triton.language as tl import triton.language as tl
@@ -24,15 +26,20 @@ def _fused_replay_state_indices_kernel(
req_pool_indices_ptr, # (total_bs,) int64 — static replay buffer req_pool_indices_ptr, # (total_bs,) int64 — static replay buffer
mamba_map_ptr, # (req_pool_size,) int32 — req_index_to_mamba_index_mapping mamba_map_ptr, # (req_pool_size,) int32 — req_index_to_mamba_index_mapping
out_ptr, # (total_bs,) int32 — state_indices_list[bs - 1] out_ptr, # (total_bs,) int32 — state_indices_list[bs - 1]
v2p_ptr, # (num_slots + 1,) int64 — mamba virtual->physical, or unused
valid_bs, valid_bs,
total_bs, total_bs,
BS_UPPER: tl.constexpr, BS_UPPER: tl.constexpr,
HAS_V2P: tl.constexpr,
): ):
offs = tl.arange(0, BS_UPPER) offs = tl.arange(0, BS_UPPER)
in_range = offs < total_bs in_range = offs < total_bs
valid = offs < valid_bs valid = offs < valid_bs
req = tl.load(req_pool_indices_ptr + offs, mask=valid, other=0) req = tl.load(req_pool_indices_ptr + offs, mask=valid, other=0)
idx = tl.load(mamba_map_ptr + req, mask=valid, other=0) idx = tl.load(mamba_map_ptr + req, mask=valid, other=0)
if HAS_V2P:
# Must gather before the padding sentinel, as the reference chain does.
idx = tl.load(v2p_ptr + idx, mask=valid, other=0)
out_val = tl.where(valid, idx.to(tl.int32), -1) out_val = tl.where(valid, idx.to(tl.int32), -1)
tl.store(out_ptr + offs, out_val, mask=in_range) tl.store(out_ptr + offs, out_val, mask=in_range)
# Preserve the reference chain's side effect: padded rows of the static # Preserve the reference chain's side effect: padded rows of the static
@@ -49,6 +56,7 @@ def fused_replay_state_indices(
out_state_indices: torch.Tensor, out_state_indices: torch.Tensor,
valid_bs: int, valid_bs: int,
total_bs: int, total_bs: int,
v2p: Optional[torch.Tensor] = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""Fill the captured replay state-indices buffer in one launch. """Fill the captured replay state-indices buffer in one launch.
@@ -58,9 +66,11 @@ def fused_replay_state_indices(
the ``-1`` sentinel (mamba kernels skip ``state_idx < 0``) and their the ``-1`` sentinel (mamba kernels skip ``state_idx < 0``) and their
``req_pool_indices`` entries are zeroed. ``req_pool_indices`` entries are zeroed.
Callers must supply an identity virtual->physical mapping (the static ``v2p`` is the mamba virtual->physical slot table, for a pool whose slot
hybrid pool); the unified pool's allocator translate is not a flat table ids are virtual (the unified memory pool). Pass None when the mapping
gather and has to take the reference chain. already yields physical slots (the static hybrid pool). The unified
allocator runs the mamba sub-pool at page_size 1, so its translate is a
plain table gather and folds into this launch.
Returns the filled ``out_state_indices[:total_bs]`` view. Returns the filled ``out_state_indices[:total_bs]`` view.
""" """
@@ -68,8 +78,10 @@ def fused_replay_state_indices(
req_pool_indices, req_pool_indices,
mamba_index_mapping, mamba_index_mapping,
out_state_indices, out_state_indices,
v2p,
valid_bs, valid_bs,
total_bs, total_bs,
BS_UPPER=triton.next_power_of_2(total_bs), BS_UPPER=triton.next_power_of_2(total_bs),
HAS_V2P=v2p is not None,
) )
return out_state_indices[:total_bs] return out_state_indices[:total_bs]
@@ -18,6 +18,7 @@ _TRITON_KERNELS = [
("virtual_slot", "alloc_bind_inplace"), ("virtual_slot", "alloc_bind_inplace"),
("virtual_slot", "free_unbind_inplace"), ("virtual_slot", "free_unbind_inplace"),
("virtual_slot", "bind_inplace"), ("virtual_slot", "bind_inplace"),
("virtual_slot", "write_loc_to_kernel_ids"),
] ]
for _mod, _fn in _TRITON_KERNELS: for _mod, _fn in _TRITON_KERNELS:
register_kernel( register_kernel(
@@ -15,6 +15,8 @@
from __future__ import annotations from __future__ import annotations
from typing import Optional
import torch import torch
import triton import triton
import triton.language as tl import triton.language as tl
@@ -189,3 +191,134 @@ def bind_inplace(
return return
grid = (triton.cdiv(N, ALLOC_BIND_BLOCK),) grid = (triton.cdiv(N, ALLOC_BIND_BLOCK),)
bind_inplace_kernel[grid](v, p, v2p, p2v, N, BLOCK=ALLOC_BIND_BLOCK) bind_inplace_kernel[grid](v, p, v2p, p2v, N, BLOCK=ALLOC_BIND_BLOCK)
WRITE_LOC_BLOCK = 512
@triton.jit
def write_loc_to_kernel_id_kernel(
loc_ptr, # in: [N] int64 — WIDENED virtual token ids
v2p_ptr, # in: [num_pages + 1] int64 — virtual->physical page table
out_ptr, # out: [N] int64 — kernel-facing ids
N, # runtime: live element count
W, # runtime: lanes to write; [N, W) get 0
stride, # runtime: pool_page_size * kernel_page_multiplier
PAGE_SIZE: tl.constexpr,
DCP_SIZE: tl.constexpr,
DCP_RANK: tl.constexpr,
BLOCK: tl.constexpr,
):
"""``kernel_id(t) = v2p[t // ps] * ps * mult + t % ps``, clamped at 0.
Under DCP the incoming id is WIDENED: ``loc % dcp_size`` names its owner
and ``loc // dcp_size`` is the row. Ids this rank does not own resolve to
kernel id 0, the padding sink every write kernel skips.
Triton truncates ``//`` toward zero where torch floors it, so a negative
loc is tested explicitly rather than left to the division; it resolves to
0, as the torch path does.
Writing ``W > N`` lanes fills ``[N, W)`` with 0, the padding sink, so a
caller may hand in a capture-stable buffer wider than this batch and have
the stale tail cleared in the same launch.
"""
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
in_range = offs < W
mask = offs < N
loc = tl.load(loc_ptr + offs, mask=mask, other=0).to(tl.int64)
keep = mask & (loc >= 0)
if DCP_SIZE > 1:
keep = keep & ((loc % DCP_SIZE) == DCP_RANK)
loc = loc // DCP_SIZE
page = loc // PAGE_SIZE if PAGE_SIZE > 1 else loc
offset = loc % PAGE_SIZE if PAGE_SIZE > 1 else 0
# `keep` already excludes negatives, so the gather index is in range.
phys = tl.load(v2p_ptr + tl.where(keep, page, 0), mask=mask, other=0).to(tl.int64)
ids = tl.maximum(phys * stride + offset, 0)
tl.store(out_ptr + offs, tl.where(keep, ids, 0), mask=in_range)
def write_loc_to_kernel_ids(
*,
loc: torch.Tensor,
v2p: torch.Tensor,
page_size: int,
stride: int,
dcp_size: int = 1,
dcp_rank: int = 0,
out: Optional[torch.Tensor] = None,
out_width: Optional[int] = None,
) -> torch.Tensor:
"""One launch for the whole write-loc conversion; see the kernel.
``out`` is written in place when given (a captured graph records the
gather against a fixed ``data_ptr``), else a fresh int64 tensor is
returned. Cuda-graph safe: no ``.item()``, no host sync, no allocation on
the ``out=`` path.
``out_width`` writes that many lanes rather than ``loc.numel()``, zeroing
the ones past the batch; pass the captured tier's width to clear a stale
tail here.
"""
N = int(loc.numel())
# Flat-indexed as `ptr + offs`, so a strided view is mis-addressed.
assert loc.is_contiguous(), (
f"write_loc_to_kernel_ids: loc must be contiguous, got shape "
f"{tuple(loc.shape)} stride {tuple(loc.stride())}"
)
if out is None:
out = torch.empty_like(loc, dtype=torch.int64)
width = N if out_width is None else int(out_width)
assert out.dtype == torch.int64, (
f"write_loc_to_kernel_ids: out dtype must be int64 (matches v2p), "
f"got {out.dtype}"
)
if out_width is None:
# `out` mirrors `loc` whatever its shape; a 2-D page table is legal.
assert out.shape == loc.shape and out.is_contiguous(), (
f"write_loc_to_kernel_ids: out shape {tuple(out.shape)} must match "
f"loc shape {tuple(loc.shape)}"
)
else:
assert out.dim() == 1 and out.is_contiguous() and out.numel() >= width, (
f"write_loc_to_kernel_ids: out_width needs a packed 1-D out of at "
f"least {width}, got {tuple(out.shape)}"
)
assert width >= N, (
f"write_loc_to_kernel_ids: out_width {width} is under the batch's "
f"{N} locs, which would drop live rows"
)
if width == 0:
return out
if not loc.is_cuda:
# Pure-torch reference; the allocator's unit tests run on CPU.
big = loc.to(torch.int64)
keep = big >= 0
if dcp_size > 1:
keep = keep & (big % dcp_size == dcp_rank)
big = torch.div(big, dcp_size, rounding_mode="floor")
page = torch.where(keep, torch.div(big, page_size, rounding_mode="floor"), 0)
offset = big % page_size if page_size > 1 else 0
ids = (v2p[page] * stride + offset).clamp_(min=0)
out[:N].copy_(torch.where(keep, ids, torch.zeros_like(ids)))
if width > N:
out[N:width].zero_()
return out
grid = (triton.cdiv(width, WRITE_LOC_BLOCK),)
write_loc_to_kernel_id_kernel[grid](
loc,
v2p,
out,
N,
width,
stride,
PAGE_SIZE=page_size,
DCP_SIZE=dcp_size,
DCP_RANK=dcp_rank,
BLOCK=WRITE_LOC_BLOCK,
)
return out
@@ -249,6 +249,12 @@ def handle_unified_memory_pool(server_args: Any) -> None:
"not translate speculative verify indices to the unified " "not translate speculative verify indices to the unified "
"pool's kernel-facing space yet." "pool's kernel-facing space yet."
) )
assert not cfg.enable_two_batch_overlap, (
"--enable-unified-memory does not support --enable-two-batch-overlap: "
"TBO's replay split hands each child a view without the pre-translate "
"write loc, so a captured decode replay raises. "
"TODO(ch-wan): carry out_cache_loc_virtual into the child view."
)
assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), ( assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), (
"--enable-unified-memory is not yet compatible with hierarchical / " "--enable-unified-memory is not yet compatible with hierarchical / "
"host-tiered KV cache (--enable-hierarchical-cache / --enable-lmcache): " "host-tiered KV cache (--enable-hierarchical-cache / --enable-lmcache): "
@@ -693,6 +693,11 @@ class TboForwardBatchPreparer:
) )
output_dict[key] = old_value[start_token_index:end_token_index] output_dict[key] = old_value[start_token_index:end_token_index]
if batch.out_cache_loc_virtual is not None:
output_dict["out_cache_loc_virtual"] = batch.out_cache_loc_virtual[
start_token_index:end_token_index
]
attention_tp_size = get_parallel().attn_tp_size attention_tp_size = get_parallel().attn_tp_size
_tbo_padded_len = ( _tbo_padded_len = (
(end_token_index - start_token_index - 1) // attention_tp_size + 1 (end_token_index - start_token_index - 1) // attention_tp_size + 1
@@ -70,15 +70,11 @@ class MambaAttnBackendBase(AttentionBackend):
# backends on runners without a real model_config. # backends on runners without a real model_config.
self._model_runner = model_runner self._model_runner = model_runner
self._mamba_chunk_size: Optional[int] = None self._mamba_chunk_size: Optional[int] = None
# Fused replay-prep state-indices fast path (fused_replay_state_indices): pool = self.req_to_token_pool
# requires the static hybrid pool whose v2p translate is the identity —
# the unified pool overrides translate_mamba_indices with an allocator
# lookup that is not a flat table gather.
self._fused_state_indices_ok = ( self._fused_state_indices_ok = (
str(self.device).startswith("cuda") torch.device(self.device).type == "cuda"
and isinstance(self.req_to_token_pool, HybridReqToTokenPool) and isinstance(pool, HybridReqToTokenPool)
and type(self.req_to_token_pool).translate_mamba_indices and pool.mamba_translate_is_fusable
is HybridReqToTokenPool.translate_mamba_indices
) )
self.forward_metadata: ForwardMetadata = None self.forward_metadata: ForwardMetadata = None
self.state_indices_list = [] self.state_indices_list = []
@@ -643,6 +639,7 @@ class MambaAttnBackendBase(AttentionBackend):
out_state_indices=self.state_indices_list[bs - 1], out_state_indices=self.state_indices_list[bs - 1],
valid_bs=bs - int(num_padding), valid_bs=bs - int(num_padding),
total_bs=bs, total_bs=bs,
v2p=self.req_to_token_pool.mamba_v2p_table,
) )
else: else:
# Make sure forward metadata is correctly handled for padding reqs # Make sure forward metadata is correctly handled for padding reqs
@@ -748,22 +748,16 @@ class TritonAttnBackend(AttentionBackend):
def _fill_cuda_graph_write_locs( def _fill_cuda_graph_write_locs(
self, forward_batch: ForwardBatch, bs: int self, forward_batch: ForwardBatch, bs: int
) -> Optional[torch.Tensor]: ) -> Optional[torch.Tensor]:
"""Copy the cuda-graph WRITE loc into the capture-stable buffer and """Runs BEFORE graph.replay(), so it reads the live post-compaction
return the ``[:n]`` view; no-op for non-unified pools. v2p; no-op for non-unified pools."""
# The buffer exists only for a translating pool; return before naming it.
Runs BEFORE graph.replay() so it reads the live post-compaction v2p.
The capture batch is runner-built with zeros, which is safe because
slot 0 is the reserved sink in every id space.
"""
if not self.kv_index_translator.is_translating: if not self.kv_index_translator.is_translating:
return None return None
out_cache_loc = forward_batch.out_cache_loc return self.kv_index_translator.fill_capture_write_loc(
n = out_cache_loc.shape[0] out=self.cuda_graph_out_cache_loc_full_physical,
# Zero the padded tail first: a smaller replay batch leaves [n:] holding forward_batch=forward_batch,
# stale ids that the captured store would write; send them to slot 0 (sink). width=self.cuda_graph_out_cache_loc_full_physical.numel(),
self.cuda_graph_out_cache_loc_full_physical[n:].zero_() )
self.cuda_graph_out_cache_loc_full_physical[:n].copy_(out_cache_loc)
return self.cuda_graph_out_cache_loc_full_physical[:n]
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend.""" """Init auxiliary variables for triton attention backend."""
@@ -766,19 +766,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if self.kv_index_translator.is_translating and ( if self.kv_index_translator.is_translating and (
forward_mode.is_decode_or_idle() or forward_mode.is_target_verify() forward_mode.is_decode_or_idle() or forward_mode.is_target_verify()
): ):
out_cache_loc = forward_batch.out_cache_loc # The captured kernel consumes the whole buffer, so the tail a
n = out_cache_loc.shape[0] # shorter replay leaves must go to slot 0 rather than live pages.
dst = self.cuda_graph_out_cache_loc_kernel[:n] self._decode_kernel_loc = self.kv_index_translator.fill_capture_write_loc(
dst.copy_(out_cache_loc) out=self.cuda_graph_out_cache_loc_kernel,
# Replay-prep receives the RAW (unpadded) out_cache_loc forward_batch=forward_batch,
# (build_replay_fb_view), but the captured write kernel consumes the width=self.cuda_graph_out_cache_loc_kernel.numel(),
# full captured tier of this buffer. Zero the tail so pad rows write )
# to the sink (row 0) instead of stale kernel-facing locs left by
# earlier larger replays — a stale tail scatters pad-row garbage into
# live KV pages. Mirrors the runner's PaddingPolicy.ZERO on its own
# out_cache_loc slot.
self.cuda_graph_out_cache_loc_kernel[n:].zero_()
self._decode_kernel_loc = dst
else: else:
self._decode_kernel_loc = None self._decode_kernel_loc = None
@@ -331,10 +331,13 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
loc: torch.Tensor, loc: torch.Tensor,
*, *,
out: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None,
out_width: Optional[int] = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""Widened virtual WRITE loc -> kernel-facing id. DCP is rejected for this """Widened virtual WRITE loc -> kernel-facing id. DCP is rejected for this
composite at argument validation, so it coincides with the read translate.""" composite at argument validation, so it coincides with the read translate."""
return self.full_attn_allocator.translate_write_loc_for_kernel(loc, out=out) return self.full_attn_allocator.translate_write_loc_for_kernel(
loc, out=out, out_width=out_width
)
@property @property
def swa_kernel_page_multiplier(self) -> int: def swa_kernel_page_multiplier(self) -> int:
@@ -260,9 +260,12 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
loc: torch.Tensor, loc: torch.Tensor,
*, *,
out: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None,
out_width: Optional[int] = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""Widened virtual WRITE loc -> DENSE id; see the sub-allocator's copy.""" """Widened virtual WRITE loc -> DENSE id; see the sub-allocator's copy."""
return self.full_attn_allocator.translate_write_loc_for_kernel(loc, out=out) return self.full_attn_allocator.translate_write_loc_for_kernel(
loc, out=out, out_width=out_width
)
def translate_kv_indices_for_transfer( def translate_kv_indices_for_transfer(
self, kv_indices: torch.Tensor self, kv_indices: torch.Tensor
@@ -43,6 +43,7 @@ from sglang.kernels.ops.memory.virtual_slot import (
alloc_bind_inplace, alloc_bind_inplace,
bind_inplace, bind_inplace,
free_unbind_inplace, free_unbind_inplace,
write_loc_to_kernel_ids,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
@@ -990,39 +991,47 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
clamp to kernel-facing id 0, the page-0 sink. int64 out; a consumer whose clamp to kernel-facing id 0, the page-0 sink. int64 out; a consumer whose
kernel ABI wants int32 narrows where it fills that buffer. kernel ABI wants int32 narrows where it fills that buffer.
""" """
ps = self.pool_page_size
stride = ps * self.kernel_page_multiplier
with record_function("MultiEndedAlloc.translate_kv_loc_for_kernel"): with record_function("MultiEndedAlloc.translate_kv_loc_for_kernel"):
pages = virt_tokens if ps == 1 else virt_tokens // ps return self._translate_loc_fused(virt_tokens, dcp_size=1, out=out)
offsets = None if ps == 1 else virt_tokens % ps
if out is None: def _translate_loc_fused(
phys = self.virtual_to_physical[pages] self,
ids = phys * stride if offsets is None else phys * stride + offsets loc: torch.Tensor,
return ids.clamp_(min=0) *,
dcp_size: int,
dcp_rank: int = 0,
out: Optional[torch.Tensor] = None,
out_width: Optional[int] = None,
) -> torch.Tensor:
"""One launch for the read and write conversions alike; see
`write_loc_to_kernel_ids`."""
if out is not None:
assert out.dtype == torch.int64, ( assert out.dtype == torch.int64, (
f"translate_kv_loc_for_kernel: out= dtype must be int64 (matches v2p), " f"translate_kv_loc_for_kernel: out= dtype must be int64 (matches v2p), "
f"got {out.dtype}" f"got {out.dtype}"
) )
assert out.shape == virt_tokens.shape, ( if out_width is None:
assert out.shape == loc.shape, (
f"translate_kv_loc_for_kernel: out= shape {tuple(out.shape)} must " f"translate_kv_loc_for_kernel: out= shape {tuple(out.shape)} must "
f"match virt_tokens shape {tuple(virt_tokens.shape)}" f"match virt_tokens shape {tuple(loc.shape)}"
)
return write_loc_to_kernel_ids(
loc=loc,
v2p=self.virtual_to_physical,
page_size=self.pool_page_size,
stride=self.pool_page_size * self.kernel_page_multiplier,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
out=out,
out_width=out_width,
) )
if pages.dtype != torch.int64:
pages = pages.to(torch.int64)
if pages is virt_tokens:
out.copy_(torch.take(self.virtual_to_physical, pages))
else:
torch.take(self.virtual_to_physical, pages, out=out)
out.mul_(stride)
if offsets is not None:
out.add_(offsets)
return out.clamp_(min=0)
def translate_write_loc_for_kernel( def translate_write_loc_for_kernel(
self, self,
widened_loc: torch.Tensor, widened_loc: torch.Tensor,
*, *,
out: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None,
out_width: Optional[int] = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""Widened virtual WRITE loc (`out_cache_loc`) -> kernel-facing id. """Widened virtual WRITE loc (`out_cache_loc`) -> kernel-facing id.
@@ -1032,16 +1041,14 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
""" """
parallel = get_parallel() parallel = get_parallel()
dcp_size = parallel.attn_dcp_size if self.shards_under_dcp else 1 dcp_size = parallel.attn_dcp_size if self.shards_under_dcp else 1
if dcp_size == 1:
return self.translate_kv_loc_for_kernel(widened_loc, out=out)
with record_function("MultiEndedAlloc.translate_write_loc_for_kernel"): with record_function("MultiEndedAlloc.translate_write_loc_for_kernel"):
owned = (widened_loc % dcp_size) == parallel.attn_dcp_rank return self._translate_loc_fused(
dense = self.translate_kv_loc_for_kernel(widened_loc // dcp_size) widened_loc,
dense = torch.where(owned, dense, torch.zeros_like(dense)) dcp_size=dcp_size,
if out is not None: dcp_rank=parallel.attn_dcp_rank,
out.copy_(dense) out=out,
return out out_width=out_width,
return dense )
# -- alloc -- # -- alloc --
@@ -432,20 +432,55 @@ class KVIndexTranslator:
def rebind_write_loc(self, forward_batch) -> None: def rebind_write_loc(self, forward_batch) -> None:
"""Phase 1 of the WRITE contract: translate the batch's write loc to """Phase 1 of the WRITE contract: translate the batch's write loc to
FULL-side kernel-facing ids exactly once, at ForwardBatch FULL-side kernel-facing ids, once, at ForwardBatch construction.
construction. No-op on non-unified pools.
REBIND, never mutate: the translate returns a FRESH tensor, so the REBIND, never mutate: the translate returns a FRESH tensor, so the
ScheduleBatch's aliased tensor stays VIRTUAL for the radix / accept / ScheduleBatch's aliased tensor stays VIRTUAL for the radix / accept /
in-flight machinery that reads it. in-flight machinery that reads it. The pre-translate tensor stays on
the batch for `fill_capture_write_loc`.
""" """
self._index_table_memo = None self._index_table_memo = None
if not self.is_translating or forward_batch.out_cache_loc is None: if not self.is_translating or forward_batch.out_cache_loc is None:
return return
forward_batch.out_cache_loc_virtual = forward_batch.out_cache_loc
forward_batch.out_cache_loc = self._translate_write_full( forward_batch.out_cache_loc = self._translate_write_full(
forward_batch.out_cache_loc forward_batch.out_cache_loc
) )
def fill_capture_write_loc(
self,
*,
out: torch.Tensor,
forward_batch,
width: Optional[int] = None,
) -> Optional[torch.Tensor]:
"""Translate this batch's WRITE loc straight into ``out``, a backend's
capture-stable buffer, and return the live ``[:n]`` view. One launch
fills the live prefix and clears the tail a shorter replay leaves;
None when this pool needs no translation.
Must run at metadata-init time: `out` is reused every step, so filling
it sooner would race a still-pending previous step under overlap
scheduling.
"""
if not self.is_translating:
return None
virtual = forward_batch.out_cache_loc_virtual
if virtual is None:
loc = forward_batch.out_cache_loc
if loc is None:
return None
# The runner builds the capture batch outside `init_new`, so no
# rebind marked its virtual source; bake it holding sink ids.
width = int(loc.numel()) if width is None else int(width)
out[:width].zero_()
return out[: int(loc.numel())]
n = int(virtual.numel())
width = n if width is None else int(width)
buf = out[:width]
self._translate_write_full(virtual, out=buf, out_width=width)
return buf[:n]
def sliding_window_write_loc_for( def sliding_window_write_loc_for(
self, out_cache_loc: Optional[torch.Tensor] self, out_cache_loc: Optional[torch.Tensor]
) -> Optional[torch.Tensor]: ) -> Optional[torch.Tensor]:
@@ -1385,6 +1385,28 @@ class HybridReqToTokenPool(ReqToTokenPool):
def get_mamba_indices(self, req_indices: torch.Tensor) -> torch.Tensor: def get_mamba_indices(self, req_indices: torch.Tensor) -> torch.Tensor:
return self.req_index_to_mamba_index_mapping[req_indices] return self.req_index_to_mamba_index_mapping[req_indices]
@property
def mamba_v2p_table(self) -> Optional[torch.Tensor]:
"""The mamba virtual->physical slot table, or None when the ids this
pool hands out are already physical."""
return None
@property
def mamba_translate_is_fusable(self) -> bool:
"""Whether `fused_replay_state_indices` can reproduce this pool's
`translate_mamba_indices` in its own launch.
The kernel expresses exactly two shapes: the identity, and one gather
through `mamba_v2p_table`. A subclass that replaces the translate with
anything else is excluded here rather than silently mis-served.
"""
if self.mamba_v2p_table is not None:
return True
return (
type(self).translate_mamba_indices
is HybridReqToTokenPool.translate_mamba_indices
)
def translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor: def translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor:
"""Virtual->physical mamba-slot translate. Identity for a static pool """Virtual->physical mamba-slot translate. Identity for a static pool
(slots are physical); UnifiedHybridReqToTokenPool overrides it for the (slots are physical); UnifiedHybridReqToTokenPool overrides it for the
@@ -1107,6 +1107,15 @@ class UnifiedHybridReqToTokenPool(HybridReqToTokenPool):
) )
) )
@property
def mamba_v2p_table(self) -> Optional[torch.Tensor]:
"""This pool's ids ARE virtual; page_size is 1, so the translate is the
plain gather this table serves, which keeps `mamba_translate_is_fusable`
true despite the override."""
if self.mamba_allocator is None:
return None
return self.mamba_allocator.virtual_to_physical
def translate_mamba_indices(self, virtual_ids: torch.Tensor) -> torch.Tensor: def translate_mamba_indices(self, virtual_ids: torch.Tensor) -> torch.Tensor:
"""Virtual mamba ids -> physical slot ids.""" """Virtual mamba ids -> physical slot ids."""
return self.mamba_allocator.translate(virtual_ids).to(torch.int32) return self.mamba_allocator.translate(virtual_ids).to(torch.int32)
@@ -414,6 +414,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# The original sequence length without being chunked. Qwen-1M related. # The original sequence length without being chunked. Qwen-1M related.
orig_seq_lens: Optional[torch.Tensor] = None orig_seq_lens: Optional[torch.Tensor] = None
# The write loc before `rebind_write_loc` replaced it with kernel-facing
# ids; a backend re-derives from it into its capture-stable buffer.
out_cache_loc_virtual: Optional[torch.Tensor] = None
# DSV4-NPU only: per-pool slot bundle from DSV4NPUTokenToKVPoolAllocator, # DSV4-NPU only: per-pool slot bundle from DSV4NPUTokenToKVPoolAllocator,
# consumed by the Ascend backend for PA_ND block tables. None elsewhere. # consumed by the Ascend backend for PA_ND block tables. None elsewhere.
out_cache_loc_dsv4: Optional[DSV4OutCacheLoc] = None out_cache_loc_dsv4: Optional[DSV4OutCacheLoc] = None
@@ -1836,6 +1839,8 @@ def build_inner_fb_view(
seq_lens_cpu=forward_batch.seq_lens_cpu, seq_lens_cpu=forward_batch.seq_lens_cpu,
encoder_lens=encoder_lens, encoder_lens=encoder_lens,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None), out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
# A caller may hand in another view that does not carry this field.
out_cache_loc_virtual=getattr(forward_batch, "out_cache_loc_virtual", None),
out_cache_loc_dsv4=getattr(forward_batch, "out_cache_loc_dsv4", None), out_cache_loc_dsv4=getattr(forward_batch, "out_cache_loc_dsv4", None),
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
) )
@@ -192,6 +192,7 @@ def build_replay_fb_view(
num_padding=bs - raw_bs, num_padding=bs - raw_bs,
encoder_lens=buffers.encoder_lens[:bs] if is_encoder_decoder else None, encoder_lens=buffers.encoder_lens[:bs] if is_encoder_decoder else None,
out_cache_loc=getattr(forward_batch, "out_cache_loc", None), out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
out_cache_loc_virtual=forward_batch.out_cache_loc_virtual,
out_cache_loc_dsv4=getattr(forward_batch, "out_cache_loc_dsv4", None), out_cache_loc_dsv4=getattr(forward_batch, "out_cache_loc_dsv4", None),
# The mamba-track registry slot (VIRTUAL ids) is the v2p translate SOURCE # The mamba-track registry slot (VIRTUAL ids) is the v2p translate SOURCE
# for the backend, which copies the result into its own static buffer and # for the backend, which copies the result into its own static buffer and
@@ -432,6 +432,89 @@ class TestFusedFp8WriteGate(CustomTestCase):
) )
@unittest.skipUnless(torch.cuda.is_available(), "the fused translate is Triton")
class TestFusedWriteLocTranslateCuda(CustomTestCase):
"""The fused write-loc translate must agree with its own CPU branch.
The two implementations must not drift: Triton truncates division toward
zero where torch floors it, so a negative loc and a tombstoned v2p row are
where a divergence would appear -- and the CPU branch is all the CPU suites
ever exercise.
"""
def test_cuda_matches_the_cpu_branch(self):
from sglang.kernels.ops.memory.virtual_slot import write_loc_to_kernel_ids
for page_size in (1, 64):
span = page_size * 4
loc = torch.tensor(
[-1, 0, 1, page_size, span, span + 1, 2 * span + 3, 5 * page_size],
dtype=torch.int64,
)
# Size the table past the highest page any loc can name, then
# scramble it and tombstone one row (-1), so neither a dropped
# clamp nor a skipped gather can coincide with the right answer.
num_pages = int(loc.max()) // page_size + 2
v2p = torch.tensor(
[(5 * i + 2) % num_pages for i in range(num_pages)] + [-1],
dtype=torch.int64,
)
v2p[1] = -1
for dcp_size, dcp_rank in ((1, 0), (2, 1), (4, 2)):
kw = dict(
page_size=page_size,
stride=page_size * 3,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
)
cpu = write_loc_to_kernel_ids(loc=loc, v2p=v2p, **kw)
gpu = write_loc_to_kernel_ids(loc=loc.cuda(), v2p=v2p.cuda(), **kw)
self.assertEqual(
cpu.tolist(),
gpu.cpu().tolist(),
f"ps={page_size} dcp_size={dcp_size} rank={dcp_rank}",
)
def test_wide_out_clears_the_stale_tail(self):
"""`out_width` past the batch must zero the tail in the same launch.
This is what lets a backend hand in its whole capture-stable buffer:
a shorter replay leaves stale kernel-facing ids past the batch, and the
captured write kernel consumes the full buffer, so an uncleared tail
scatters pad rows into live KV pages.
"""
from sglang.kernels.ops.memory.virtual_slot import write_loc_to_kernel_ids
v2p = torch.tensor([2, 5, 1, 3, 4], dtype=torch.int64, device="cuda")
loc = torch.tensor([0, 64, 128], dtype=torch.int64, device="cuda")
width = 8
# Poison the whole buffer so an unwritten or uncleared cell is visible.
buf = torch.full((width,), -999, dtype=torch.int64, device="cuda")
write_loc_to_kernel_ids(
loc=loc, v2p=v2p, page_size=64, stride=64 * 2, out=buf, out_width=width
)
self.assertEqual(buf[:3].tolist(), [2 * 128, 5 * 128, 1 * 128])
self.assertEqual(buf[3:].tolist(), [0] * (width - 3))
# And it must agree with the narrow call on the live prefix.
narrow = write_loc_to_kernel_ids(loc=loc, v2p=v2p, page_size=64, stride=64 * 2)
self.assertEqual(narrow.tolist(), buf[:3].tolist())
def test_out_is_written_in_place(self):
# The captured decode path hands in a capture-stable buffer; rebinding
# instead of filling it would leave the graph on a stale pointer.
from sglang.kernels.ops.memory.virtual_slot import write_loc_to_kernel_ids
v2p = torch.tensor([2, 5, -1, 3], dtype=torch.int64, device="cuda")
loc = torch.tensor([0, 64, 128, 192], dtype=torch.int64, device="cuda")
dst = torch.full_like(loc, -7)
ret = write_loc_to_kernel_ids(
loc=loc, v2p=v2p, page_size=64, stride=64 * 2, out=dst
)
self.assertIs(ret, dst)
self.assertEqual(dst.tolist(), [2 * 128, 5 * 128, 0, 3 * 128])
class TestDcpDecodeLayout(CustomTestCase): class TestDcpDecodeLayout(CustomTestCase):
"""Rank-local length math the decode page table above is built from.""" """Rank-local length math the decode page table above is built from."""
@@ -49,11 +49,15 @@ def _reference_chain(
out_buf: torch.Tensor, out_buf: torch.Tensor,
valid_bs: int, valid_bs: int,
total_bs: int, total_bs: int,
v2p: torch.Tensor = None,
) -> None: ) -> None:
"""Replicates the _replay_metadata reference ops, in order, in place.""" """Replicates the _replay_metadata reference ops, in order, in place."""
req_pool_indices[valid_bs:total_bs] = 0 req_pool_indices[valid_bs:total_bs] = 0
mamba_indices = mapping[req_pool_indices[:total_bs]] mamba_indices = mapping[req_pool_indices[:total_bs]]
# static pool: _translate_mamba_indices is the identity if v2p is not None:
# Unified pool: virtual->physical slot translate, and it runs BEFORE
# the padding sentinel -- captured kernels read physical ids.
mamba_indices = v2p[mamba_indices].to(torch.int32)
mamba_indices[valid_bs:] = -1 mamba_indices[valid_bs:] = -1
out_buf[: len(mamba_indices)].copy_(mamba_indices) out_buf[: len(mamba_indices)].copy_(mamba_indices)
@@ -139,6 +143,79 @@ class TestFusedReplayStateIndices(CustomTestCase):
f"{name} guard tail clobbered ({case}): {buf[total_bs:].tolist()}", f"{name} guard tail clobbered ({case}): {buf[total_bs:].tolist()}",
) )
def _run_v2p_case(self, total_bs: int, num_padding: int, seed: int) -> None:
"""Same equivalence, with the unified pool's virtual slot ids.
The mapping yields VIRTUAL slots there and the kernel folds the v2p
gather in, so the two must still agree element for element.
"""
device = torch.device("cuda")
gen = torch.Generator(device="cpu").manual_seed(seed + 1000)
valid_bs = total_bs - num_padding
req_pool = torch.randint(
0, _REQ_POOL_SIZE, (total_bs + _GUARD,), generator=gen, dtype=torch.int64
)
req_pool[total_bs:] = _GUARD_SENTINEL
mapping = torch.randint(
0, _MAMBA_POOL_SIZE, (_REQ_POOL_SIZE,), generator=gen, dtype=torch.int32
)
# Scrambled table with a tombstone, so a skipped gather cannot pass.
v2p = torch.randperm(_MAMBA_POOL_SIZE + 1, generator=gen).to(torch.int64)
v2p[7] = -1
out = torch.full((total_bs + _GUARD,), _OUT_POISON, dtype=torch.int32)
req_ref, req_fused = req_pool.clone().to(device), req_pool.clone().to(device)
out_ref, out_fused = out.clone().to(device), out.clone().to(device)
mapping_d, v2p_d = mapping.to(device), v2p.to(device)
_reference_chain(
req_pool_indices=req_ref,
mapping=mapping_d,
out_buf=out_ref,
valid_bs=valid_bs,
total_bs=total_bs,
v2p=v2p_d,
)
returned = fused_replay_state_indices(
req_pool_indices=req_fused,
mamba_index_mapping=mapping_d,
out_state_indices=out_fused,
valid_bs=valid_bs,
total_bs=total_bs,
v2p=v2p_d,
)
case = f"v2p {total_bs=} {num_padding=} {seed=}"
self.assertTrue(
torch.equal(out_ref[:total_bs], out_fused[:total_bs]),
f"state indices mismatch ({case}):\n"
f" ref {out_ref[:total_bs].tolist()}\n"
f" fused {out_fused[:total_bs].tolist()}",
)
self.assertTrue(torch.equal(returned, out_fused[:total_bs]), case)
self.assertTrue(
torch.equal(req_pool_ref_tail := req_ref[total_bs:], req_fused[total_bs:]),
f"guard tail diverged ({case}): {req_pool_ref_tail.tolist()}",
)
self.assertTrue(
(out_fused[total_bs:] == _OUT_POISON).all(),
f"out guard tail clobbered ({case})",
)
def test_v2p_matrix(self):
for total_bs in (1, 7, 32, 33):
paddings = sorted(
{0, 1, total_bs // 2, total_bs - 1} & set(range(total_bs))
)
for num_padding in paddings:
for seed in (0, 1):
with self.subTest(
total_bs=total_bs, num_padding=num_padding, seed=seed
):
self._run_v2p_case(
total_bs=total_bs, num_padding=num_padding, seed=seed
)
def test_matrix(self): def test_matrix(self):
# Non-power-of-two sizes (7, 33) exercise the BS_UPPER in_range mask; # Non-power-of-two sizes (7, 33) exercise the BS_UPPER in_range mask;
# num_padding sweeps none / one / half / all-but-one padded rows. # num_padding sweeps none / one / half / all-but-one padded rows.
@@ -3100,5 +3100,82 @@ class TestDcpWidening(unittest.TestCase):
self.assertLess(base_cost * full_entry, mamba_bytes * dcp_size) self.assertLess(base_cost * full_entry, mamba_bytes * dcp_size)
class TestFusedWriteLocTranslate(unittest.TestCase):
"""`write_loc_to_kernel_ids` must equal the arithmetic it stands for.
Nothing downstream can tell the two apart except by value, so the reference
here is the definition rather than a recorded expectation. Triton truncates
division toward zero where torch floors it, so the negative-loc and
tombstoned-page cases are the ones that matter.
"""
def _reference(self, loc, v2p, page_size, stride, dcp_size, dcp_rank):
out = []
for raw in loc.tolist():
if raw < 0 or (dcp_size > 1 and raw % dcp_size != dcp_rank):
out.append(0)
continue
collapsed = raw // dcp_size
page = collapsed // page_size
offset = collapsed % page_size if page_size > 1 else 0
out.append(max(int(v2p[page]) * stride + offset, 0))
return out
def _check(self, *, page_size, multiplier, dcp_size, dcp_rank, device):
from sglang.kernels.ops.memory.virtual_slot import write_loc_to_kernel_ids
span = page_size * dcp_size
locs = [0, 1, span - 1, span, 2 * span + 3, 5 * span + dcp_rank, -1]
locs += [3 * span + dcp_rank] # lands on the tombstoned page
loc = torch.tensor(locs, dtype=torch.int64, device=device)
# Table sized past the highest page any loc can name, scrambled, with
# one tombstone (-1) so a missing clamp shows up.
num_pages = max(locs) // (page_size * dcp_size) + 2
v2p = torch.tensor(
[(5 * i + 2) % num_pages for i in range(num_pages)] + [-1],
dtype=torch.int64,
device=device,
)
v2p[min(3, num_pages - 1)] = -1
stride = page_size * multiplier
got = write_loc_to_kernel_ids(
loc=loc,
v2p=v2p,
page_size=page_size,
stride=stride,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
)
want = self._reference(
loc.cpu(), v2p.cpu(), page_size, stride, dcp_size, dcp_rank
)
self.assertEqual(got.tolist(), want, f"ps={page_size} dcp={dcp_size}")
# `out=` must write in place and agree (the cuda-graph-stable path).
dst = torch.full_like(loc, -7)
ret = write_loc_to_kernel_ids(
loc=loc,
v2p=v2p,
page_size=page_size,
stride=stride,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
out=dst,
)
self.assertIs(ret, dst)
self.assertEqual(dst.tolist(), want)
def test_matches_reference_on_cpu(self):
for page_size in (1, 64):
for dcp_size, dcp_rank in ((1, 0), (2, 1), (4, 2)):
self._check(
page_size=page_size,
multiplier=7,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
device="cpu",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -55,6 +55,7 @@ def _run_handler(*, prefill_backend, explicit):
"speculative_eagle_topk": None, "speculative_eagle_topk": None,
"enable_hierarchical_cache": False, "enable_hierarchical_cache": False,
"enable_lmcache": False, "enable_lmcache": False,
"enable_two_batch_overlap": False,
"dcp_size": 1, "dcp_size": 1,
"cuda_graph_config": cg, "cuda_graph_config": cg,
"cuda_graph_backend_prefill": prefill_backend if explicit else None, "cuda_graph_backend_prefill": prefill_backend if explicit else None,
@@ -0,0 +1,69 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""`--enable-unified-memory` refuses `--enable-two-batch-overlap`.
BUG REGRESSION. The combination launches and captures fine, then dies in the
forward path on the first captured decode replay. Nothing else rejects the
pair, so without this gate a running server crashes mid-serving.
python -m pytest test/registered/unit/server_args/test_unified_tbo_gate.py -v
"""
import unittest
from types import SimpleNamespace
from sglang.srt.arg_groups.kv_cache_hook import handle_unified_memory_pool
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _run_handler(*, unified, tbo):
"""Run just `handle_unified_memory_pool` over a minimal stand-in."""
sa = ServerArgs.__new__(ServerArgs)
for name, value in {
"enable_unified_memory": unified,
"enable_two_batch_overlap": tbo,
"disaggregation_mode": "null",
"speculative_algorithm": None,
"speculative_eagle_topk": None,
"enable_hierarchical_cache": False,
"enable_lmcache": False,
"dcp_size": 1,
"cuda_graph_config": SimpleNamespace(
prefill=SimpleNamespace(backend=Backend.DISABLED),
decode=SimpleNamespace(backend=Backend.FULL),
),
"cuda_graph_backend_prefill": Backend.DISABLED,
}.items():
object.__setattr__(sa, name, value)
handle_unified_memory_pool(sa)
class TestUnifiedTboGate(unittest.TestCase):
def test_tbo_with_unified_memory_is_refused(self):
with self.assertRaises(AssertionError) as ctx:
_run_handler(unified=True, tbo=True)
self.assertIn("two-batch-overlap", str(ctx.exception))
def test_gate_fires_only_on_the_pair(self):
"""An inverted condition here would reject every unified launch."""
_run_handler(unified=True, tbo=False)
_run_handler(unified=False, tbo=True)
if __name__ == "__main__":
unittest.main()