Revert "[AMD][DSV4] Fix unified-KV pool sizing and SWA ring accounting" (#38163)

This commit is contained in:
Liangsheng Yin
2026-09-05 17:28:46 -07:00
committed by GitHub
parent 09daea94ac
commit f5819b09bf
23 changed files with 159 additions and 791 deletions
@@ -28,14 +28,12 @@ using R2T_T = int32_t;
using F2S_T = int64_t; using F2S_T = int64_t;
using IDX_T = int64_t; using IDX_T = int64_t;
/// NOTE: for the internal use, we pack the ragged and batch id, since both not /// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536
/// exceed 65536
SGL_DEVICE __host__ PlanW pack_w(uint32_t ragged_id, uint32_t batch_id, int32_t seq_len) { SGL_DEVICE __host__ PlanW pack_w(uint32_t ragged_id, uint32_t batch_id, int32_t seq_len) {
return {static_cast<uint32_t>(ragged_id | batch_id << 16), seq_len}; return {static_cast<uint32_t>(ragged_id | batch_id << 16), seq_len};
} }
/// NOTE: for the internal use, we pack the ragged and batch id, since both not /// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536
/// exceed 65536
SGL_DEVICE uint2 unpack_w(PlanW plan) { SGL_DEVICE uint2 unpack_w(PlanW plan) {
return {static_cast<uint16_t>(plan.ragged_id), static_cast<uint16_t>(plan.ragged_id >> 16)}; return {static_cast<uint16_t>(plan.ragged_id), static_cast<uint16_t>(plan.ragged_id >> 16)};
} }
@@ -49,11 +47,9 @@ struct Prefill0Params {
uint32_t num_q_tokens; uint32_t num_q_tokens;
int32_t compress_ratio; int32_t compress_ratio;
int32_t swa_page_size; int32_t swa_page_size;
/// \brief Trailing tokens the write plan keeps resident in the compress state /// \brief Trailing tokens the write plan keeps resident in the compress state ring.
/// ring. Derived from the ring in `plan_compress_prefill`; see the bound /// Derived from the ring in `plan_compress_prefill`; see the bound there.
/// there.
int32_t mtp_pad; int32_t mtp_pad;
bool use_req_ring;
}; };
struct Prefill1Params { struct Prefill1Params {
@@ -71,7 +67,6 @@ 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 {
@@ -85,7 +80,6 @@ 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 {
@@ -161,8 +155,7 @@ __global__ __launch_bounds__(1024, 1) //
counter_w = 0; counter_w = 0;
} }
// === Stage B: min/max(extend_len) for MTP-uniform detection === // === Stage B: min/max(extend_len) for MTP-uniform detection ===
// For min, treat threads outside `batch_size` as +inf so they don't pull the // For min, treat threads outside `batch_size` as +inf so they don't pull the min down.
// min down.
const uint32_t e_for_max = static_cast<uint32_t>(extend_len); const uint32_t e_for_max = static_cast<uint32_t>(extend_len);
const uint32_t e_for_min = (tx < params.batch_size) ? e_for_max : 0xFFFFFFFFu; const uint32_t e_for_min = (tx < params.batch_size) ? e_for_max : 0xFFFFFFFFu;
warp_max[warp_id] = warp::reduce_max(e_for_max); warp_max[warp_id] = warp::reduce_max(e_for_max);
@@ -175,19 +168,17 @@ __global__ __launch_bounds__(1024, 1) //
__syncthreads(); __syncthreads();
const auto num_q = params.num_q_tokens; const auto num_q = params.num_q_tokens;
// MTP-uniform: every batch shares the same small extend_len `E`, so we can // MTP-uniform: every batch shares the same small extend_len `E`, so we can decompose
// decompose a global token id `k` into (batch_id, j) = (k / E, k % E) and // a global token id `k` into (batch_id, j) = (k / E, k % E) and skip the per-batch loop.
// skip the per-batch loop.
const bool is_mtp_extend = (s_min_extend == s_max_extend) && (s_max_extend > 0) && (s_max_extend <= 32); const bool is_mtp_extend = (s_min_extend == s_max_extend) && (s_max_extend > 0) && (s_max_extend <= 32);
// === Stage C: emit valid plans, slot allocation via shared-mem atomicAdd === // === Stage C: emit valid plans, slot allocation via shared-mem atomicAdd ===
if (is_mtp_extend) { if (is_mtp_extend) {
// Path 1: token-driven. Each global token id maps to exactly one (batch_id, // Path 1: token-driven. Each global token id maps to exactly one (batch_id, j).
// j).
const uint32_t E = s_max_extend; const uint32_t E = s_max_extend;
// num_q is the padded buffer size (graph bucket), not the work size: cap // num_q is the padded buffer size (graph bucket), not the work size: cap the
// the loop at the real token count so batch_id = k / E stays < batch_size // loop at the real token count so batch_id = k / E stays < batch_size on an
// on an underfilled replay; Stage D pads [counter, num_q) with invalid. // underfilled replay; Stage D pads [counter, num_q) with invalid.
const uint32_t num_real_q = params.batch_size * E; const uint32_t num_real_q = params.batch_size * E;
for (uint32_t k = tx; k < num_real_q; k += block_size) { for (uint32_t k = tx; k < num_real_q; k += block_size) {
const uint32_t batch_id = k / E; const uint32_t batch_id = k / E;
@@ -212,15 +203,15 @@ __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 && !params.use_req_ring) do_write = (position % sps) >= (sps - cr); if (!do_write && is_overlap) 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);
} }
} }
} else { } else {
// Path 2: general prefill (long extend_len). Iterate batches in an outer // Path 2: general prefill (long extend_len). Iterate batches in an outer loop;
// loop; the whole block sweeps each batch's tokens in parallel. // the whole block sweeps each batch's tokens in parallel.
uint32_t base_e = 0; uint32_t base_e = 0;
for (uint32_t batch_id = 0; batch_id < params.batch_size; ++batch_id) { for (uint32_t batch_id = 0; batch_id < params.batch_size; ++batch_id) {
const int32_t pl = s_prefix_len[batch_id]; const int32_t pl = s_prefix_len[batch_id];
@@ -245,7 +236,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 && !params.use_req_ring) do_write = (position % sps) >= (sps - cr); if (!do_write && is_overlap) 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);
@@ -279,7 +270,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_req_ring_loc = [&](int64_t rid, int32_t position) { const auto compute_c128_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);
}; };
@@ -292,9 +283,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 || params.use_req_ring) { if (params.compress_ratio == 128) {
plan_c.read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio; plan_c.read_page_0 = compute_c128_loc(rid, position_0) / 128;
plan_c.read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio; plan_c.read_page_1 = compute_c128_loc(rid, position_1) / 128;
} 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];
@@ -316,8 +307,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 || params.use_req_ring) { if (params.compress_ratio == 128) {
plan_w.write_loc = compute_req_ring_loc(rid, position); plan_w.write_loc = compute_c128_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]);
@@ -338,7 +329,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_req_ring_loc = [&](int64_t rid, int32_t position) { const auto compute_c128_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]);
@@ -347,10 +338,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 || params.use_req_ring) { if (params.compress_ratio == 128) {
write_loc = compute_req_ring_loc(rid, position_1); write_loc = compute_c128_loc(rid, position_1);
read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio; read_page_0 = compute_c128_loc(rid, position_0) / 128;
read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio; read_page_1 = compute_c128_loc(rid, position_1) / 128;
} 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];
@@ -375,10 +366,8 @@ __global__ void plan_compress_prefill_legacy_kernel(const Prefill1ParamsLegacy p
auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid(); auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid();
/// Per-request ring buffer slot translation: /// Per-request ring buffer slot translation:
/// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4
/// 4 /// - c128: page = rid; slot = rid * 128 + position % 128
/// - c128: page = rid; slot = rid * 128 + position %
/// 128
const auto legacy_compute_page = [&](int32_t rid, int32_t position) { const auto legacy_compute_page = [&](int32_t rid, int32_t position) {
if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1);
return rid; // c128 return rid; // c128
@@ -404,8 +393,7 @@ __global__ void plan_compress_prefill_legacy_kernel(const Prefill1ParamsLegacy p
if (!plan_w.is_invalid()) { if (!plan_w.is_invalid()) {
const auto [ragged_id, batch_id] = unpack_w(plan_w); const auto [ragged_id, batch_id] = unpack_w(plan_w);
const auto rid = static_cast<int32_t>(params.rid_ptr[batch_id]); const auto rid = static_cast<int32_t>(params.rid_ptr[batch_id]);
// `write_loc` carries (position + 1) at this stage; may not be // `write_loc` carries (position + 1) at this stage; may not be ratio-aligned
// ratio-aligned
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;
plan_w.write_loc = legacy_compute_loc(rid, position); plan_w.write_loc = legacy_compute_loc(rid, position);
@@ -419,10 +407,8 @@ __global__ void plan_compress_decode_legacy_kernel(const DecodeParamsLegacy para
const auto idx = blockIdx.x * blockDim.x + threadIdx.x; const auto idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= params.batch_size) return; if (idx >= params.batch_size) return;
/// Per-request ring buffer slot translation: /// Per-request ring buffer slot translation:
/// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4
/// 4 /// - c128: page = rid; slot = rid * 128 + position % 128
/// - c128: page = rid; slot = rid * 128 + position %
/// 128
const auto legacy_compute_page = [&](int32_t rid, int32_t position) { const auto legacy_compute_page = [&](int32_t rid, int32_t position) {
if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1);
return rid; // c128 return rid; // c128
@@ -461,8 +447,7 @@ using PrefillPlan = tvm::ffi::Tuple<tvm::ffi::Tensor, tvm::ffi::Tensor>;
* @param compress_plan `[num_q_tokens, 16]` uint8 (output) * @param compress_plan `[num_q_tokens, 16]` uint8 (output)
* @param write_plan `[num_q_tokens, 8]` uint8 (output) * @param write_plan `[num_q_tokens, 8]` uint8 (output)
* @param compress_ratio 4 for c4, 128 for c128 * @param compress_ratio 4 for c4, 128 for c128
* @param use_cuda_graph Whether the plans will be used with cuda graph (affects * @param use_cuda_graph Whether the plans will be used with cuda graph (affects padding)
* padding)
* @return (compress plan tensor, write plan tensor) * @return (compress plan tensor, write plan tensor)
*/ */
inline PrefillPlan plan_compress_prefill( inline PrefillPlan plan_compress_prefill(
@@ -476,7 +461,6 @@ 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"};
@@ -519,29 +503,27 @@ 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);
// Write pad: trailing tokens kept resident so a verify batch's committed tail // Write pad: trailing tokens kept resident so a verify batch's committed tail survives
// survives any accept length. Zero without speculation -- nothing rolls back, // any accept length. Zero without speculation -- nothing rolls back, and the ring is
// and the ring is then exactly one window wide. Otherwise the ring bounds it: // then exactly one window wide. Otherwise the ring bounds it: a write at `w` aliases
// a write at `w` aliases onto `w - ring_size`, and the earliest position a // onto `w - ring_size`, and the earliest position a future compression still needs is
// future compression still needs is `prefix_len - window_size + 2` (the next // `prefix_len - window_size + 2` (the next batch commits >= 1 token, and `run_prefill`
// batch commits >= 1 token, and `run_prefill` launches the compress kernel // launches the compress kernel before the write kernel, so a batch's own compressions
// before the write kernel, so a batch's own compressions read the pre-write // read the pre-write ring). Padding past the extend range is harmless: the loops only
// ring). Padding past the extend range is harmless: the loops only span // span `[prefix_len, seq_len)`.
// `[prefix_len, seq_len)`.
const auto mtp_pad = ring_size > window_size ? ring_size - window_size + 2 : 0; const auto mtp_pad = ring_size > window_size ? ring_size - window_size + 2 : 0;
const auto device = device_.unwrap(); const auto device = device_.unwrap();
const auto stream = LaunchKernel::resolve_device(device); const auto stream = LaunchKernel::resolve_device(device);
if (cpu_or_gpu.unwrap().device_type == kDLGPU) { if (cpu_or_gpu.unwrap().device_type == kDLGPU) {
// GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata // GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata directly
// directly on device, padding to num_q_tokens with invalid; kernel_1 then // on device, padding to num_q_tokens with invalid; kernel_1 then finalizes the
// finalizes the SWA-translated read/write locations. Used for MTP / // SWA-translated read/write locations. Used for MTP / cuda-graph capture where
// cuda-graph capture where a host sync would be expensive. // a host sync would be expensive.
RuntimeCheck(batch_size <= kMaxPrefillBatchSize, "GPU plan only support batch size up to ", kMaxPrefillBatchSize); RuntimeCheck(batch_size <= kMaxPrefillBatchSize, "GPU plan only support batch size up to ", kMaxPrefillBatchSize);
auto C = ffi::empty({num_q_tokens, sizeof(PlanC)}, kDLUInt8, device); auto C = ffi::empty({num_q_tokens, sizeof(PlanC)}, kDLUInt8, device);
auto W = ffi::empty({num_q_tokens, sizeof(PlanW)}, kDLUInt8, device); auto W = ffi::empty({num_q_tokens, sizeof(PlanW)}, kDLUInt8, device);
@@ -555,11 +537,9 @@ 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 // kernel_1 sees the already-padded buffers, so num_c == num_w == num_padded == num_q_tokens.
// == num_q_tokens.
const auto params1 = Prefill1Params{ const auto params1 = Prefill1Params{
.plan_c = static_cast<PlanC*>(C.data_ptr()), .plan_c = static_cast<PlanC*>(C.data_ptr()),
.plan_w = static_cast<PlanW*>(W.data_ptr()), .plan_w = static_cast<PlanW*>(W.data_ptr()),
@@ -575,7 +555,6 @@ 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);
@@ -603,7 +582,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 && !use_req_ring && position % swa_page_size >= (swa_page_size - compress_ratio); return is_overlap && 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;
@@ -652,7 +631,6 @@ 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);
@@ -667,8 +645,7 @@ 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>();
@@ -690,7 +667,6 @@ 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);
@@ -705,7 +681,6 @@ 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,7 +100,6 @@ 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,
@@ -136,8 +135,6 @@ def create_paged_compress_data_kernel(
pos = tl.maximum(pos, 0) pos = tl.maximum(pos, 0)
if compress_ratio == 128: if compress_ratio == 128:
state_loc = rid * ring_size + (pos % ring_size) state_loc = rid * ring_size + (pos % ring_size)
elif use_req_ring:
state_loc = rid * ring_size + (pos % ring_size)
else: else:
loc = tl.load( loc = tl.load(
req_to_token_ptr req_to_token_ptr
@@ -185,7 +182,6 @@ 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]
@@ -209,7 +205,6 @@ 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,7 +162,6 @@ 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
@@ -170,7 +169,7 @@ class CompressorDecodePlan(NamedTuple):
module = _jit_compress_plan_module() module = _jit_compress_plan_module()
fn = module.plan_decode fn = module.plan_decode
args = ( plan_d = fn(
req_pool_indices, req_pool_indices,
req_to_token, req_to_token,
full_to_state, full_to_state,
@@ -179,7 +178,6 @@ class CompressorDecodePlan(NamedTuple):
int(swa_page_size), int(swa_page_size),
int(ring_size), int(ring_size),
) )
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
@@ -249,7 +247,6 @@ 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(
@@ -277,7 +274,7 @@ class CompressorPrefillPlan(NamedTuple):
module = _jit_compress_plan_module() module = _jit_compress_plan_module()
fn = module.plan_prefill fn = module.plan_prefill
args = ( plan_c, plan_w = fn(
req_pool_indices, req_pool_indices,
req_to_token, req_to_token,
full_to_state, full_to_state,
@@ -288,11 +285,7 @@ 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),
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,
@@ -1768,18 +1768,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if total_prefix_len is None: if total_prefix_len is None:
total_prefix_len = prefix_len total_prefix_len = prefix_len
is_new_req_slot = req.kv.req_pool_idx is None
req_pool_indices = self.req_to_token_pool.alloc([req]) req_pool_indices = self.req_to_token_pool.alloc([req])
assert req_pool_indices is not None, ( assert req_pool_indices is not None, (
"req_pool_indices is full! There is a bug in memory estimation." "req_pool_indices is full! There is a bug in memory estimation."
) )
if is_new_req_slot:
clear_c4_req_states = getattr(
self.token_to_kv_pool, "clear_c4_req_states", None
)
if clear_c4_req_states is not None:
clear_c4_req_states(req_pool_indices)
fill_len = self._pre_alloc_fill_len(req) fill_len = self._pre_alloc_fill_len(req)
req.kv.kv_committed_len = fill_len req.kv.kv_committed_len = fill_len
@@ -144,9 +144,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 or ( if self.ratio == 128:
self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False)
):
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
) )
@@ -168,9 +166,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 or ( if self.ratio == 128:
self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False)
):
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
) )
@@ -275,9 +271,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 or ( if self.ratio == 128:
self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False)
):
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
) )
@@ -292,9 +286,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 or ( if self.ratio == 128:
self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False)
):
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
@@ -26,7 +26,9 @@ from sglang.srt.layers.utils.cp_utils import (
cp_all_gather_rerange_finish, cp_all_gather_rerange_finish,
cp_all_gather_rerange_launch, cp_all_gather_rerange_launch,
) )
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.deepseek_v4_compress_state import (
CompressStatePool,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_v2 import _is_hip from sglang.srt.models.deepseek_v2 import _is_hip
@@ -262,7 +264,6 @@ 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:
@@ -272,8 +273,6 @@ def create_paged_compressor_data(
positions = positions.masked_fill(positions < 0, 0) positions = positions.masked_fill(positions < 0, 0)
if compress_ratio == 128: if compress_ratio == 128:
state_loc = req_pool_indices * ring_size + positions % ring_size state_loc = req_pool_indices * ring_size + positions % ring_size
elif use_req_ring:
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]
swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(loc) swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(loc)
@@ -295,7 +294,6 @@ 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,7 +441,6 @@ 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()
@@ -468,7 +467,6 @@ 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,
) )
@@ -481,7 +479,6 @@ 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,
) )
+5 -47
View File
@@ -57,7 +57,6 @@ import dataclasses
import logging import logging
import re import re
import sys import sys
import time
from array import array from array import array
from concurrent.futures import Future from concurrent.futures import Future
from enum import Enum, auto from enum import Enum, auto
@@ -101,7 +100,10 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import ( from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
NewTokenRatioTracker, NewTokenRatioTracker,
) )
from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend from sglang.srt.mem_cache.allocation import (
alloc_for_decode,
alloc_for_extend,
)
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
@@ -159,9 +161,6 @@ _MM_HASH_MASK = (1 << 64) - 1
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Throttle for the unified-KV SWA bottleneck diagnostic (seconds).
_last_swa_bottleneck_log = 0.0
ReturnHiddenStatesMode = Union[bool, Literal["last"]] ReturnHiddenStatesMode = Union[bool, Literal["last"]]
@@ -3078,50 +3077,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
shortfalls retract gracefully instead of tripping fail-loud alloc shortfalls retract gracefully instead of tripping fail-loud alloc
errors.""" errors."""
num_tokens = self.new_tokens_required_next_decode(selected_indices) num_tokens = self.new_tokens_required_next_decode(selected_indices)
allocator = self.token_to_kv_pool_allocator return self.token_to_kv_pool_allocator.check_decode_capacity(
ok = allocator.check_decode_capacity(
num_tokens=num_tokens, tree_cache=self.tree_cache num_tokens=num_tokens, tree_cache=self.tree_cache
) )
if not ok and getattr(allocator.get_kvcache(), "_unified_kv", False):
self._log_unified_swa_bottleneck(allocator, num_tokens, selected_indices)
return ok
def _log_unified_swa_bottleneck(self, allocator, num_tokens, selected_indices):
"""Diagnostic (unified-KV only): when check_decode_mem is short, compare
the SWA token bookkeeping against the real per-slot ring utilization.
Throttled to avoid log floods during retract storms."""
global _last_swa_bottleneck_log
now = time.monotonic()
if now - _last_swa_bottleneck_log < 1.0:
return
_last_swa_bottleneck_log = now
try:
full_avail = allocator.full_available_size()
swa_avail = allocator.swa_available_size()
reqs = (
self.reqs
if selected_indices is None
else [self.reqs[i] for i in selected_indices]
)
active_slots = len(
{int(r.kv.req_pool_idx) for r in reqs if r.kv.req_pool_idx is not None}
)
unified = getattr(allocator.get_kvcache(), "unified_kv_pool", None)
if unified is not None:
num_slots = unified.num_slots
ring_util = (
active_slots * unified.swa_ring_size / max(unified.swa_pages, 1)
)
else:
num_slots, ring_util = -1, -1.0
logger.warning(
"[SWA-BOTTLENECK] check_decode_mem short: "
f"need={num_tokens}, full_avail={full_avail}, swa_avail={swa_avail}, "
f"active_slots={active_slots}/{num_slots}, "
f"ring_util_upper={ring_util:.4f}"
)
except Exception as e: # diagnostics must never break scheduling
logger.warning(f"[SWA-BOTTLENECK] logging failed: {e}")
def retract_decode(self) -> Tuple[List[Req], float, List[Req]]: def retract_decode(self) -> Tuple[List[Req], float, List[Req]]:
"""Retract the decoding requests when there is not enough memory.""" """Retract the decoding requests when there is not enough memory."""
+9 -39
View File
@@ -5,7 +5,10 @@ from array import array
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
from sglang.srt.runtime_context import get_disagg, get_schedule from sglang.srt.runtime_context import (
get_disagg,
get_schedule,
)
from sglang.srt.utils import get_bool_env_var, is_hip from sglang.srt.utils import get_bool_env_var, is_hip
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG") _ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
@@ -660,16 +663,8 @@ class PrefillAdder:
@property @property
def rem_swa_tokens(self): def rem_swa_tokens(self):
allocator = self.token_to_kv_pool_allocator
if getattr(allocator.get_kvcache(), "_unified_kv", False):
# Unified-KV: SWA is a per-request ring, not a tree-reusable token
# pool. swa_available_size() already reports ring capacity
# (free_slots * ring_cost). tree swa_evictable is in the old linear
# token unit and freeing it does not release ring space, so exclude
# it here to keep a single consistent accounting unit.
return allocator.swa_available_size() - self.rem_swa_token_offset
return ( return (
allocator.swa_available_size() self.token_to_kv_pool_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
) )
@@ -716,13 +711,6 @@ 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 getattr(allocator.get_kvcache(), "_unified_kv", False):
# Unified-KV: each request occupies exactly one fixed SWA ring slot,
# independent of context / chunk length; a host-hit prefix reuses the
# same ring. Budget the fixed per-slot ring cost (paired with the
# ring-based swa_available_size on the allocator).
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:
@@ -850,9 +838,6 @@ 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,
host_hit_len: int = 0,
storage_hit_len: 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)
@@ -876,14 +861,6 @@ 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:
# Unified-KV: SWA is a fixed per-request ring slot reserved once at
# first admission and already reflected in swa_available_size() on
# later rounds. Charging it again on a chunked continuation would
# double-count the slot and over-throttle admission, so skip it.
_unified = getattr(
self.token_to_kv_pool_allocator.get_kvcache(), "_unified_kv", False
)
if not (_unified 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
) )
@@ -1021,15 +998,9 @@ 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 and not getattr( if self.is_hybrid_swa:
self.token_to_kv_pool_allocator.get_kvcache(), "_unified_kv", False
):
# 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
# Unified-KV: rem_swa_tokens is ring capacity (free_slots * ring
# cost), not a linear per-chunk token budget, and this request's
# ring slot is already reserved -- mixing units here would wrongly
# truncate the chunk, so skip the SWA clamp.
_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
) )
@@ -1068,7 +1039,6 @@ 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
@@ -1278,7 +1248,7 @@ 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 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),
@@ -1314,7 +1284,7 @@ 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 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),
@@ -3,7 +3,14 @@ from __future__ import annotations
import logging import logging
from collections import deque from collections import deque
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Callable, Deque, List, Optional, Tuple from typing import (
TYPE_CHECKING,
Callable,
Deque,
List,
Optional,
Tuple,
)
import torch import torch
@@ -25,7 +32,10 @@ from sglang.srt.observability.scheduler_stage_metrics import (
scheduler_stage_method, scheduler_stage_method,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import ceil_align, raise_error_or_warn from sglang.srt.utils.common import (
ceil_align,
raise_error_or_warn,
)
from sglang.srt.utils.watchdog import WatchdogRaw from sglang.srt.utils.watchdog import WatchdogRaw
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -142,23 +152,6 @@ 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
kv = allocator.get_kvcache()
if getattr(kv, "_unified_kv", False):
# Unified-KV DSV4: SWA is a fixed per-request ring, reused per request
# and released together with the req_pool slot (which has its own
# leak check). swa_available_size() is deliberately non-binding (it
# always reports the full ring so it never throttles admission), and
# cached radix prefixes still report swa_evictable even though the
# completed request already freed its ring slot. The token-pool
# invariant (available + evictable + protected + session == total)
# therefore does not model this pool -- skip it to avoid a spurious
# leak. Ring-slot leaks are still caught by the req_to_token check.
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
@@ -394,10 +387,7 @@ class SchedulerInvariantChecker:
# Sub-allocators to check: a flat allocator is its own single sub; a # Sub-allocators to check: a flat allocator is its own single sub; a
# hybrid-SWA wrapper exposes full_attn_allocator + swa_attn_allocator. # hybrid-SWA wrapper exposes full_attn_allocator + swa_attn_allocator.
# DSV4-HiSparse nests the real SWA allocator under logical_attn_allocator,
# so unwrap first (no-op for a plain/flat allocator).
alloc = self.token_to_kv_pool_allocator alloc = self.token_to_kv_pool_allocator
alloc = getattr(alloc, "logical_attn_allocator", alloc)
sub_allocs = ( sub_allocs = (
[alloc] [alloc]
if getattr(alloc, "free_pages", None) is not None if getattr(alloc, "free_pages", None) is not None
@@ -2,7 +2,14 @@ from __future__ import annotations
import dataclasses import dataclasses
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple from typing import (
TYPE_CHECKING,
Any,
Callable,
List,
Optional,
Tuple,
)
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
UnifiedMambaSWATokenToKVPoolAllocator, UnifiedMambaSWATokenToKVPoolAllocator,
@@ -294,16 +301,6 @@ 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()
# Unified-KV DSV4: SWA is a fixed per-request ring, released with the
# req_pool slot. Cached radix prefixes still report swa_evictable even
# though the completed request already freed its ring slot, and
# swa_available_size() is non-binding (always the full ring). Counting
# that evictable here would double-count against the ring and drive
# swa_num_used / swa_token_usage negative. The ring holds nothing
# evictable, so zero it out to keep the usage stats coherent.
_swa_kv = self.token_to_kv_pool_allocator.get_kvcache()
if getattr(_swa_kv, "_unified_kv", False):
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
) )
+1 -13
View File
@@ -230,7 +230,6 @@ def alloc_req_slots(
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
reqs: list[Req], reqs: list[Req],
tree_cache: BasePrefixCache | None, tree_cache: BasePrefixCache | None,
token_to_kv_pool=None,
) -> list[int]: ) -> list[int]:
"""Allocate request slots from the pool. """Allocate request slots from the pool.
@@ -261,7 +260,6 @@ def alloc_req_slots(
tree_cache.evict_for_alloc( tree_cache.evict_for_alloc(
EvictParams(num_tokens=0, mamba_num=mamba_num) EvictParams(num_tokens=0, mamba_num=mamba_num)
) )
newly_allocated = [req.kv.req_pool_idx is None for req in reqs]
req_pool_indices = req_to_token_pool.alloc(reqs) req_pool_indices = req_to_token_pool.alloc(reqs)
if req_pool_indices is None: if req_pool_indices is None:
raise RuntimeError( raise RuntimeError(
@@ -269,13 +267,6 @@ def alloc_req_slots(
"Please set a smaller number for `--max-running-requests`. " "Please set a smaller number for `--max-running-requests`. "
f"{req_to_token_pool.available_size()=}, {num_reqs=}, " f"{req_to_token_pool.available_size()=}, {num_reqs=}, "
) )
new_req_pool_indices = [
idx for idx, is_new in zip(req_pool_indices, newly_allocated) if is_new
]
clear_c4_req_states = getattr(token_to_kv_pool, "clear_c4_req_states", None)
if new_req_pool_indices and clear_c4_req_states is not None:
clear_c4_req_states(new_req_pool_indices)
return req_pool_indices return req_pool_indices
@@ -320,10 +311,7 @@ def alloc_for_extend(
# Allocate req slots (raises RuntimeError if the pool is exhausted) # Allocate req slots (raises RuntimeError if the pool is exhausted)
req_pool_indices = alloc_req_slots( req_pool_indices = alloc_req_slots(
batch.req_to_token_pool, batch.req_to_token_pool, batch.reqs, batch.tree_cache
batch.reqs,
batch.tree_cache,
token_to_kv_pool=batch.token_to_kv_pool_allocator.get_kvcache(),
) )
req_pool_indices_cpu = torch.tensor( req_pool_indices_cpu = torch.tensor(
req_pool_indices, dtype=torch.int64, pin_memory=pin_memory req_pool_indices, dtype=torch.int64, pin_memory=pin_memory
@@ -343,10 +343,6 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def get_kvcache(self): def get_kvcache(self):
return self._kvcache return self._kvcache
@property
def swa_ring_cost_tokens(self) -> int:
return self.logical_attn_allocator.swa_ring_cost_tokens
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices) return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices)
+3 -75
View File
@@ -1,5 +1,3 @@
import logging
import torch import torch
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
@@ -10,8 +8,6 @@ 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:
@@ -41,7 +37,6 @@ 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
@@ -109,46 +104,10 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.swa_free_group = [] self.swa_free_group = []
self._kvcache = kvcache self._kvcache = kvcache
# Unified-KV (DSV4): SWA is a per-request ring addressed by state_slot
# (== req_pool_idx) + position inside the DSV4 kernels. The paged SWA
# indices / full_to_swa_index_mapping produced here are NOT consumed on
# that path, so treating SWA as a linearly-consumed token pool
# over-throttles admission and decode retract. Instead account for it as
# a fixed per-request ring slot; the real bound is concurrency
# (num_req_slots), already enforced by req_to_token_pool /
# max_running_requests.
self._unified = getattr(kvcache, "_unified_kv", False)
self._req_to_token_pool = req_to_token_pool
if self._unified:
ring_size = getattr(kvcache, "unified_swa_ring_size", self.page_size)
self._swa_ring_cost = (
(ring_size + self.page_size - 1) // self.page_size
) * self.page_size
logger.info(
"[SWA-BOOKKEEPING] unified ring accounting enabled: "
f"num_slots={getattr(kvcache, 'num_req_slots', '?')}, "
f"swa_ring_size={ring_size}, "
f"ring_cost_tokens={self._swa_ring_cost}, "
f"unified_swa_pages={getattr(kvcache, 'unified_swa_pages', '?')} | "
f"legacy paged size_swa={self._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_ring_cost_tokens(self) -> int:
"""Unified: paged SWA cost of one request's ring slot (0 otherwise)."""
return self._swa_ring_cost
def available_size(self): def available_size(self):
if self._unified:
# 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(),
@@ -158,12 +117,6 @@ 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._unified:
# Ring-based availability: free request slots * per-slot ring cost.
# Fall back to non-binding if the req pool wasn't wired in.
if self._req_to_token_pool is None:
return self.full_attn_allocator.available_size()
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
@@ -218,15 +171,11 @@ 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:
full_ok = ( return (
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
if self._unified: <= self.swa_attn_allocator.available_size() // self.page_size
# 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(
@@ -246,20 +195,6 @@ 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._unified:
# Unified SWA ring is slot-addressed and not paged here: allocate only
# the full-attention KV and skip the vestigial SWA allocator / mapping
# (unused by the DSV4 kernels).
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(
@@ -356,13 +291,6 @@ 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._unified:
# See alloc_extend: unified SWA ring is slot-addressed, allocate full
# only and skip the vestigial SWA allocator / mapping.
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(
@@ -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, Sequence, Tuple from typing import List, Literal, NamedTuple, Optional, Tuple
import torch import torch
@@ -564,11 +564,6 @@ 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
# Keep the legacy SWA-addressed pool large enough on non-unified paths.
# Unified request-addressed sizing is set exactly after resolving the
# unified-kv gate below.
c4_ring_size = self.get_ring_size(4)
c4_state_pool_size = max(c4_state_pool_size, self.num_req_slots * c4_ring_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:
@@ -629,9 +624,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
) )
self._unified_kv = is_unified_kv_triton() self._unified_kv = is_unified_kv_triton()
if self._unified_kv:
# Unified C4 state is request-scoped: no SWA-derived over-allocation.
self.c4_state_pool_size = self.num_req_slots * c4_ring_size
if self._unified_kv: if self._unified_kv:
self.swa_kv_pool = None self.swa_kv_pool = None
@@ -1058,37 +1050,6 @@ 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:
"""Reset newly allocated unified C4 attention and indexer state rings.
Only the request-owned rows are touched. The extra sentinel/ring padding
allocated by :class:`CompressStatePool` remains intact.
"""
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:
@@ -1115,13 +1076,7 @@ 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 equivalent cleanup: draft states are written in position order,
and every rejected position is overwritten before it can become the prior
state of a later accepted token. C128 cleanup is required because its
compression boundary can consume a previously written draft slot directly.
"""
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
@@ -204,7 +204,9 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import (
SpecAuxHiddenStateConfig, SpecAuxHiddenStateConfig,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import (
MemoryPoolConfig,
)
class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True): class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True):
@@ -334,35 +336,6 @@ 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
# Unified-KV DSV4: SWA is a fixed per-request ring, so the allocator
# reports ring capacity (free_req_slots * ring_cost) from
# swa_available_size(), while swa_max_total_num_tokens was sized from
# the (vestigial, unallocated) full_token-scaled SWA pool. The idle
# pool-leak invariant requires swa total == swa available, so
# reconcile the reported SWA total to the allocator's actual idle ring
# capacity. Safe: on unified_kv swa_kv_pool is None, so no real buffer
# is resized -- this only fixes token accounting / usage reporting.
if (
self.is_hybrid_swa
and not self.is_draft_worker
and getattr(pools.token_to_kv_pool, "_unified_kv", False)
):
alloc = pools.token_to_kv_pool_allocator
if hasattr(alloc, "swa_available_size"):
ring_capacity = int(alloc.swa_available_size())
# Only reconcile downward to the (smaller) ring capacity. A
# value >= the current total means swa_available_size() hit a
# non-binding fallback (e.g. req_to_token pool not wired), in
# which case leave the reported total untouched.
if 0 < ring_capacity < swa_max_total_num_tokens:
logger.info(
"Unified-KV: reconciling swa_max_total_num_tokens "
f"{swa_max_total_num_tokens} -> {ring_capacity} "
"(fixed per-request SWA ring capacity)."
)
swa_max_total_num_tokens = 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"
@@ -372,7 +345,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=swa_max_total_num_tokens, swa_max_total_num_tokens=sizes.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,
@@ -1033,7 +1006,9 @@ class KVCacheConfigurator:
extra_max_context_len: int, extra_max_context_len: int,
pre_alloc_size: int, pre_alloc_size: int,
) -> ReqToTokenPool: ) -> ReqToTokenPool:
from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool from sglang.srt.disaggregation.decode import (
HybridMambaDecodeReqToTokenPool,
)
req_to_token_pool = HybridMambaDecodeReqToTokenPool( req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
@@ -1321,7 +1296,9 @@ class KVCacheConfigurator:
assert swa_page_size == 256, "In paged swa mode, page_size must be 256." assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
if self.is_draft_worker: if self.is_draft_worker:
from sglang.srt.models.deepseek_v4_nextn import COMPRESS_RATIO_NEXTN_LAYER from sglang.srt.models.deepseek_v4_nextn import (
COMPRESS_RATIO_NEXTN_LAYER,
)
compression_ratios = [ compression_ratios = [
COMPRESS_RATIO_NEXTN_LAYER COMPRESS_RATIO_NEXTN_LAYER
@@ -1436,7 +1413,9 @@ class KVCacheConfigurator:
full_max_total_num_tokens: Optional[int], full_max_total_num_tokens: Optional[int],
swa_max_total_num_tokens: Optional[int], swa_max_total_num_tokens: Optional[int],
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
kwargs = {} kwargs = {}
if self.is_hybrid_swa_compress: if self.is_hybrid_swa_compress:
@@ -1502,7 +1481,9 @@ class KVCacheConfigurator:
def _build_ascend_mla_kv_pool( def _build_ascend_mla_kv_pool(
self, *, max_total_num_tokens: int, is_dsa_model: bool self, *, max_total_num_tokens: int, is_dsa_model: bool
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMLATokenToKVPool,
)
token_to_kv_pool = NPUMLATokenToKVPool( token_to_kv_pool = NPUMLATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
@@ -1520,7 +1501,9 @@ class KVCacheConfigurator:
return token_to_kv_pool return token_to_kv_pool
def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
token_to_kv_pool = NPUMHATokenToKVPool( token_to_kv_pool = NPUMHATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
@@ -1960,11 +1943,12 @@ 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:
from sglang.srt.mem_cache.sparsity import parse_hisparse_config from sglang.srt.mem_cache.sparsity import (
parse_hisparse_config,
)
hisparse_cfg = parse_hisparse_config() hisparse_cfg = parse_hisparse_config()
token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator(
@@ -408,7 +408,9 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
indexer_ratio = parse_hisparse_config().host_to_device_ratio indexer_ratio = parse_hisparse_config().host_to_device_ratio
from sglang.srt.mem_cache.kv_cache_configurator import _should_elide_dsa_index_k from sglang.srt.mem_cache.kv_cache_configurator import (
_should_elide_dsa_index_k,
)
if allocate_all_layers or not _should_elide_dsa_index_k( if allocate_all_layers or not _should_elide_dsa_index_k(
is_draft_worker=kvc.is_draft_worker is_draft_worker=kvc.is_draft_worker
@@ -845,10 +847,7 @@ 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). The decode reserves a draft worker, mirroring dflash's cell_size scaling); bias = 0.
bias is the sum of request-scoped fixed pools that do not scale with
full_token: the c128 state pool and, on the unified_kv path, the fixed SWA
per-request ring (bf16, see _fixed_swa_bytes).
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
@@ -905,27 +904,6 @@ 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)
# Unified-KV uses a different physical layout than the fp8 path:
# * KV is stored bf16 over the full latent (attn_head_dim * 2 bytes),
# not the fp8(nope) + bf16(rope) + scales 584-byte cell.
# * SWA is a fixed per-request ring (num_req_slots * ring_size),
# independent of full_token, so it is a fixed *bias* rather than a
# per-token term. Gate on the same switch the pool itself uses so the
# sizing and the allocation never drift apart.
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
# Mirror DeepSeekV4TokenToKVPool: swa_ring_size = sliding_window +
# (speculative_num_draft_tokens - 1).
spec_num_draft = get_spec().speculative_num_draft_tokens or 1
self._swa_ring_size = self.swa_page_size + (
(spec_num_draft - 1) if self.is_speculative else 0
)
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(
@@ -940,8 +918,7 @@ 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._spec_infl = (target_layers + draft_layers) / target_layers self.bytes_per_full_token *= (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)
@@ -994,10 +971,6 @@ 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
@@ -1022,41 +995,17 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
c4_frac = 1 / (4 * self.c4_shrink_factor) c4_frac = 1 / (4 * self.c4_shrink_factor)
return ( return (
# Unified_kv: SWA is a fixed per-request ring (see _fixed_swa_bytes), self.swa_ratio * kv_bytes * self.num_layers_total
# not a per-token pool, so it is excluded from the per-token coeff.
(
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
# Unified_kv: the c4 (attn + indexer) compress-state is a ring buffer + self.swa_ratio * c4_state_ratio * c4_state_bytes * self.num_layers_ca4
# addressed off the SWA slot ((swa_loc // swa_page_size) * ring_size),
# and the unified SWA pool is a fixed per-request ring
# (swa_pages = num_req_slots * swa_ring_size), so the state ring is
# request-scoped, not full_token-scoped. It is therefore a fixed bias
# (see _fixed_c4_state_bytes), not a per-token term. On the non-unified
# path the SWA pool scales with full_token, so it stays per-token.
+ (
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
@@ -1067,14 +1016,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
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,
# Unified_kv sizes the c4 state ring from the fixed SWA ring c4_state_pool_size=swa_tokens // self.swa_page_size * self.c4_ring_size,
# (request-scoped), finalized once max_running_requests is known -- so
# it must not scale with full_token here (mirrors c128_state below).
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,
) )
@@ -1105,64 +1047,18 @@ 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 _unified_c4_state_pool_size(self, max_running_requests: int) -> int: def _get_c128_state_fixed_bytes_for_token_capacity(
"""Exact request-scoped C4 ring size for the unified address contract. self, token_capacity: int
) -> int:
Unified C4 state locations are
``req_pool_idx * c4_ring_size + position % c4_ring_size``.
"""
num_req_slots = self._get_num_req_slots(max_running_requests)
return num_req_slots * self.c4_ring_size
def _fixed_c4_state_bytes(self, max_running_requests: int) -> int:
"""Unified_kv c4 (attn + indexer) compress-state is a fixed per-request
ring, sized by concurrency rather than full_token. Return its byte
footprint across all c4 layers. Returns 0 on the non-unified path (where
the c4 state pool scales with the SWA pool and is accounted per-token)."""
if not self._unified or self.num_layers_ca4 == 0:
return 0
c4_state_dtype_size, _ = _get_dsv4_compress_state_dtype_sizes()
attn_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
# CompressStatePool allocates `size + ring_size + 1` rows, padded to the
# compress ratio (see CompressStatePool.__init__). Mirror that here so the
# reserved bias covers the real allocation.
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 * 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:
"""Approximate ModelRunner._resolve_max_num_reqs closely enough to size
the request-scoped fixed pools (c128 state, unified SWA ring). Over-
estimating is safe: a larger fixed bias yields a smaller full_token."""
if self.requested_max_running_requests_per_worker is not None: if self.requested_max_running_requests_per_worker is not None:
return self.requested_max_running_requests_per_worker return self._get_c128_state_fixed_bytes(
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)
return min(estimated, full_token // 2)
def _fixed_swa_bytes(self, max_running_requests: int) -> int:
"""Unified_kv SWA is a fixed per-request ring, sized by concurrency
(num_req_slots) rather than by full_token. Return its bf16 byte
footprint across all full layers, inflated for the draft worker the same
way as the per-token coeff. Returns 0 on the non-unified path (where SWA
is already accounted per-token)."""
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)
estimated = int(token_capacity / self.context_len * 512)
estimated = max(min(estimated, 4096), 2048)
max_running_requests = min(estimated, token_capacity // 2)
return self._get_c128_state_fixed_bytes(max_running_requests)
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
@@ -1193,13 +1089,6 @@ 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
# Unified_kv: the c4 state ring is request-scoped (fixed SWA pool), so
# finalize it here from the now-known concurrency. On the non-unified path
# it was already sized from full_token in _compute_dsv4_sizes.
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(
@@ -1209,34 +1098,25 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
"page_size must be multiple of 128 for compressed attention" "page_size must be multiple of 128 for compressed attention"
) )
max_running_requests_per_worker = self._resolve_max_running_requests_per_worker( if self.requested_max_running_requests_per_worker is not None:
available_bytes
)
c128_state_fixed_bytes = self._get_c128_state_fixed_bytes( c128_state_fixed_bytes = self._get_c128_state_fixed_bytes(
max_running_requests_per_worker self.requested_max_running_requests_per_worker
) )
swa_ring_fixed_bytes = self._fixed_swa_bytes(max_running_requests_per_worker) else:
c4_state_fixed_bytes = self._fixed_c4_state_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)
) )
available_bytes_for_tokens = max( available_bytes_for_tokens = max(available_bytes - c128_state_fixed_bytes, 0)
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: unified={self._unified}, " f"DSV4 memory calculation: "
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)
@@ -7,11 +7,7 @@ import pytest
import torch import torch
import triton import triton
from sglang.kernels.ops.attention.dsv4 import ( from sglang.kernels.ops.attention.dsv4 import compress_forward
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 (
@@ -126,91 +122,6 @@ 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:
@@ -9,7 +9,10 @@ from sglang.srt.managers.schedule_policy import (
PrefillAdder, PrefillAdder,
estimate_prefill_extend_tile_metrics, estimate_prefill_extend_tile_metrics,
) )
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefResult, IncLockRefResult from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefResult,
IncLockRefResult,
)
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils.common import Range from sglang.srt.utils.common import Range
@@ -68,13 +71,6 @@ class TestPrefillAdder(CustomTestCase):
allocator.swa_available_size.return_value = swa_available_size allocator.swa_available_size.return_value = swa_available_size
allocator.available_size.return_value = available_size allocator.available_size.return_value = available_size
allocator.size_swa = size_swa allocator.size_swa = size_swa
# get_kvcache().[_unified_kv] gates the unified-KV SWA-ring accounting
# path in schedule_policy.add_chunked_req / rem_swa_tokens. A bare
# MagicMock auto-creates any attribute access as a truthy Mock, so
# without this the getattr(..., "_unified_kv", False) default never
# triggers and these tests silently exercise the unified-KV branch
# instead of the standard hybrid-SWA one they intend to cover.
allocator.get_kvcache.return_value._unified_kv = False
return allocator return allocator
def create_running_batch(self, reqs=None) -> MagicMock: def create_running_batch(self, reqs=None) -> MagicMock:
@@ -23,9 +23,6 @@ class _FakeAllocator:
self.alloc_calls = [] self.alloc_calls = []
self.extend_calls = [] self.extend_calls = []
def get_kvcache(self):
return None
def available_size(self): def available_size(self):
return 1 << 30 return 1 << 30
@@ -1,119 +0,0 @@
"""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.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 _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)),
)
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
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 = True
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)
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_alloc_clears_new_slots_but_not_reused_slots(self):
req_pool = ReqToTokenPool(3, 16, "cpu", enable_memory_saver=False)
token_pool = MagicMock()
reused = _request()
# First admission: a brand-new slot, so its C4 ring must be cleared.
(reused_idx,) = alloc_req_slots(
req_pool, [reused], None, token_to_kv_pool=token_pool
)
token_pool.clear_c4_req_states.assert_called_once_with([reused_idx])
# Chunked continuation reuses the same slot -- clearing it here would
# wipe the state captured by the previous chunk.
token_pool.clear_c4_req_states.reset_mock()
reused.kv.req_pool_idx = reused_idx
reused.kv.kv_committed_len = 1
reused.kv.kv_allocated_len = 1
reused.kv.holds_kv = True
reused.inflight_middle_chunks = 1
self.assertEqual(
alloc_req_slots(req_pool, [reused], None, token_to_kv_pool=token_pool),
[reused_idx],
)
token_pool.clear_c4_req_states.assert_not_called()
# Mixed batch: only the newly allocated slot is cleared.
fresh = _request()
indices = alloc_req_slots(
req_pool, [reused, fresh], None, token_to_kv_pool=token_pool
)
self.assertEqual(indices[0], reused_idx)
self.assertNotEqual(indices[1], reused_idx)
token_pool.clear_c4_req_states.assert_called_once_with([indices[1]])
if __name__ == "__main__":
unittest.main()
@@ -131,7 +131,6 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
queue = DecodePreallocQueue.__new__(DecodePreallocQueue) queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
queue.req_to_token_pool = req_to_token_pool queue.req_to_token_pool = req_to_token_pool
queue.token_to_kv_pool_allocator = allocator queue.token_to_kv_pool_allocator = allocator
queue.token_to_kv_pool = None
queue.tree_cache = SimpleNamespace( queue.tree_cache = SimpleNamespace(
evictable_size=MagicMock(return_value=0), evictable_size=MagicMock(return_value=0),
protected_size=MagicMock(return_value=0), protected_size=MagicMock(return_value=0),
@@ -36,11 +36,6 @@ 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 branches on self._unified to skip the vestigial paged SWA
# allocator on the unified-KV path. This stub exercises the standard
# hybrid-SWA path, so pin it False rather than letting the attribute go
# missing (SimpleNamespace raises instead of defaulting).
_unified=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),