[AMD][DSV4] Reland unified-KV pool sizing and SWA ring accounting, fully gated (#38192)

Co-authored-by: hnyls2002 <lsyincs@gmail.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
yuttian1
2026-09-07 13:13:04 -07:00
committed by GitHub
co-authored by hnyls2002 Liangsheng Yin
parent 6287ebf43a
commit 570087ceda
20 changed files with 763 additions and 91 deletions
@@ -50,6 +50,7 @@ struct Prefill0Params {
/// \brief Trailing tokens the write plan keeps resident in the compress state ring. /// \brief Trailing tokens the write plan keeps resident in the compress state ring.
/// Derived from the ring in `plan_compress_prefill`; see the bound there. /// Derived from the ring in `plan_compress_prefill`; see the bound there.
int32_t mtp_pad; int32_t mtp_pad;
bool use_req_ring;
}; };
struct Prefill1Params { struct Prefill1Params {
@@ -67,6 +68,7 @@ struct Prefill1Params {
int32_t swa_page_size; int32_t swa_page_size;
int32_t ring_size; int32_t ring_size;
int32_t compress_ratio; int32_t compress_ratio;
bool use_req_ring;
}; };
struct DecodeParams { struct DecodeParams {
@@ -80,6 +82,7 @@ struct DecodeParams {
int32_t swa_page_size; int32_t swa_page_size;
int32_t ring_size; int32_t ring_size;
int32_t compress_ratio; int32_t compress_ratio;
bool use_req_ring;
}; };
struct Prefill1ParamsLegacy { struct Prefill1ParamsLegacy {
@@ -203,7 +206,7 @@ __global__ __launch_bounds__(1024, 1) //
const int32_t last_c_pos = (sl / cr) * cr; const int32_t last_c_pos = (sl / cr) * cr;
const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad);
bool do_write = position >= first_w_pos; bool do_write = position >= first_w_pos;
if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); if (!do_write && is_overlap && !params.use_req_ring) do_write = (position % sps) >= (sps - cr);
if (do_write) { if (do_write) {
const uint32_t out_idx = atomicAdd(&counter_w, 1u); const uint32_t out_idx = atomicAdd(&counter_w, 1u);
params.plan_w[out_idx] = pack_w(ragged_id, batch_id, position + 1); params.plan_w[out_idx] = pack_w(ragged_id, batch_id, position + 1);
@@ -236,7 +239,7 @@ __global__ __launch_bounds__(1024, 1) //
} }
bool do_write = position >= first_w_pos; bool do_write = position >= first_w_pos;
if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); if (!do_write && is_overlap && !params.use_req_ring) do_write = (position % sps) >= (sps - cr);
if (do_write) { if (do_write) {
const uint32_t out_idx = atomicAdd(&counter_w, 1u); const uint32_t out_idx = atomicAdd(&counter_w, 1u);
params.plan_w[out_idx] = pack_w(ragged_id, static_cast<uint32_t>(batch_id), position + 1); params.plan_w[out_idx] = pack_w(ragged_id, static_cast<uint32_t>(batch_id), position + 1);
@@ -270,7 +273,7 @@ __global__ void plan_compress_prefill_kernel_1(const Prefill1Params params) {
const auto ring_offset = swa_loc % params.ring_size; const auto ring_offset = swa_loc % params.ring_size;
return swa_page * params.ring_size + ring_offset; return swa_page * params.ring_size + ring_offset;
}; };
const auto compute_c128_loc = [&](int64_t rid, int32_t position) { const auto compute_req_ring_loc = [&](int64_t rid, int32_t position) {
return static_cast<int32_t>(rid * params.ring_size + position % params.ring_size); return static_cast<int32_t>(rid * params.ring_size + position % params.ring_size);
}; };
@@ -283,9 +286,9 @@ __global__ void plan_compress_prefill_kernel_1(const Prefill1Params params) {
const auto position_1 = static_cast<int32_t>(plan_c.seq_len - 1); const auto position_1 = static_cast<int32_t>(plan_c.seq_len - 1);
// only used for c4, harmless for c128 // only used for c4, harmless for c128
const auto position_0 = max(position_1 - params.compress_ratio, 0); const auto position_0 = max(position_1 - params.compress_ratio, 0);
if (params.compress_ratio == 128) { if (params.compress_ratio == 128 || params.use_req_ring) {
plan_c.read_page_0 = compute_c128_loc(rid, position_0) / 128; plan_c.read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio;
plan_c.read_page_1 = compute_c128_loc(rid, position_1) / 128; plan_c.read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio;
} else { } else {
const auto raw_loc_0 = mapping[position_0]; const auto raw_loc_0 = mapping[position_0];
const auto raw_loc_1 = mapping[position_1]; const auto raw_loc_1 = mapping[position_1];
@@ -307,8 +310,8 @@ __global__ void plan_compress_prefill_kernel_1(const Prefill1Params params) {
// `seq_len` (`write_loc`) may not be aligned here // `seq_len` (`write_loc`) may not be aligned here
const auto position = static_cast<int32_t>(plan_w.write_loc - 1); const auto position = static_cast<int32_t>(plan_w.write_loc - 1);
plan_w.ragged_id = ragged_id; plan_w.ragged_id = ragged_id;
if (params.compress_ratio == 128) { if (params.compress_ratio == 128 || params.use_req_ring) {
plan_w.write_loc = compute_c128_loc(rid, position); plan_w.write_loc = compute_req_ring_loc(rid, position);
} else { } else {
const auto raw_loc = mapping[position]; const auto raw_loc = mapping[position];
plan_w.write_loc = compute_loc(params.f2s_ptr[raw_loc]); plan_w.write_loc = compute_loc(params.f2s_ptr[raw_loc]);
@@ -329,7 +332,7 @@ __global__ void plan_compress_decode_kernel(const DecodeParams params) {
const auto ring_offset = swa_loc % params.ring_size; const auto ring_offset = swa_loc % params.ring_size;
return swa_page * params.ring_size + ring_offset; return swa_page * params.ring_size + ring_offset;
}; };
const auto compute_c128_loc = [&](int64_t rid, int32_t position) { const auto compute_req_ring_loc = [&](int64_t rid, int32_t position) {
return static_cast<int32_t>(rid * params.ring_size + position % params.ring_size); return static_cast<int32_t>(rid * params.ring_size + position % params.ring_size);
}; };
const auto seq_len = static_cast<int32_t>(params.seq_ptr[idx]); const auto seq_len = static_cast<int32_t>(params.seq_ptr[idx]);
@@ -338,10 +341,10 @@ __global__ void plan_compress_decode_kernel(const DecodeParams params) {
int32_t write_loc; int32_t write_loc;
int32_t read_page_0; int32_t read_page_0;
int32_t read_page_1; int32_t read_page_1;
if (params.compress_ratio == 128) { if (params.compress_ratio == 128 || params.use_req_ring) {
write_loc = compute_c128_loc(rid, position_1); write_loc = compute_req_ring_loc(rid, position_1);
read_page_0 = compute_c128_loc(rid, position_0) / 128; read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio;
read_page_1 = compute_c128_loc(rid, position_1) / 128; read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio;
} else { } else {
const auto raw_loc_0 = mapping[position_0]; const auto raw_loc_0 = mapping[position_0];
const auto raw_loc_1 = mapping[position_1]; const auto raw_loc_1 = mapping[position_1];
@@ -461,6 +464,7 @@ inline PrefillPlan plan_compress_prefill(
const int32_t compress_ratio, const int32_t compress_ratio,
const int32_t swa_page_size, const int32_t swa_page_size,
const int32_t ring_size, const int32_t ring_size,
const bool use_req_ring,
const bool use_cuda_graph) { const bool use_cuda_graph) {
auto B = SymbolicSize{"batch_size"}; auto B = SymbolicSize{"batch_size"};
auto N = SymbolicSize{"num_q_tokens"}; auto N = SymbolicSize{"num_q_tokens"};
@@ -503,6 +507,7 @@ inline PrefillPlan plan_compress_prefill(
const auto batch_size = static_cast<uint32_t>(B.unwrap()); const auto batch_size = static_cast<uint32_t>(B.unwrap());
constexpr auto kMaxTokens = static_cast<uint32_t>(std::numeric_limits<uint16_t>::max()); constexpr auto kMaxTokens = static_cast<uint32_t>(std::numeric_limits<uint16_t>::max());
RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); RuntimeCheck(compress_ratio == 4 || compress_ratio == 128);
RuntimeCheck(!use_req_ring || compress_ratio == 4);
RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens);
// `swa_page_size` >= `ring_size` >= `compress_ratio` // `swa_page_size` >= `ring_size` >= `compress_ratio`
RuntimeCheck(swa_page_size % ring_size == 0 && ring_size % compress_ratio == 0); RuntimeCheck(swa_page_size % ring_size == 0 && ring_size % compress_ratio == 0);
@@ -537,6 +542,7 @@ inline PrefillPlan plan_compress_prefill(
.compress_ratio = compress_ratio, .compress_ratio = compress_ratio,
.swa_page_size = swa_page_size, .swa_page_size = swa_page_size,
.mtp_pad = mtp_pad, .mtp_pad = mtp_pad,
.use_req_ring = use_req_ring,
}; };
LaunchKernel(1, kMaxPrefillBatchSize, device)(plan_compress_prefill_kernel0, params0); LaunchKernel(1, kMaxPrefillBatchSize, device)(plan_compress_prefill_kernel0, params0);
// kernel_1 sees the already-padded buffers, so num_c == num_w == num_padded == num_q_tokens. // kernel_1 sees the already-padded buffers, so num_c == num_w == num_padded == num_q_tokens.
@@ -555,6 +561,7 @@ inline PrefillPlan plan_compress_prefill(
.swa_page_size = swa_page_size, .swa_page_size = swa_page_size,
.ring_size = ring_size, .ring_size = ring_size,
.compress_ratio = compress_ratio, .compress_ratio = compress_ratio,
.use_req_ring = use_req_ring,
}; };
const auto block_size_1 = 256; const auto block_size_1 = 256;
const auto num_blocks_1 = div_ceil(params1.num_work, block_size_1); const auto num_blocks_1 = div_ceil(params1.num_work, block_size_1);
@@ -582,7 +589,7 @@ inline PrefillPlan plan_compress_prefill(
RuntimeCheck(0 < extend_len && extend_len <= seq_len); RuntimeCheck(0 < extend_len && extend_len <= seq_len);
const auto should_write = [=](int32_t position) { const auto should_write = [=](int32_t position) {
if (position >= first_w_pos) return true; if (position >= first_w_pos) return true;
return is_overlap && position % swa_page_size >= (swa_page_size - compress_ratio); return is_overlap && !use_req_ring && position % swa_page_size >= (swa_page_size - compress_ratio);
}; };
for (const auto j : irange(extend_len)) { for (const auto j : irange(extend_len)) {
const int32_t position = prefix_len + j; const int32_t position = prefix_len + j;
@@ -631,6 +638,7 @@ inline PrefillPlan plan_compress_prefill(
.swa_page_size = swa_page_size, .swa_page_size = swa_page_size,
.ring_size = ring_size, .ring_size = ring_size,
.compress_ratio = compress_ratio, .compress_ratio = compress_ratio,
.use_req_ring = use_req_ring,
}; };
const auto block_size = 256; const auto block_size = 256;
const auto num_blocks = div_ceil(params.num_work, block_size); const auto num_blocks = div_ceil(params.num_work, block_size);
@@ -645,7 +653,8 @@ inline tvm::ffi::Tensor plan_compress_decode(
const tvm::ffi::TensorView seq_lens, // CPU/GPU const tvm::ffi::TensorView seq_lens, // CPU/GPU
const int32_t compress_ratio, const int32_t compress_ratio,
const int32_t swa_page_size, const int32_t swa_page_size,
const int32_t ring_size) { const int32_t ring_size,
const bool use_req_ring) {
auto B = SymbolicSize{"batch_size"}; auto B = SymbolicSize{"batch_size"};
auto device_ = SymbolicDevice{}; auto device_ = SymbolicDevice{};
device_.set_options<kDLGPU>(); device_.set_options<kDLGPU>();
@@ -667,6 +676,7 @@ inline tvm::ffi::Tensor plan_compress_decode(
.with_device(device_) .with_device(device_)
.verify(seq_lens); .verify(seq_lens);
RuntimeCheck(!use_req_ring || compress_ratio == 4);
const auto batch_size = static_cast<uint32_t>(B.unwrap()); const auto batch_size = static_cast<uint32_t>(B.unwrap());
const auto device = device_.unwrap(); const auto device = device_.unwrap();
auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device); auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device);
@@ -681,6 +691,7 @@ inline tvm::ffi::Tensor plan_compress_decode(
.swa_page_size = swa_page_size, .swa_page_size = swa_page_size,
.ring_size = ring_size, .ring_size = ring_size,
.compress_ratio = compress_ratio, .compress_ratio = compress_ratio,
.use_req_ring = use_req_ring,
}; };
const auto block_size = 256; const auto block_size = 256;
const auto num_blocks = div_ceil(batch_size, block_size); const auto num_blocks = div_ceil(batch_size, block_size);
@@ -100,6 +100,7 @@ def create_paged_compress_data_kernel(
stride_out_1_1: tl.constexpr, stride_out_1_1: tl.constexpr,
compress_ratio: tl.constexpr, compress_ratio: tl.constexpr,
is_overlap: tl.constexpr, is_overlap: tl.constexpr,
use_req_ring: tl.constexpr,
swa_page_size: tl.constexpr, swa_page_size: tl.constexpr,
ring_size: tl.constexpr, ring_size: tl.constexpr,
BLOCK: tl.constexpr, BLOCK: tl.constexpr,
@@ -133,7 +134,7 @@ def create_paged_compress_data_kernel(
else: else:
pos = write_overlap_pos pos = write_overlap_pos
pos = tl.maximum(pos, 0) pos = tl.maximum(pos, 0)
if compress_ratio == 128: if compress_ratio == 128 or use_req_ring:
state_loc = rid * ring_size + (pos % ring_size) state_loc = rid * ring_size + (pos % ring_size)
else: else:
loc = tl.load( loc = tl.load(
@@ -182,6 +183,7 @@ def triton_create_paged_compress_data(
extend_seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor,
req_to_token: torch.Tensor, req_to_token: torch.Tensor,
full_to_swa_index_mapping: torch.Tensor, full_to_swa_index_mapping: torch.Tensor,
use_req_ring: bool = False,
block: int = 128, block: int = 128,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
batch_size = req_pool_indices.shape[0] batch_size = req_pool_indices.shape[0]
@@ -205,6 +207,7 @@ def triton_create_paged_compress_data(
stride_out_1_1=out_1.stride(1), # type: ignore stride_out_1_1=out_1.stride(1), # type: ignore
compress_ratio=compress_ratio, # type: ignore compress_ratio=compress_ratio, # type: ignore
is_overlap=1 if is_overlap else 0, # type: ignore is_overlap=1 if is_overlap else 0, # type: ignore
use_req_ring=1 if use_req_ring else 0, # type: ignore
swa_page_size=swa_page_size, # type: ignore swa_page_size=swa_page_size, # type: ignore
ring_size=ring_size, # type: ignore ring_size=ring_size, # type: ignore
BLOCK=block, # type: ignore BLOCK=block, # type: ignore
@@ -162,6 +162,7 @@ class CompressorDecodePlan(NamedTuple):
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
swa_page_size: int, swa_page_size: int,
ring_size: int, ring_size: int,
use_req_ring: bool = False,
) -> CompressorDecodePlan: ) -> CompressorDecodePlan:
if _is_xpu: if _is_xpu:
fn = plan_compress_decode fn = plan_compress_decode
@@ -169,7 +170,7 @@ class CompressorDecodePlan(NamedTuple):
module = _jit_compress_plan_module() module = _jit_compress_plan_module()
fn = module.plan_decode fn = module.plan_decode
plan_d = fn( args = (
req_pool_indices, req_pool_indices,
req_to_token, req_to_token,
full_to_state, full_to_state,
@@ -178,6 +179,10 @@ class CompressorDecodePlan(NamedTuple):
int(swa_page_size), int(swa_page_size),
int(ring_size), int(ring_size),
) )
assert not (_is_xpu and use_req_ring), (
"use_req_ring is not supported by the XPU compress plan builder"
)
plan_d = fn(*args) if _is_xpu else fn(*args, bool(use_req_ring))
return CompressorDecodePlan(compress_ratio, torch.from_dlpack(plan_d)) return CompressorDecodePlan(compress_ratio, torch.from_dlpack(plan_d))
@staticmethod @staticmethod
@@ -247,6 +252,7 @@ class CompressorPrefillPlan(NamedTuple):
ring_size: int, ring_size: int,
num_q_tokens: int, num_q_tokens: int,
use_cuda_graph: bool = False, use_cuda_graph: bool = False,
use_req_ring: bool = False,
) -> CompressorPrefillPlan: ) -> CompressorPrefillPlan:
is_gpu_input = seq_lens.device.type in ["cuda", "xpu"] is_gpu_input = seq_lens.device.type in ["cuda", "xpu"]
pin_buffer = torch.empty( pin_buffer = torch.empty(
@@ -274,7 +280,7 @@ class CompressorPrefillPlan(NamedTuple):
module = _jit_compress_plan_module() module = _jit_compress_plan_module()
fn = module.plan_prefill fn = module.plan_prefill
plan_c, plan_w = fn( args = (
req_pool_indices, req_pool_indices,
req_to_token, req_to_token,
full_to_state, full_to_state,
@@ -285,7 +291,14 @@ class CompressorPrefillPlan(NamedTuple):
int(compress_ratio), int(compress_ratio),
int(swa_page_size), int(swa_page_size),
int(ring_size), int(ring_size),
bool(use_cuda_graph), )
assert not (_is_xpu and use_req_ring), (
"use_req_ring is not supported by the XPU compress plan builder"
)
plan_c, plan_w = (
fn(*args, bool(use_cuda_graph))
if _is_xpu
else fn(*args, bool(use_req_ring), bool(use_cuda_graph))
) )
return CompressorPrefillPlan( return CompressorPrefillPlan(
compress_ratio, compress_ratio,
+18 -2
View File
@@ -27,7 +27,7 @@ from collections import deque
from concurrent.futures import Future from concurrent.futures import Future
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple
import numpy as np import numpy as np
import torch import torch
@@ -74,6 +74,7 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.managers.schedule_policy import match_prefix_for_req from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache, BasePrefixCache,
DecLockRefParams, DecLockRefParams,
@@ -139,6 +140,9 @@ class DecodeReqToTokenPool:
#running <= 8, #pre-allocated + #transfer <= pre_alloc_size, so we can use the free memory to pre-allocate requests to unblock prefill. #running <= 8, #pre-allocated + #transfer <= pre_alloc_size, so we can use the free memory to pre-allocate requests to unblock prefill.
""" """
# Mirrors ReqToTokenPool.register_on_alloc_rows.
_on_alloc_rows: Optional[Callable[[List[int]], None]] = None
def __init__( def __init__(
self, self,
size: int, size: int,
@@ -204,6 +208,8 @@ class DecodeReqToTokenPool:
return None return None
select_index = self.free_slots[:need_size] select_index = self.free_slots[:need_size]
self.free_slots = self.free_slots[need_size:] self.free_slots = self.free_slots[need_size:]
if self._on_alloc_rows is not None and select_index:
self._on_alloc_rows(select_index)
offset = 0 offset = 0
for r in reqs: for r in reqs:
if not r.kv.holds_kv: if not r.kv.holds_kv:
@@ -221,6 +227,10 @@ class DecodeReqToTokenPool:
self.free_slots = list(range(1, self._alloc_size)) self.free_slots = list(range(1, self._alloc_size))
self.req_generation.zero_() self.req_generation.zero_()
def register_on_alloc_rows(self, hook: Callable[[List[int]], None]) -> None:
assert self._on_alloc_rows is None
self._on_alloc_rows = hook
class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
def __init__( def __init__(
@@ -1711,7 +1721,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
window_size = self.scheduler.sliding_window_size or 0 window_size = self.scheduler.sliding_window_size or 0
swa_total = self.token_to_kv_pool_allocator.size_swa swa_total = self.token_to_kv_pool_allocator.size_swa
swa_available = self.token_to_kv_pool_allocator.swa_available_size() swa_available = self.token_to_kv_pool_allocator.swa_available_size()
swa_evictable = self.tree_cache.swa_evictable_size() # Per-request SWA ring: cached prefixes still report swa_evictable, but
# evicting them frees no ring space.
swa_evictable = (
0
if is_swa_req_ring(self.token_to_kv_pool_allocator)
else self.tree_cache.swa_evictable_size()
)
swa_used = swa_total - swa_available - swa_evictable swa_used = swa_total - swa_available - swa_evictable
swa_growth_potential = max(0, n_active * window_size - swa_used) swa_growth_potential = max(0, n_active * window_size - swa_used)
swa_reserved_tokens = min(reserved_tokens, swa_growth_potential) swa_reserved_tokens = min(reserved_tokens, swa_growth_potential)
@@ -118,6 +118,9 @@ class CompressorHip(_CompressorBase):
assert isinstance(backend, DeepseekV4HipRadixBackend) assert isinstance(backend, DeepseekV4HipRadixBackend)
token_to_kv_pool = backend.token_to_kv_pool token_to_kv_pool = backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
req_ring_state = self.ratio == 128 or (
self.ratio == 4 and token_to_kv_pool._unified_kv
)
state_pool = self._get_state_pool(backend) state_pool = self._get_state_pool(backend)
prefix_lens = forward_batch.extend_prefix_lens_cpu prefix_lens = forward_batch.extend_prefix_lens_cpu
@@ -144,7 +147,7 @@ class CompressorHip(_CompressorBase):
pre_state_indices = self.compute_state_len_indices( pre_state_indices = self.compute_state_len_indices(
seq_len=prefix_lens[i], ratio=self.ratio seq_len=prefix_lens[i], ratio=self.ratio
).to(device) ).to(device)
if self.ratio == 128: if req_ring_state:
state_loc = state_pool.translate_from_req_position_to_state_loc( state_loc = state_pool.translate_from_req_position_to_state_loc(
req_pool_indices[i], pre_state_indices req_pool_indices[i], pre_state_indices
) )
@@ -166,7 +169,7 @@ class CompressorHip(_CompressorBase):
post_state_len = post_state_indices.size(0) post_state_len = post_state_indices.size(0)
assert post_state_len <= valid_kv_len assert post_state_len <= valid_kv_len
if self.ratio == 128: if req_ring_state:
post_state_loc = state_pool.translate_from_req_position_to_state_loc( post_state_loc = state_pool.translate_from_req_position_to_state_loc(
req_pool_indices[i], post_state_indices req_pool_indices[i], post_state_indices
) )
@@ -260,6 +263,9 @@ class CompressorHip(_CompressorBase):
state_pool = self._get_state_pool(attn_backend) state_pool = self._get_state_pool(attn_backend)
token_to_kv_pool = attn_backend.token_to_kv_pool token_to_kv_pool = attn_backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
req_ring_state = self.ratio == 128 or (
self.ratio == 4 and token_to_kv_pool._unified_kv
)
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
req_to_token = attn_backend.req_to_token_pool.req_to_token req_to_token = attn_backend.req_to_token_pool.req_to_token
seq_lens = forward_batch.seq_lens seq_lens = forward_batch.seq_lens
@@ -271,7 +277,7 @@ class CompressorHip(_CompressorBase):
seq_lens = seq_lens_2d.view(-1) seq_lens = seq_lens_2d.view(-1)
req_pool_indices = req_pool_indices.repeat_interleave(draft_tokens) req_pool_indices = req_pool_indices.repeat_interleave(draft_tokens)
if self.ratio == 128: if req_ring_state:
state_locs = state_pool.translate_from_req_position_to_state_loc( state_locs = state_pool.translate_from_req_position_to_state_loc(
req_pool_indices, seq_lens - 1 req_pool_indices, seq_lens - 1
) )
@@ -286,7 +292,7 @@ class CompressorHip(_CompressorBase):
-compress_bulk_len, 0, device=seq_lens.device -compress_bulk_len, 0, device=seq_lens.device
) )
compress_indices.clamp_(min=-1) compress_indices.clamp_(min=-1)
if self.ratio == 128: if req_ring_state:
compress_indices_state = ( compress_indices_state = (
state_pool.translate_from_req_position_to_state_loc( state_pool.translate_from_req_position_to_state_loc(
req_pool_indices[:, None], compress_indices req_pool_indices[:, None], compress_indices
@@ -264,6 +264,7 @@ def create_paged_compressor_data(
) -> FusedCompressMetadata: ) -> FusedCompressMetadata:
swa_page_size = token_to_kv_pool.swa_page_size swa_page_size = token_to_kv_pool.swa_page_size
ring_size = token_to_kv_pool.get_ring_size(compress_ratio=compress_ratio) ring_size = token_to_kv_pool.get_ring_size(compress_ratio=compress_ratio)
use_req_ring = compress_ratio == 4 and token_to_kv_pool._unified_kv
# assert ring_size % compress_ratio == 0 # assert ring_size % compress_ratio == 0
def clip_down(positions: torch.Tensor) -> torch.Tensor: def clip_down(positions: torch.Tensor) -> torch.Tensor:
@@ -271,7 +272,7 @@ def create_paged_compressor_data(
def get_raw_loc(positions: torch.Tensor) -> torch.Tensor: def get_raw_loc(positions: torch.Tensor) -> torch.Tensor:
positions = positions.masked_fill(positions < 0, 0) positions = positions.masked_fill(positions < 0, 0)
if compress_ratio == 128: if compress_ratio == 128 or use_req_ring:
state_loc = req_pool_indices * ring_size + positions % ring_size state_loc = req_pool_indices * ring_size + positions % ring_size
else: else:
loc = req_to_token[req_pool_indices, positions] loc = req_to_token[req_pool_indices, positions]
@@ -294,6 +295,7 @@ def create_paged_compressor_data(
extend_seq_lens=extend_lens, extend_seq_lens=extend_lens,
req_to_token=req_to_token, req_to_token=req_to_token,
full_to_swa_index_mapping=token_to_kv_pool.full_to_swa_index_mapping, full_to_swa_index_mapping=token_to_kv_pool.full_to_swa_index_mapping,
use_req_ring=use_req_ring,
) )
plan_kwargs: dict plan_kwargs: dict
@@ -441,6 +441,7 @@ def create_paged_compressor_data(
swa_page_size = token_to_kv_pool.swa_page_size swa_page_size = token_to_kv_pool.swa_page_size
ring_size = token_to_kv_pool.get_ring_size(compress_ratio=compress_ratio) ring_size = token_to_kv_pool.get_ring_size(compress_ratio=compress_ratio)
use_req_ring = compress_ratio == 4 and token_to_kv_pool._unified_kv
# NOTE: This is actually a proxy, which encounter some bug with tvm-ffi. # NOTE: This is actually a proxy, which encounter some bug with tvm-ffi.
# As a workaround, we use `.detach()` to get the real tensor. # As a workaround, we use `.detach()` to get the real tensor.
full_to_swa = token_to_kv_pool.full_to_swa_index_mapping.detach() full_to_swa = token_to_kv_pool.full_to_swa_index_mapping.detach()
@@ -467,6 +468,7 @@ def create_paged_compressor_data(
full_to_state=full_to_swa, full_to_state=full_to_swa,
swa_page_size=swa_page_size, swa_page_size=swa_page_size,
ring_size=ring_size, ring_size=ring_size,
use_req_ring=use_req_ring,
num_q_tokens=num_q_tokens, num_q_tokens=num_q_tokens,
use_cuda_graph=use_prefill_cuda_graph, use_cuda_graph=use_prefill_cuda_graph,
) )
@@ -479,6 +481,7 @@ def create_paged_compressor_data(
seq_lens=seq_lens.to(torch.int64), seq_lens=seq_lens.to(torch.int64),
swa_page_size=swa_page_size, swa_page_size=swa_page_size,
ring_size=ring_size, ring_size=ring_size,
use_req_ring=use_req_ring,
) )
+33 -5
View File
@@ -50,6 +50,7 @@ from sglang.srt.mem_cache.allocator.hisparse import (
from sglang.srt.mem_cache.allocator.swa import ( from sglang.srt.mem_cache.allocator.swa import (
PureSWATokenToKVPoolAllocator, PureSWATokenToKVPoolAllocator,
SWATokenToKVPoolAllocator, SWATokenToKVPoolAllocator,
is_swa_req_ring,
) )
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
UnifiedMambaSWATokenToKVPoolAllocator, UnifiedMambaSWATokenToKVPoolAllocator,
@@ -500,6 +501,8 @@ class PrefillAdder:
self.prefill_tile_block_m = prefill_tile_block_m self.prefill_tile_block_m = prefill_tile_block_m
self.tree_cache = tree_cache self.tree_cache = tree_cache
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
# Per-request SWA ring: one fixed slot per request, not a token budget.
self._swa_req_ring = is_swa_req_ring(token_to_kv_pool_allocator)
self.running_batch = running_batch self.running_batch = running_batch
self.new_token_ratio = new_token_ratio self.new_token_ratio = new_token_ratio
self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens
@@ -659,8 +662,13 @@ class PrefillAdder:
@property @property
def rem_swa_tokens(self): def rem_swa_tokens(self):
allocator = self.token_to_kv_pool_allocator
if self._swa_req_ring:
# swa_available_size() already reports ring capacity; tree
# swa_evictable is in linear token units and frees no ring space.
return allocator.swa_available_size() - self.rem_swa_token_offset
return ( return (
self.token_to_kv_pool_allocator.swa_available_size() allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size() + self.tree_cache.swa_evictable_size()
- self.rem_swa_token_offset - self.rem_swa_token_offset
) )
@@ -707,6 +715,10 @@ class PrefillAdder:
where alloc = min(extend, rem_chunk); the min() cap keeps the two terms where alloc = min(extend, rem_chunk); the min() cap keeps the two terms
from double-counting extend, so budget <= extend + max_new_tokens + page. from double-counting extend, so budget <= extend + max_new_tokens + page.
""" """
allocator = self.token_to_kv_pool_allocator
if self._swa_req_ring:
# One ring slot per request, in the same unit as swa_available_size.
return allocator.swa_ring_cost_tokens
if self.rem_chunk_tokens is not None: if self.rem_chunk_tokens is not None:
alloc = min(extend_input_len, self.rem_chunk_tokens) alloc = min(extend_input_len, self.rem_chunk_tokens)
else: else:
@@ -834,6 +846,7 @@ class PrefillAdder:
max_new_tokens: int, max_new_tokens: int,
retracted_stain: bool, retracted_stain: bool,
mamba_gap_reserve: int = 0, mamba_gap_reserve: int = 0,
is_chunked_continuation: bool = False,
): ):
# TODO(lsyin): check this workaround logic, which only ensures the prefill will not out of memory, and may be too conservative # TODO(lsyin): check this workaround logic, which only ensures the prefill will not out of memory, and may be too conservative
extend_input_len = self.ceil_paged_tokens(extend_input_len) extend_input_len = self.ceil_paged_tokens(extend_input_len)
@@ -857,6 +870,9 @@ class PrefillAdder:
self.rem_input_tokens -= extend_input_len self.rem_input_tokens -= extend_input_len
if self.is_hybrid_swa: if self.is_hybrid_swa:
# The ring slot is reserved once at first admission; charging it
# again on a continuation would double-count and over-throttle.
if not (self._swa_req_ring and is_chunked_continuation):
self.rem_swa_token_offset += self._swa_budget_for_req( self.rem_swa_token_offset += self._swa_budget_for_req(
extend_input_len, max_new_tokens extend_input_len, max_new_tokens
) )
@@ -994,9 +1010,10 @@ class PrefillAdder:
_rem_tokens = self._get_dllm_remain_tokens() _rem_tokens = self._get_dllm_remain_tokens()
else: else:
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens)) _rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
if self.is_hybrid_swa: if self.is_hybrid_swa and not self._swa_req_ring:
# alloc_extend needs extend_num_tokens + page_size per request, # alloc_extend needs extend_num_tokens + page_size per request,
# so reserve one page here to avoid OOM # so reserve one page here to avoid OOM.
# Ring mode skips it: rem_swa_tokens counts slots, not chunk tokens.
_rem_tokens = min( _rem_tokens = min(
_rem_tokens, int(self.rem_swa_tokens) - self.page_size _rem_tokens, int(self.rem_swa_tokens) - self.page_size
) )
@@ -1035,6 +1052,7 @@ class PrefillAdder:
), ),
req.retracted_stain, req.retracted_stain,
mamba_gap_reserve=self._mamba_gap_budget_for_req(req), mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
is_chunked_continuation=True,
) )
# Return if chunked prefill not finished # Return if chunked prefill not finished
@@ -1238,7 +1256,13 @@ class PrefillAdder:
self._swa_new_tokens(req), self._swa_new_tokens(req),
swa_host_hit_length=req.swa_host_hit_length, swa_host_hit_length=req.swa_host_hit_length,
) )
if swa_needed >= self.rem_swa_tokens: # Ring-slot capacity is exact, so needing exactly what is left still
# fits; the legacy SWA-token path keeps its conservative `>=`.
if (
swa_needed > self.rem_swa_tokens
if self._swa_req_ring
else swa_needed >= self.rem_swa_tokens
):
if not self._swa_req_never_fits( if not self._swa_req_never_fits(
real_input_tokens, real_input_tokens,
self._swa_new_tokens(req), self._swa_new_tokens(req),
@@ -1274,7 +1298,11 @@ class PrefillAdder:
self._swa_new_tokens(req), self._swa_new_tokens(req),
swa_host_hit_length=req.swa_host_hit_length, swa_host_hit_length=req.swa_host_hit_length,
) )
if swa_needed >= self.rem_swa_tokens: if (
swa_needed > self.rem_swa_tokens
if self._swa_req_ring
else swa_needed >= self.rem_swa_tokens
):
if not self._swa_req_never_fits( if not self._swa_req_never_fits(
real_input_tokens, real_input_tokens,
self._swa_new_tokens(req), self._swa_new_tokens(req),
@@ -21,6 +21,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
SchedulerPoolStatsObserver, SchedulerPoolStatsObserver,
) )
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
UnifiedMambaSWATokenToKVPoolAllocator, UnifiedMambaSWATokenToKVPoolAllocator,
) )
@@ -152,6 +153,15 @@ class SchedulerInvariantChecker:
def _check_swa_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]: def _check_swa_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]:
allocator = self.token_to_kv_pool_allocator allocator = self.token_to_kv_pool_allocator
if is_swa_req_ring(allocator):
# Per-request SWA ring: there is no token pool to conserve; ring-slot
# leaks are caught by the req_to_token check instead.
return False, (
"[swa] unified ring (leak-check skipped): "
f"available={ps.swa_available_size}, "
f"evictable={ps.swa_evictable_size}, "
f"total={self.swa_tokens_per_layer}"
)
swa_available = ps.swa_available_size swa_available = ps.swa_available_size
if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator): if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator):
# Tri-pool: same floating-boundary phantom as the full pool -- use the # Tri-pool: same floating-boundary phantom as the full pool -- use the
@@ -11,6 +11,7 @@ from typing import (
Tuple, Tuple,
) )
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
UnifiedMambaSWATokenToKVPoolAllocator, UnifiedMambaSWATokenToKVPoolAllocator,
) )
@@ -301,6 +302,10 @@ class SchedulerPoolStatsObserver:
swa_available_size = allocator.swa_available_size() swa_available_size = allocator.swa_available_size()
full_evictable_size = self.tree_cache.full_evictable_size() full_evictable_size = self.tree_cache.full_evictable_size()
swa_evictable_size = self.tree_cache.swa_evictable_size() swa_evictable_size = self.tree_cache.swa_evictable_size()
# Per-request SWA ring: released with the req slot, yet cached radix
# prefixes still report swa_evictable; counting it drives usage negative.
if is_swa_req_ring(self.token_to_kv_pool_allocator):
swa_evictable_size = 0
full_num_used = self.full_tokens_per_layer - ( full_num_used = self.full_tokens_per_layer - (
full_available_size + full_evictable_size full_available_size + full_evictable_size
) )
+91 -4
View File
@@ -1,3 +1,5 @@
import logging
import torch import torch
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
@@ -8,6 +10,8 @@ from sglang.srt.utils import is_npu
from sglang.srt.utils.common import get_num_new_pages from sglang.srt.utils.common import get_num_new_pages
from sglang.srt.utils.invariants import Bucket, Invariant, IsTrue, expect from sglang.srt.utils.invariants import Bucket, Invariant, IsTrue, expect
logger = logging.getLogger(__name__)
_is_npu = is_npu() _is_npu = is_npu()
if _is_npu: if _is_npu:
@@ -28,6 +32,10 @@ _SWA_PEER_RELEASED = Invariant("swa.peer_released", Bucket.GUARD, IsTrue())
class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
"""Allocator for SWA hybrid KV cache.""" """Allocator for SWA hybrid KV cache."""
# Per-request SWA ring (BaseSWAKVPool.swa_req_ring_size). Class default so
# subclasses that bypass this __init__ read False.
_swa_req_ring = False
def __init__( def __init__(
self, self,
size: int, size: int,
@@ -37,6 +45,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
device: str, device: str,
kvcache: BaseSWAKVPool, kvcache: BaseSWAKVPool,
need_sort: bool, need_sort: bool,
req_to_token_pool=None,
): ):
assert isinstance(kvcache, BaseSWAKVPool) assert isinstance(kvcache, BaseSWAKVPool)
self._size_full = size self._size_full = size
@@ -104,10 +113,45 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.swa_free_group = [] self.swa_free_group = []
self._kvcache = kvcache self._kvcache = kvcache
# Per-request SWA ring: the paged SWA indices built here are unused and
# SWA capacity is bounded by req slots, not tokens.
ring_size = kvcache.swa_req_ring_size
self._swa_req_ring = ring_size is not None
self._req_to_token_pool = req_to_token_pool
if self._swa_req_ring:
assert req_to_token_pool is not None, (
"per-request SWA ring: capacity is counted in req slots"
)
self._swa_ring_cost = (
(ring_size + self.page_size - 1) // self.page_size
) * self.page_size
# Total SWA capacity is every req slot's ring; all slots are free here.
self._size_swa = req_to_token_pool.available_size() * self._swa_ring_cost
logger.info(
"SWA per-request ring accounting enabled: "
f"ring_size={ring_size}, ring_cost_tokens={self._swa_ring_cost}, "
f"size_swa={self._size_swa} (paged size_swa={size_swa} bypassed)"
)
else:
self._swa_ring_cost = 0
self.clear() self.clear()
self._kvcache.register_mapping(self.full_to_swa_index_mapping) self._kvcache.register_mapping(self.full_to_swa_index_mapping)
@property
def swa_req_ring(self) -> bool:
return self._swa_req_ring
@property
def swa_ring_cost_tokens(self) -> int:
return self._swa_ring_cost
def available_size(self): def available_size(self):
if self._swa_req_ring:
# The SWA ring is pre-allocated per slot and reused by decode, so it
# never constrains token growth; full attention is the real limiter.
return self.full_attn_allocator.available_size()
return min( return min(
self.full_attn_allocator.available_size(), self.full_attn_allocator.available_size(),
self.swa_attn_allocator.available_size(), self.swa_attn_allocator.available_size(),
@@ -117,6 +161,9 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return self.full_attn_allocator.available_size() return self.full_attn_allocator.available_size()
def swa_available_size(self): def swa_available_size(self):
if self._swa_req_ring:
# Ring-based availability: free request slots * per-slot ring cost.
return self._req_to_token_pool.available_size() * self._swa_ring_cost
return self.swa_attn_allocator.available_size() return self.swa_attn_allocator.available_size()
# Slot-conservation views for the leak invariant. On the non-shared allocator # Slot-conservation views for the leak invariant. On the non-shared allocator
@@ -142,7 +189,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def debug_print(self) -> str: def debug_print(self) -> str:
msg = "" msg = ""
msg += f"#swa-available-size: {self.swa_attn_allocator.available_size()}, " msg += f"#swa-available-size: {self.swa_available_size()}, "
msg += ( msg += (
f"#full-attn-available-size: {self.full_attn_allocator.available_size()}, " f"#full-attn-available-size: {self.full_attn_allocator.available_size()}, "
) )
@@ -171,11 +218,15 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return alloc_full_indices return alloc_full_indices
def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool: def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool:
return ( full_ok = (
num_full_pages num_full_pages
<= self.full_attn_allocator.available_size() // self.page_size <= self.full_attn_allocator.available_size() // self.page_size
and num_swa_pages )
<= self.swa_attn_allocator.available_size() // self.page_size if self._swa_req_ring:
# SWA ring rows are pre-allocated per slot; no per-token SWA paging.
return full_ok
return full_ok and (
num_swa_pages <= self.swa_attn_allocator.available_size() // self.page_size
) )
def alloc_extend( def alloc_extend(
@@ -195,6 +246,18 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if not self.new_pages_available(num_new_pages, num_new_pages): if not self.new_pages_available(num_new_pages, num_new_pages):
return None return None
if self._swa_req_ring:
# Ring mode pages full KV only; full_to_swa_index_mapping stays unwritten.
return self.full_attn_allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
num_new_pages=num_new_pages,
)
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc) swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
alloc_full_indices = self.full_attn_allocator.alloc_extend( alloc_full_indices = self.full_attn_allocator.alloc_extend(
@@ -245,6 +308,18 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if not self.new_pages_available(num_full_pages, num_swa_pages): if not self.new_pages_available(num_full_pages, num_swa_pages):
return None return None
if self._swa_req_ring:
# See alloc_extend: full KV only.
return self.full_attn_allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
num_new_pages=num_full_pages,
)
alloc_full_indices = self.full_attn_allocator.alloc_extend( alloc_full_indices = self.full_attn_allocator.alloc_extend(
prefix_lens, prefix_lens,
prefix_lens_cpu, prefix_lens_cpu,
@@ -291,6 +366,12 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
last_loc: torch.Tensor, # last_loc for full layers last_loc: torch.Tensor, # last_loc for full layers
): ):
assert self.page_size > 1 assert self.page_size > 1
if self._swa_req_ring:
# See alloc_extend: slot-addressed ring, so full-attention KV only.
return self.full_attn_allocator.alloc_decode(
seq_lens, seq_lens_cpu, last_loc
)
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc) swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
alloc_full_indices = self.full_attn_allocator.alloc_decode( alloc_full_indices = self.full_attn_allocator.alloc_decode(
@@ -453,6 +534,8 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
size_full = int(config.full_max_total_num_tokens) size_full = int(config.full_max_total_num_tokens)
size_swa = int(config.swa_max_total_num_tokens) size_swa = int(config.swa_max_total_num_tokens)
self._size_full = size_full self._size_full = size_full
if not self._swa_req_ring:
# Ring capacity follows the req slot count, not the token config.
self._size_swa = size_swa self._size_swa = size_swa
for alloc, sz in ( for alloc, sz in (
(self.full_attn_allocator, size_full), (self.full_attn_allocator, size_full),
@@ -625,3 +708,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
def clear(self): def clear(self):
self.swa_attn_allocator.clear() self.swa_attn_allocator.clear()
self.free_group = None self.free_group = None
def is_swa_req_ring(allocator) -> bool:
return isinstance(allocator, SWATokenToKVPoolAllocator) and allocator.swa_req_ring
@@ -1,5 +1,5 @@
import abc import abc
from typing import List, Tuple from typing import List, Optional, Tuple
import torch import torch
@@ -15,6 +15,9 @@ class BaseSWAKVPool(KVCache):
""" """
swa_kv_pool: KVCache swa_kv_pool: KVCache
# Set when SWA KV is a per-request ring of this many tokens (addressed by
# req_pool_idx) rather than a paged token pool; SWA is then not budgeted per token.
swa_req_ring_size: Optional[int] = None
@abc.abstractmethod @abc.abstractmethod
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor) -> None: def register_mapping(self, full_to_swa_index_mapping: torch.Tensor) -> None:
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging import logging
from contextlib import nullcontext from contextlib import nullcontext
from typing import List, Literal, NamedTuple, Optional, Tuple from typing import List, Literal, NamedTuple, Optional, Sequence, Tuple
import torch import torch
@@ -63,6 +63,12 @@ def get_compress_state_write_pad(compress_ratio: int, ring_size: int) -> int:
return ring_size - window_size + 2 if ring_size > window_size else 0 return ring_size - window_size + 2 if ring_size > window_size else 0
def get_swa_ring_size(sliding_window: int, is_speculative: bool = False) -> int:
# A verify batch writes its draft tokens ahead of the committed position.
spec_extra = (get_spec().speculative_num_draft_tokens - 1) if is_speculative else 0
return sliding_window + spec_extra
class DeepSeekV4SingleKVPool(KVCache): class DeepSeekV4SingleKVPool(KVCache):
def __init__( def __init__(
self, self,
@@ -566,6 +572,18 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.c4_size = c4_size self.c4_size = c4_size
self.c4_logical_size = c4_logical_size self.c4_logical_size = c4_logical_size
self.c128_size = c128_size self.c128_size = c128_size
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton,
)
# Resolve the unified-kv gate before any sizing so the two cannot drift.
self._unified_kv = is_unified_kv_triton()
c4_ring_size = self.get_ring_size(4)
if self._unified_kv:
# Unified C4 state is request-addressed: one ring per req slot,
# so the caller-supplied, SWA-scaled size does not apply here.
c4_state_pool_size = self.num_req_slots * c4_ring_size
# Non-unified (fp8) keeps the caller-supplied, SWA-addressed size.
self.c4_state_pool_size = c4_state_pool_size self.c4_state_pool_size = c4_state_pool_size
c128_ring_size = self.get_ring_size(128) c128_ring_size = self.get_ring_size(128)
if ONLINE_C128: if ONLINE_C128:
@@ -621,20 +639,12 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
c4_page_size = page_size // 4 c4_page_size = page_size // 4
c128_page_size = page_size // 128 c128_page_size = page_size // 128
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton,
)
self._unified_kv = is_unified_kv_triton()
if self._unified_kv: if self._unified_kv:
self.swa_kv_pool = None self.swa_kv_pool = None
self.c4_kv_pool = None self.c4_kv_pool = None
self.c128_kv_pool = None self.c128_kv_pool = None
spec_extra = ( swa_ring_size = get_swa_ring_size(
(get_spec().speculative_num_draft_tokens - 1) self.sliding_window, get_spec().speculative_algorithm is not None
if get_spec().speculative_algorithm is not None
else 0
) )
self.unified_kv_pool = DeepSeekV4UnifiedKVPool( self.unified_kv_pool = DeepSeekV4UnifiedKVPool(
stage_ratios=stage_ratios, stage_ratios=stage_ratios,
@@ -646,12 +656,13 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
device=device, device=device,
memory_saver_adapter=self.memory_saver_adapter, memory_saver_adapter=self.memory_saver_adapter,
custom_mem_pool=self.custom_mem_pool, custom_mem_pool=self.custom_mem_pool,
swa_ring_size=self.sliding_window + spec_extra, swa_ring_size=swa_ring_size,
) )
self.unified_swa_window = self.sliding_window self.unified_swa_window = self.sliding_window
self.unified_swa_ring_size = self.sliding_window + spec_extra self.unified_swa_ring_size = swa_ring_size
self.unified_swa_pages = self.unified_kv_pool.swa_pages self.unified_swa_pages = self.unified_kv_pool.swa_pages
self.swa_req_ring_size = self.unified_swa_ring_size
else: else:
self.unified_kv_pool = None self.unified_kv_pool = None
self.swa_kv_pool = self._make_kv_pool( self.swa_kv_pool = self._make_kv_pool(
@@ -1052,6 +1063,32 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
assert self.online_c128_mtp_pending_seq_lens is not None assert self.online_c128_mtp_pending_seq_lens is not None
return self.online_c128_mtp_pending_seq_lens return self.online_c128_mtp_pending_seq_lens
def clear_c4_req_states(self, req_pool_indices: Sequence[int]) -> None:
if not self._unified_kv or not req_pool_indices:
return
pools = [
pool
for pool in self.compress_state_pools + self.indexer_compress_state_pools
if pool is not None and pool.ratio == 4
]
if not pools:
return
ring_size = self.get_ring_size(4)
device = pools[0].kv_score_buffer.kv_score.device
req_indices = torch.as_tensor(req_pool_indices, dtype=torch.long, device=device)
state_locs = (
req_indices[:, None] * ring_size
+ torch.arange(ring_size, dtype=torch.long, device=device)
).flatten()
for pool in pools:
state = pool.kv_score_buffer.kv_score
half = state.shape[-1] // 2
state[state_locs, :half] = 0
state[state_locs, half:] = float("-inf")
def clear_c128_req_state(self, req_pool_idx: int) -> None: def clear_c128_req_state(self, req_pool_idx: int) -> None:
"""Reset request-scoped C128 state for one req slot.""" """Reset request-scoped C128 state for one req slot."""
for pool in self.compress_state_pools: for pool in self.compress_state_pools:
@@ -1078,7 +1115,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
accept_lens: torch.Tensor, accept_lens: torch.Tensor,
num_draft_tokens: int, num_draft_tokens: int,
) -> None: ) -> None:
"""Clear offline C128 ring slots written for rejected speculative tokens.""" """Clear offline C128 ring slots written for rejected speculative tokens.
C4 needs no counterpart: its draft states are overwritten in position order
before any read; a C128 compression boundary can read a stale draft slot."""
if ONLINE_C128 or num_draft_tokens <= 1 or req_pool_indices.numel() == 0: if ONLINE_C128 or num_draft_tokens <= 1 or req_pool_indices.numel() == 0:
return return
@@ -48,6 +48,7 @@ from sglang.srt.mem_cache.allocator.hisparse import (
from sglang.srt.mem_cache.allocator.swa import ( from sglang.srt.mem_cache.allocator.swa import (
PureSWATokenToKVPoolAllocator, PureSWATokenToKVPoolAllocator,
SWATokenToKVPoolAllocator, SWATokenToKVPoolAllocator,
is_swa_req_ring,
) )
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
UnifiedSWATokenToKVPoolAllocator, UnifiedSWATokenToKVPoolAllocator,
@@ -336,6 +337,18 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
) )
swa_max_total_num_tokens = sizes.swa_max_total_num_tokens
alloc = pools.token_to_kv_pool_allocator
if not self.is_draft_worker and is_swa_req_ring(alloc):
# Per-request SWA ring: the sizer's swa token count describes the
# vestigial paged pool; the allocator knows the real ring total.
swa_max_total_num_tokens = alloc.size_swa
logger.info(
"SWA ring: swa_max_total_num_tokens "
f"{sizes.swa_max_total_num_tokens} -> {swa_max_total_num_tokens} "
"(fixed per-request SWA ring capacity)."
)
logger.info( logger.info(
f"Memory pool end. " f"Memory pool end. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB" f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
@@ -345,7 +358,7 @@ class KVCacheConfigurator:
max_total_num_tokens=sizes.max_total_num_tokens, max_total_num_tokens=sizes.max_total_num_tokens,
max_running_requests=sizes.max_running_requests, max_running_requests=sizes.max_running_requests,
full_max_total_num_tokens=sizes.full_max_total_num_tokens, full_max_total_num_tokens=sizes.full_max_total_num_tokens,
swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, swa_max_total_num_tokens=swa_max_total_num_tokens,
req_to_token_pool=pools.req_to_token_pool, req_to_token_pool=pools.req_to_token_pool,
token_to_kv_pool=pools.token_to_kv_pool, token_to_kv_pool=pools.token_to_kv_pool,
token_to_kv_pool_allocator=pools.token_to_kv_pool_allocator, token_to_kv_pool_allocator=pools.token_to_kv_pool_allocator,
@@ -1348,6 +1361,12 @@ class KVCacheConfigurator:
enable_hisparse=get_memory().enable_hisparse, enable_hisparse=get_memory().enable_hisparse,
online_mtp_max_draft_tokens=(max_speculative_num_draft_tokens() or 0), online_mtp_max_draft_tokens=(max_speculative_num_draft_tokens() or 0),
) )
if not self.is_draft_worker and token_to_kv_pool._unified_kv:
# The draft pool has no C4 layers and shares this req pool, so only
# the target registers the per-slot C4 reset.
req_to_token_pool.register_on_alloc_rows(
token_to_kv_pool.clear_c4_req_states
)
return token_to_kv_pool return token_to_kv_pool
def _build_oot_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_oot_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
@@ -1979,6 +1998,7 @@ class KVCacheConfigurator:
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
need_sort=need_sort, need_sort=need_sort,
req_to_token_pool=req_to_token_pool,
) )
else: else:
if get_memory().enable_hisparse: if get_memory().enable_hisparse:
@@ -2275,6 +2295,12 @@ class KVCacheConfigurator:
max_tokens = self._apply_token_constraints(config.max_total_num_tokens) max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
if cap_tokens is not None: if cap_tokens is not None:
max_tokens = min(max_tokens, cap_tokens) max_tokens = min(max_tokens, cap_tokens)
# calculate_pool_sizes_from_max_tokens takes a token count, not a byte
# budget; it cannot re-subtract the fixed pools, so capacity must not rise.
assert max_tokens <= config.max_total_num_tokens, (
f"token constraints must not raise capacity: {max_tokens} > "
f"{config.max_total_num_tokens}"
)
if max_tokens != config.max_total_num_tokens: if max_tokens != config.max_total_num_tokens:
# Token-capped re-derivation: the profiled budget no longer # Token-capped re-derivation: the profiled budget no longer
# applies; the recalced config's unified_total_bytes stays None # applies; the recalced config's unified_total_bytes stays None
+10 -1
View File
@@ -31,7 +31,7 @@ import os
from contextlib import contextmanager, nullcontext from contextlib import contextmanager, nullcontext
from dataclasses import dataclass, fields from dataclasses import dataclass, fields
from functools import cached_property from functools import cached_property
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple, Union
import numpy as np import numpy as np
import torch import torch
@@ -259,6 +259,9 @@ class ReqToTokenPool:
"""A memory pool that maps a request to its token locations.""" """A memory pool that maps a request to its token locations."""
enable_mamba_extra_buffer_lazy: bool = False enable_mamba_extra_buffer_lazy: bool = False
# Class default: some decode pools borrow another __init__ (see
# DecodeReqToTokenPool) but inherit alloc_rows.
_on_alloc_rows: Optional[Callable[[List[int]], None]] = None
def __init__( def __init__(
self, self,
@@ -322,6 +325,8 @@ class ReqToTokenPool:
select_index = self.free_slots[-need_size:] select_index = self.free_slots[-need_size:]
del self.free_slots[-need_size:] del self.free_slots[-need_size:]
self.req_generation[select_index] += 1 self.req_generation[select_index] += 1
if self._on_alloc_rows is not None:
self._on_alloc_rows(select_index)
return select_index return select_index
def free_rows(self, indices: List[int]) -> None: def free_rows(self, indices: List[int]) -> None:
@@ -347,6 +352,10 @@ class ReqToTokenPool:
assert self._aux_cache is None assert self._aux_cache is None
self._aux_cache = aux_cache self._aux_cache = aux_cache
def register_on_alloc_rows(self, hook: Callable[[List[int]], None]) -> None:
assert self._on_alloc_rows is None
self._on_alloc_rows = hook
def reset_aux_cache_allocator(self) -> None: def reset_aux_cache_allocator(self) -> None:
if self._aux_cache is not None: if self._aux_cache is not None:
self._aux_cache.reset_allocator() self._aux_cache.reset_allocator()
@@ -38,6 +38,7 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
get_compress_state_ring_size, get_compress_state_ring_size,
get_compress_state_write_pad, get_compress_state_write_pad,
get_dsv4_indexer_bytes_per_token, get_dsv4_indexer_bytes_per_token,
get_swa_ring_size,
) )
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
@@ -875,7 +876,8 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
Splits available memory across full / swa / c4 / c128 + c4_state / c128_state Splits available memory across full / swa / c4 / c128 + c4_state / c128_state
pools. coeff is bytes_per_full_token (inflated by (T+D)/T when speculative pools. coeff is bytes_per_full_token (inflated by (T+D)/T when speculative
decode reserves a draft worker, mirroring dflash's cell_size scaling); bias = 0. decode reserves a draft worker, mirroring dflash's cell_size scaling). bias
is the request-scoped fixed pools that do not scale with full_token.
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
@@ -932,6 +934,16 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.num_layers_ca4 = sum(1 for r in self.compression_ratios if r == 4) self.num_layers_ca4 = sum(1 for r in self.compression_ratios if r == 4)
self.num_layers_ca128 = sum(1 for r in self.compression_ratios if r == 128) self.num_layers_ca128 = sum(1 for r in self.compression_ratios if r == 128)
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton,
)
self._unified = is_unified_kv_triton()
self.attn_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
# swa_page_size is the model's sliding window (cfg.window_size).
self._swa_ring_size = get_swa_ring_size(self.swa_page_size, self.is_speculative)
self._spec_infl = 1.0
if self.is_speculative: if self.is_speculative:
# Ring is sized once here, so it must serve the largest adaptive tier. # Ring is sized once here, so it must serve the largest adaptive tier.
self._assert_ring_serves_draft_tokens( self._assert_ring_serves_draft_tokens(
@@ -946,7 +958,8 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
# bytes_per_full_token: tokens = avail / (bpft * (T+D)/T). # bytes_per_full_token: tokens = avail / (bpft * (T+D)/T).
draft_layers = 1 draft_layers = 1
target_layers = self.num_layers_total target_layers = self.num_layers_total
self.bytes_per_full_token *= (target_layers + draft_layers) / target_layers self._spec_infl = (target_layers + draft_layers) / target_layers
self.bytes_per_full_token *= self._spec_infl
# Online c128 keeps a single in-progress (max, sum, kv) state per index # Online c128 keeps a single in-progress (max, sum, kv) state per index
# and assumes a strict forward-only schedule. Speculative decode (MTP) # and assumes a strict forward-only schedule. Speculative decode (MTP)
@@ -999,6 +1012,10 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
) )
def _get_bytes_per_full_token(self) -> float: def _get_bytes_per_full_token(self) -> float:
if self._unified:
# Unified_kv stores the whole latent in bf16.
kv_bytes = self.attn_head_dim * 2
else:
kv_bytes = self.qk_nope_head_dim + self.qk_rope_head_dim * 2 + 8 kv_bytes = self.qk_nope_head_dim + self.qk_rope_head_dim * 2 + 8
attn_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim attn_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
@@ -1023,28 +1040,52 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
c4_frac = 1 / (4 * self.c4_shrink_factor) c4_frac = 1 / (4 * self.c4_shrink_factor)
return ( return (
self.swa_ratio * kv_bytes * self.num_layers_total # Ring mode: SWA is a fixed per-request pool (see _fixed_swa_bytes).
(
0.0
if self._unified
else self.swa_ratio * kv_bytes * self.num_layers_total
)
+ c4_frac * kv_bytes * self.num_layers_ca4 + c4_frac * kv_bytes * self.num_layers_ca4
+ 1 / 128 * kv_bytes * self.num_layers_ca128 + 1 / 128 * kv_bytes * self.num_layers_ca128
+ 1 / 4 * self.indexer_bytes_per_token * self.num_layers_ca4 + 1 / 4 * self.indexer_bytes_per_token * self.num_layers_ca4
+ self.swa_ratio * c4_state_ratio * c4_state_bytes * self.num_layers_ca4 # Ring mode: C4 state is per-request too (see _fixed_c4_state_bytes).
+ (
0.0
if self._unified
else self.swa_ratio
* c4_state_ratio
* c4_state_bytes
* self.num_layers_ca4
)
+ c128_state_ratio * c128_state_bytes * self.num_layers_ca128 + c128_state_ratio * c128_state_bytes * self.num_layers_ca128
+ self.swa_ratio + (
0.0
if self._unified
else self.swa_ratio
* c4_state_ratio * c4_state_ratio
* c4_indexer_state_bytes * c4_indexer_state_bytes
* self.num_layers_ca4 * self.num_layers_ca4
) )
)
def _compute_dsv4_sizes(self, full_token: int, page_size: int) -> _DSV4PoolSizes: def _compute_dsv4_sizes(self, full_token: int, page_size: int) -> _DSV4PoolSizes:
full_token = full_token // page_size * page_size full_token = full_token // page_size * page_size
swa_tokens = int(full_token * self.swa_ratio) // page_size * page_size swa_tokens = int(full_token * self.swa_ratio) // page_size * page_size
if not self._unified:
# Ring mode: the paged SWA pool is vestigial, so its floor does not apply.
self.validate_swa_pool_size(swa_tokens, self.sliding_window_size, page_size) self.validate_swa_pool_size(swa_tokens, self.sliding_window_size, page_size)
return _DSV4PoolSizes( return _DSV4PoolSizes(
full_max_total_num_tokens=full_token, full_max_total_num_tokens=full_token,
swa_max_total_num_tokens=swa_tokens, swa_max_total_num_tokens=swa_tokens,
c4_max_total_num_tokens=full_token // (4 * self.c4_shrink_factor), c4_max_total_num_tokens=full_token // (4 * self.c4_shrink_factor),
c128_max_total_num_tokens=full_token // 128, c128_max_total_num_tokens=full_token // 128,
c4_state_pool_size=swa_tokens // self.swa_page_size * self.c4_ring_size, # Unified_kv: request-scoped, finalized once concurrency is known.
c4_state_pool_size=(
0
if self._unified
else swa_tokens // self.swa_page_size * self.c4_ring_size
),
c128_state_pool_size=0, c128_state_pool_size=0,
) )
@@ -1075,18 +1116,48 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
state_rows * state_last_dim * c128_state_dtype_size * self.num_layers_ca128 state_rows * state_last_dim * c128_state_dtype_size * self.num_layers_ca128
) )
def _get_c128_state_fixed_bytes_for_token_capacity( def _unified_c4_state_pool_size(self, max_running_requests: int) -> int:
self, token_capacity: int # Unified C4 state loc is req_pool_idx * c4_ring_size + pos % c4_ring_size.
) -> int: num_req_slots = self._get_num_req_slots(max_running_requests)
if self.requested_max_running_requests_per_worker is not None: return num_req_slots * self.c4_ring_size
return self._get_c128_state_fixed_bytes(
self.requested_max_running_requests_per_worker
)
estimated = int(token_capacity / self.context_len * 512) def _fixed_c4_state_bytes(self, max_running_requests: int) -> int:
if not self._unified or self.num_layers_ca4 == 0:
return 0
c4_state_dtype_size, _ = _get_dsv4_compress_state_dtype_sizes()
# Mirror CompressStatePool.__init__: it allocates `size + ring_size + 1`
# rows, padded to the compress ratio.
state_rows = self._unified_c4_state_pool_size(max_running_requests)
state_rows = ceil_div(state_rows + self.c4_ring_size + 1, 4) * 4
# overlap c4: last_dim = 2 * (1 + overlap) * head_dim = 4 * head_dim.
core_bytes = 4 * self.attn_head_dim * c4_state_dtype_size
indexer_bytes = 4 * self.indexer_head_dim * c4_state_dtype_size
return state_rows * (core_bytes + indexer_bytes) * self.num_layers_ca4
def _resolve_max_running_requests_per_worker(self, available_bytes: int) -> int:
# Approximates ModelRunner._resolve_max_num_reqs. Over-estimating is safe:
# a larger fixed bias yields a smaller full_token.
if self.requested_max_running_requests_per_worker is not None:
return self.requested_max_running_requests_per_worker
full_token = int(available_bytes / self.bytes_per_full_token)
estimated = int(full_token / self.context_len * 512)
estimated = max(min(estimated, 4096), 2048) estimated = max(min(estimated, 4096), 2048)
max_running_requests = min(estimated, token_capacity // 2) return min(estimated, full_token // 2)
return self._get_c128_state_fixed_bytes(max_running_requests)
def _fixed_swa_bytes(self, max_running_requests: int) -> int:
if not self._unified:
return 0
num_req_slots = self._get_num_req_slots(max_running_requests)
ring_bytes = (
num_req_slots
* self._swa_ring_size
* self.attn_head_dim
* 2 # bf16
* self.num_layers_total
)
return int(ring_bytes * self._spec_infl)
def _to_config(self, sizes: _DSV4PoolSizes) -> MemoryPoolConfig: def _to_config(self, sizes: _DSV4PoolSizes) -> MemoryPoolConfig:
full = sizes.full_max_total_num_tokens full = sizes.full_max_total_num_tokens
@@ -1117,6 +1188,11 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
config.c128_state_pool_size = num_req_slots config.c128_state_pool_size = num_req_slots
else: else:
config.c128_state_pool_size = num_req_slots * self.c128_ring_size config.c128_state_pool_size = num_req_slots * self.c128_ring_size
# Ring mode: C4 state is request-scoped, so size it from the known concurrency.
if self._unified and self.num_layers_ca4 > 0:
config.c4_state_pool_size = self._unified_c4_state_pool_size(
config.max_running_requests
)
return config return config
def calculate_pool_sizes( def calculate_pool_sizes(
@@ -1126,25 +1202,34 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
"page_size must be multiple of 128 for compressed attention" "page_size must be multiple of 128 for compressed attention"
) )
if self.requested_max_running_requests_per_worker is not None: max_running_requests_per_worker = self._resolve_max_running_requests_per_worker(
c128_state_fixed_bytes = self._get_c128_state_fixed_bytes( available_bytes
self.requested_max_running_requests_per_worker
) )
else: c128_state_fixed_bytes = self._get_c128_state_fixed_bytes(
full_token = int(available_bytes / self.bytes_per_full_token) max_running_requests_per_worker
c128_state_fixed_bytes = ( )
self._get_c128_state_fixed_bytes_for_token_capacity(full_token) swa_ring_fixed_bytes = self._fixed_swa_bytes(max_running_requests_per_worker)
c4_state_fixed_bytes = self._fixed_c4_state_bytes(
max_running_requests_per_worker
) )
available_bytes_for_tokens = max(available_bytes - c128_state_fixed_bytes, 0) available_bytes_for_tokens = max(
available_bytes
- c128_state_fixed_bytes
- swa_ring_fixed_bytes
- c4_state_fixed_bytes,
0,
)
full_token = int(available_bytes_for_tokens / self.bytes_per_full_token) full_token = int(available_bytes_for_tokens / self.bytes_per_full_token)
sizes = self._compute_dsv4_sizes(full_token, page_size) sizes = self._compute_dsv4_sizes(full_token, page_size)
logger.info( logger.info(
f"DSV4 memory calculation: " f"DSV4 memory calculation: unified={self._unified}, "
f"bytes_per_full_token={self.bytes_per_full_token:.2f}, " f"bytes_per_full_token={self.bytes_per_full_token:.2f}, "
f"available_bytes={available_bytes / (1 << 30):.2f} GB, " f"available_bytes={available_bytes / (1 << 30):.2f} GB, "
f"c128_state_fixed={c128_state_fixed_bytes / (1 << 30):.2f} GB, " f"c128_state_fixed={c128_state_fixed_bytes / (1 << 30):.2f} GB, "
f"swa_ring_fixed={swa_ring_fixed_bytes / (1 << 30):.2f} GB, "
f"c4_state_fixed={c4_state_fixed_bytes / (1 << 30):.2f} GB, "
f"full_token={sizes.full_max_total_num_tokens}" f"full_token={sizes.full_max_total_num_tokens}"
) )
return self._to_config(sizes) return self._to_config(sizes)
@@ -1152,6 +1237,8 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
def calculate_pool_sizes_from_max_tokens( def calculate_pool_sizes_from_max_tokens(
self, max_total_num_tokens: int, page_size: int self, max_total_num_tokens: int, page_size: int
) -> MemoryPoolConfig: ) -> MemoryPoolConfig:
# Token count, not a byte budget: the fixed pools are not re-subtracted, so
# the input must not exceed what calculate_pool_sizes derived for it.
assert page_size % 128 == 0, ( assert page_size % 128 == 0, (
"page_size must be multiple of 128 for compressed attention" "page_size must be multiple of 128 for compressed attention"
) )
@@ -7,7 +7,11 @@ import pytest
import torch import torch
import triton import triton
from sglang.kernels.ops.attention.dsv4 import compress_forward from sglang.kernels.ops.attention.dsv4 import (
CompressorDecodePlan,
CompressorPrefillPlan,
compress_forward,
)
from sglang.srt.utils import get_device from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kernels.deepseek_v4.common import ( from sglang.test.kernels.deepseek_v4.common import (
@@ -122,6 +126,91 @@ def _make_inputs(
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
@pytest.mark.parametrize("ring_size", [8, 16])
@pytest.mark.parametrize(
("gpu_inputs", "use_cuda_graph"),
[(False, False), (True, False), (True, True)],
)
def test_unified_request_ring_plans_ignore_full_to_state(
ring_size: int, gpu_inputs: bool, use_cuda_graph: bool
) -> None:
"""C4 plans must address state by request slot, not the full-cache map."""
device = torch.device(get_device())
req_pool_indices = torch.tensor([2, 5], dtype=torch.int64, device=device)
req_to_token = torch.zeros((6, 16), dtype=torch.int32, device=device)
full_to_state = torch.zeros(1, dtype=torch.int64, device=device)
seq_lens = torch.tensor([8, 12], dtype=torch.int64)
extend_lens = torch.tensor([4, 4], dtype=torch.int64)
if gpu_inputs:
seq_lens = seq_lens.to(device)
extend_lens = extend_lens.to(device)
prefill = CompressorPrefillPlan.generate(
compress_ratio=RATIO,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
extend_lens=extend_lens,
req_to_token=req_to_token,
full_to_state=full_to_state,
swa_page_size=256,
ring_size=ring_size,
num_q_tokens=8,
use_cuda_graph=use_cuda_graph,
use_req_ring=True,
)
plan_c = prefill.plan_c.view(torch.int32).reshape(-1, 4).cpu()
plan_w = prefill.plan_w.view(torch.int32).reshape(-1, 2).cpu()
valid_c = plan_c[plan_c[:, 2] >= 0]
got_reads = {
int(row[1].item()) & 0xFFFF: (int(row[2].item()), int(row[3].item()))
for row in valid_c
}
expected_reads = {
3: (
(2 * ring_size + 3 % ring_size) // RATIO,
(2 * ring_size + 7 % ring_size) // RATIO,
),
7: (
(5 * ring_size + 7 % ring_size) // RATIO,
(5 * ring_size + 11 % ring_size) // RATIO,
),
}
assert got_reads == expected_reads
valid_w = plan_w[plan_w[:, 1] >= 0]
got_writes = {int(row[0].item()): int(row[1].item()) for row in valid_w}
expected_writes = {
**{j: 2 * ring_size + (4 + j) % ring_size for j in range(4)},
**{4 + j: 5 * ring_size + (8 + j) % ring_size for j in range(4)},
}
assert got_writes == expected_writes
assert {got_writes[j] for j in range(4)}.isdisjoint(
{got_writes[j] for j in range(4, 8)}
)
decode = CompressorDecodePlan.generate(
compress_ratio=RATIO,
req_pool_indices=req_pool_indices,
req_to_token=req_to_token,
full_to_state=full_to_state,
seq_lens=torch.tensor([8, 12], dtype=torch.int64, device=device),
swa_page_size=256,
ring_size=ring_size,
use_req_ring=True,
)
got_decode = decode.plan_d.view(torch.int32).reshape(-1, 4).cpu()
expected_decode = torch.tensor(
[
[8, 2 * ring_size + 7 % ring_size, *expected_reads[3]],
[12, 5 * ring_size + 11 % ring_size, *expected_reads[7]],
],
dtype=torch.int32,
)
assert torch.equal(got_decode, expected_decode)
assert got_decode[0, 1] != got_decode[1, 1]
@pytest.mark.parametrize("mode", ["legacy", "paged"]) @pytest.mark.parametrize("mode", ["legacy", "paged"])
@pytest.mark.parametrize("seq_len", [4, 8, 32, 256, 1024]) @pytest.mark.parametrize("seq_len", [4, 8, 32, 256, 1024])
def test_prefill_no_context(mode: str, seq_len: int) -> None: def test_prefill_no_context(mode: str, seq_len: int) -> None:
@@ -0,0 +1,162 @@
"""CPU/mock tests for unified DSV4 C4 request-state lifecycle."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
import torch
from sglang.srt.disaggregation.decode import DecodeReqToTokenPool
from sglang.srt.mem_cache.allocation import alloc_req_slots
from sglang.srt.mem_cache.deepseek_v4_compress_state import KVAndScore
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _request(req_pool_idx=None, *, reused=False):
return SimpleNamespace(
kv=SimpleNamespace(
req_pool_idx=req_pool_idx,
kv_committed_len=1 if reused else 0,
kv_allocated_len=1 if reused else 0,
holds_kv=reused,
),
inflight_middle_chunks=1 if reused else 0,
)
def _mark_reused(req):
req.kv.kv_committed_len = 1
req.kv.kv_allocated_len = 1
req.kv.holds_kv = True
req.inflight_middle_chunks = 1
def _c4_pool(rows: int, width: int, ring_size: int):
return SimpleNamespace(
ratio=4,
ring_size=ring_size,
kv_score_buffer=KVAndScore(torch.full((rows, width), 7.0)),
)
def _token_pool(unified: bool, ring_size: int = 8):
logical_rows = 4 * ring_size
physical_rows = logical_rows + ring_size + 4
attn = _c4_pool(physical_rows, width=12, ring_size=ring_size)
indexer = _c4_pool(physical_rows, width=8, ring_size=ring_size)
c128 = SimpleNamespace(
ratio=128,
ring_size=128,
kv_score_buffer=KVAndScore(torch.full((physical_rows, 8), 9.0)),
)
token_pool = object.__new__(DeepSeekV4TokenToKVPool)
token_pool._unified_kv = unified
token_pool.compress_state_pools = [attn, c128]
token_pool.indexer_compress_state_pools = [indexer, None]
token_pool.get_ring_size = MagicMock(return_value=ring_size)
return token_pool, attn, indexer, c128, logical_rows
class TestUnifiedC4StateLifecycle(unittest.TestCase):
def test_pool_size_is_exact_request_ring_product(self):
configurator = object.__new__(DSV4PoolConfigurator)
configurator.disaggregation_mode = "decode"
configurator.disaggregation_decode_extra_slots = 3
configurator.c4_ring_size = 16
self.assertEqual(configurator._unified_c4_state_pool_size(10), 14 * 16)
def test_clear_resets_only_selected_request_rings(self):
ring_size = 8
token_pool, attn, indexer, c128, logical_rows = _token_pool(
unified=True, ring_size=ring_size
)
token_pool.clear_c4_req_states([1, 3])
selected = torch.tensor(list(range(8, 16)) + list(range(24, 32)))
untouched = torch.tensor(list(range(0, 8)) + list(range(16, 24)))
for pool in (attn, indexer):
state = pool.kv_score_buffer.kv_score
half = state.shape[-1] // 2
self.assertTrue(
torch.equal(
state[selected, :half], torch.zeros_like(state[selected, :half])
)
)
self.assertTrue(torch.isneginf(state[selected, half:]).all())
self.assertTrue((state[untouched] == 7).all())
self.assertTrue((state[logical_rows:] == 7).all())
self.assertTrue((c128.kv_score_buffer.kv_score == 9).all())
def test_clear_is_noop_off_the_unified_path(self):
"""The non-unified (fp8) pool addresses C4 state by SWA page, so a
req-slot reset must not touch it."""
token_pool, attn, indexer, _, _ = _token_pool(unified=False)
token_pool.clear_c4_req_states([1, 3])
for pool in (attn, indexer):
self.assertTrue((pool.kv_score_buffer.kv_score == 7).all())
def test_req_pool_hook_fires_for_new_slots_only(self):
req_pool = ReqToTokenPool(3, 16, "cpu", enable_memory_saver=False)
hook = MagicMock()
req_pool.register_on_alloc_rows(hook)
reused = _request()
# First admission: a brand-new slot, so its C4 ring must be cleared.
(reused_idx,) = alloc_req_slots(req_pool, [reused], None)
hook.assert_called_once_with([reused_idx])
# Chunked continuation reuses the same slot -- clearing it here would
# wipe the state captured by the previous chunk.
hook.reset_mock()
_mark_reused(reused)
self.assertEqual(alloc_req_slots(req_pool, [reused], None), [reused_idx])
hook.assert_not_called()
# Mixed batch: only the newly allocated slot is reported.
fresh = _request()
indices = alloc_req_slots(req_pool, [reused, fresh], None)
self.assertEqual(indices[0], reused_idx)
self.assertNotEqual(indices[1], reused_idx)
hook.assert_called_once_with([indices[1]])
def test_decode_req_pool_hook_fires_for_new_slots_only(self):
"""PD decode pre-allocates through DecodeReqToTokenPool, which has its
own alloc; it must report fresh rows the same way."""
req_pool = DecodeReqToTokenPool(
2, 16, "cpu", enable_memory_saver=False, pre_alloc_size=2
)
hook = MagicMock()
req_pool.register_on_alloc_rows(hook)
first = _request()
(first_idx,) = req_pool.alloc([first])
hook.assert_called_once_with([first_idx])
hook.reset_mock()
_mark_reused(first)
second = _request()
indices = req_pool.alloc([first, second])
self.assertEqual(indices[0], first_idx)
hook.assert_called_once_with([indices[1]])
hook.reset_mock()
self.assertEqual(req_pool.alloc([first]), [first_idx])
hook.assert_not_called()
def test_req_pool_without_hook_is_unchanged(self):
req_pool = ReqToTokenPool(2, 16, "cpu", enable_memory_saver=False)
(idx,) = alloc_req_slots(req_pool, [_request()], None)
self.assertGreater(idx, 0)
if __name__ == "__main__":
unittest.main()
@@ -36,6 +36,8 @@ def _make_self(*, page_size: int, full_available: int, swa_available: int):
return SimpleNamespace( return SimpleNamespace(
page_size=page_size, page_size=page_size,
# alloc_extend reads _swa_req_ring; pin the paged-SWA path.
_swa_req_ring=False,
full_attn_allocator=SimpleNamespace( full_attn_allocator=SimpleNamespace(
available_size=lambda: full_available, available_size=lambda: full_available,
alloc_extend=MagicMock(return_value=full_indices), alloc_extend=MagicMock(return_value=full_indices),
@@ -1044,9 +1044,9 @@ class TestSWAPoolFloor(CustomTestCase):
) )
self.assertEqual(config.swa_max_total_num_tokens, 3072) self.assertEqual(config.swa_max_total_num_tokens, 3072)
def _dsv4_sizes(self, max_tokens, page_size): def _dsv4_sizes(self, max_tokens, page_size, unified=False):
"""Exercise the DSV4 size arithmetic without a full V4 model fixture: """Exercise the DSV4 size arithmetic without a full V4 model fixture:
_compute_dsv4_sizes reads only these five attributes.""" _compute_dsv4_sizes reads only these six attributes."""
from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator
cfg = object.__new__(DSV4PoolConfigurator) cfg = object.__new__(DSV4PoolConfigurator)
@@ -1055,6 +1055,7 @@ class TestSWAPoolFloor(CustomTestCase):
cfg.swa_page_size = 128 cfg.swa_page_size = 128
cfg.c4_ring_size = 8 cfg.c4_ring_size = 8
cfg.c4_shrink_factor = 1 cfg.c4_shrink_factor = 1
cfg._unified = unified
return cfg._compute_dsv4_sizes(max_tokens, page_size) return cfg._compute_dsv4_sizes(max_tokens, page_size)
def test_dsv4_rejects_single_page_pool(self): def test_dsv4_rejects_single_page_pool(self):
@@ -1068,6 +1069,76 @@ class TestSWAPoolFloor(CustomTestCase):
sizes = self._dsv4_sizes(max_tokens=32768, page_size=256) sizes = self._dsv4_sizes(max_tokens=32768, page_size=256)
self.assertEqual(sizes.full_max_total_num_tokens, 32768) self.assertEqual(sizes.full_max_total_num_tokens, 32768)
self.assertEqual(sizes.swa_max_total_num_tokens, 3072) self.assertEqual(sizes.swa_max_total_num_tokens, 3072)
# Non-unified: the c4 state pool scales with the paged SWA pool.
self.assertEqual(sizes.c4_state_pool_size, 3072 // 128 * 8)
def test_dsv4_token_cap_never_grows_total_footprint(self):
"""Regression: the token-cap path subtracts no fixed-pool bias, so
capping the budget-derived token count must still shrink the total."""
cfg = self._dsv4_configurator_for_budget()
page_size = 128
budget = 256 * (1 << 30)
base = cfg.calculate_pool_sizes(budget, page_size)
base_bytes = self._dsv4_total_bytes(cfg, base.max_total_num_tokens)
self.assertLessEqual(base_bytes, budget)
for numerator in (999, 900, 500, 100, 1):
capped_tokens = (
base.max_total_num_tokens * numerator // 1000 // page_size * page_size
)
if capped_tokens <= 0:
continue
capped = cfg.calculate_pool_sizes_from_max_tokens(capped_tokens, page_size)
capped_bytes = self._dsv4_total_bytes(cfg, capped.max_total_num_tokens)
with self.subTest(numerator=numerator):
self.assertLessEqual(capped_bytes, base_bytes)
# White-box 671B-class shape: the byte arithmetic runs without a model fixture.
def _dsv4_configurator_for_budget(self):
from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator
cfg = object.__new__(DSV4PoolConfigurator)
cfg.qk_nope_head_dim, cfg.qk_rope_head_dim = 128, 64
cfg.attn_head_dim = 192
cfg.indexer_head_dim = 128
cfg.num_layers_total = 61
cfg.num_layers_ca4 = 61
cfg.num_layers_ca128 = 61
cfg.c4_ring_size = 8
cfg.c128_ring_size = 128
cfg._swa_ring_size = 128
cfg._spec_infl = 1.0
cfg.context_len = 65536
cfg.bytes_per_full_token = 576.0
cfg.requested_max_running_requests_per_worker = None
cfg.swa_ratio = 0.1
cfg.sliding_window_size = 4096
cfg.swa_page_size = 128
cfg.c4_shrink_factor = 1
cfg.online_c128_mtp_max_draft_tokens = 0
cfg.disaggregation_mode = None
cfg.disaggregation_decode_extra_slots = 0
cfg._unified = True
return cfg
# Token pool plus the three request-scoped fixed pools, sized from the
# concurrency resolve_max_num_reqs derives from this token count.
def _dsv4_total_bytes(self, cfg, tokens):
estimated = max(min(int(tokens / cfg.context_len * 512), 4096), 2048)
max_running_requests = min(estimated, tokens // 2)
return int(
tokens * cfg.bytes_per_full_token
+ cfg._fixed_swa_bytes(max_running_requests)
+ cfg._fixed_c4_state_bytes(max_running_requests)
+ cfg._get_c128_state_fixed_bytes(max_running_requests)
)
def test_dsv4_unified_c4_state_not_token_scaled(self):
# Unified-KV sizes the c4 state ring from max_running_requests in
# finalize_with_max_running_requests, so it must not scale here.
sizes = self._dsv4_sizes(max_tokens=32768, page_size=256, unified=True)
self.assertEqual(sizes.full_max_total_num_tokens, 32768)
self.assertEqual(sizes.swa_max_total_num_tokens, 3072)
self.assertEqual(sizes.c4_state_pool_size, 0)
if __name__ == "__main__": if __name__ == "__main__":