diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh index e341d4ba4..7e7b1005e 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh @@ -28,14 +28,12 @@ using R2T_T = int32_t; using F2S_T = int64_t; using IDX_T = int64_t; -/// NOTE: for the internal use, we pack the ragged and batch id, since both not -/// exceed 65536 +/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 SGL_DEVICE __host__ PlanW pack_w(uint32_t ragged_id, uint32_t batch_id, int32_t seq_len) { return {static_cast(ragged_id | batch_id << 16), seq_len}; } -/// NOTE: for the internal use, we pack the ragged and batch id, since both not -/// exceed 65536 +/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 SGL_DEVICE uint2 unpack_w(PlanW plan) { return {static_cast(plan.ragged_id), static_cast(plan.ragged_id >> 16)}; } @@ -49,11 +47,9 @@ struct Prefill0Params { uint32_t num_q_tokens; int32_t compress_ratio; int32_t swa_page_size; - /// \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. + /// \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. int32_t mtp_pad; - bool use_req_ring; }; struct Prefill1Params { @@ -71,7 +67,6 @@ struct Prefill1Params { int32_t swa_page_size; int32_t ring_size; int32_t compress_ratio; - bool use_req_ring; }; struct DecodeParams { @@ -85,7 +80,6 @@ struct DecodeParams { int32_t swa_page_size; int32_t ring_size; int32_t compress_ratio; - bool use_req_ring; }; struct Prefill1ParamsLegacy { @@ -161,8 +155,7 @@ __global__ __launch_bounds__(1024, 1) // counter_w = 0; } // === 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 - // min down. + // For min, treat threads outside `batch_size` as +inf so they don't pull the min down. const uint32_t e_for_max = static_cast(extend_len); const uint32_t e_for_min = (tx < params.batch_size) ? e_for_max : 0xFFFFFFFFu; warp_max[warp_id] = warp::reduce_max(e_for_max); @@ -175,19 +168,17 @@ __global__ __launch_bounds__(1024, 1) // __syncthreads(); const auto num_q = params.num_q_tokens; - // MTP-uniform: every batch shares the same small extend_len `E`, so we can - // decompose a global token id `k` into (batch_id, j) = (k / E, k % E) and - // skip the per-batch loop. + // MTP-uniform: every batch shares the same small extend_len `E`, so we can decompose + // a global token id `k` into (batch_id, j) = (k / E, k % E) and skip the per-batch loop. 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 === if (is_mtp_extend) { - // Path 1: token-driven. Each global token id maps to exactly one (batch_id, - // j). + // Path 1: token-driven. Each global token id maps to exactly one (batch_id, j). const uint32_t E = s_max_extend; - // num_q is the padded buffer size (graph bucket), not the work size: cap - // the loop at the real token count so batch_id = k / E stays < batch_size - // on an underfilled replay; Stage D pads [counter, num_q) with invalid. + // num_q is the padded buffer size (graph bucket), not the work size: cap the + // loop at the real token count so batch_id = k / E stays < batch_size on an + // underfilled replay; Stage D pads [counter, num_q) with invalid. const uint32_t num_real_q = params.batch_size * E; for (uint32_t k = tx; k < num_real_q; k += block_size) { 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 first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); 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) { const uint32_t out_idx = atomicAdd(&counter_w, 1u); params.plan_w[out_idx] = pack_w(ragged_id, batch_id, position + 1); } } } else { - // Path 2: general prefill (long extend_len). Iterate batches in an outer - // loop; the whole block sweeps each batch's tokens in parallel. + // Path 2: general prefill (long extend_len). Iterate batches in an outer loop; + // the whole block sweeps each batch's tokens in parallel. uint32_t base_e = 0; for (uint32_t batch_id = 0; batch_id < params.batch_size; ++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; - 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) { const uint32_t out_idx = atomicAdd(&counter_w, 1u); params.plan_w[out_idx] = pack_w(ragged_id, static_cast(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; 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(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(plan_c.seq_len - 1); // only used for c4, harmless for c128 const auto position_0 = max(position_1 - params.compress_ratio, 0); - if (params.compress_ratio == 128 || params.use_req_ring) { - plan_c.read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio; - plan_c.read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio; + if (params.compress_ratio == 128) { + plan_c.read_page_0 = compute_c128_loc(rid, position_0) / 128; + plan_c.read_page_1 = compute_c128_loc(rid, position_1) / 128; } else { const auto raw_loc_0 = mapping[position_0]; 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 const auto position = static_cast(plan_w.write_loc - 1); plan_w.ragged_id = ragged_id; - if (params.compress_ratio == 128 || params.use_req_ring) { - plan_w.write_loc = compute_req_ring_loc(rid, position); + if (params.compress_ratio == 128) { + plan_w.write_loc = compute_c128_loc(rid, position); } else { const auto raw_loc = mapping[position]; 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; 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(rid * params.ring_size + position % params.ring_size); }; const auto seq_len = static_cast(params.seq_ptr[idx]); @@ -347,10 +338,10 @@ __global__ void plan_compress_decode_kernel(const DecodeParams params) { int32_t write_loc; int32_t read_page_0; int32_t read_page_1; - if (params.compress_ratio == 128 || params.use_req_ring) { - write_loc = compute_req_ring_loc(rid, position_1); - read_page_0 = compute_req_ring_loc(rid, position_0) / params.compress_ratio; - read_page_1 = compute_req_ring_loc(rid, position_1) / params.compress_ratio; + if (params.compress_ratio == 128) { + write_loc = compute_c128_loc(rid, position_1); + read_page_0 = compute_c128_loc(rid, position_0) / 128; + read_page_1 = compute_c128_loc(rid, position_1) / 128; } else { const auto raw_loc_0 = mapping[position_0]; 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(); /// Per-request ring buffer slot translation: - /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % - /// 4 - /// - c128: page = rid; slot = rid * 128 + position % - /// 128 + /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 + /// - c128: page = rid; slot = rid * 128 + position % 128 const auto legacy_compute_page = [&](int32_t rid, int32_t position) { if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); return rid; // c128 @@ -404,8 +393,7 @@ __global__ void plan_compress_prefill_legacy_kernel(const Prefill1ParamsLegacy p if (!plan_w.is_invalid()) { const auto [ragged_id, batch_id] = unpack_w(plan_w); const auto rid = static_cast(params.rid_ptr[batch_id]); - // `write_loc` carries (position + 1) at this stage; may not be - // ratio-aligned + // `write_loc` carries (position + 1) at this stage; may not be ratio-aligned const auto position = static_cast(plan_w.write_loc) - 1; plan_w.ragged_id = ragged_id; 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; if (idx >= params.batch_size) return; /// Per-request ring buffer slot translation: - /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % - /// 4 - /// - c128: page = rid; slot = rid * 128 + position % - /// 128 + /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 + /// - c128: page = rid; slot = rid * 128 + position % 128 const auto legacy_compute_page = [&](int32_t rid, int32_t position) { if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); return rid; // c128 @@ -461,8 +447,7 @@ using PrefillPlan = tvm::ffi::Tuple; * @param compress_plan `[num_q_tokens, 16]` uint8 (output) * @param write_plan `[num_q_tokens, 8]` uint8 (output) * @param compress_ratio 4 for c4, 128 for c128 - * @param use_cuda_graph Whether the plans will be used with cuda graph (affects - * padding) + * @param use_cuda_graph Whether the plans will be used with cuda graph (affects padding) * @return (compress plan tensor, write plan tensor) */ inline PrefillPlan plan_compress_prefill( @@ -476,7 +461,6 @@ inline PrefillPlan plan_compress_prefill( const int32_t compress_ratio, const int32_t swa_page_size, const int32_t ring_size, - const bool use_req_ring, const bool use_cuda_graph) { auto B = SymbolicSize{"batch_size"}; auto N = SymbolicSize{"num_q_tokens"}; @@ -519,29 +503,27 @@ inline PrefillPlan plan_compress_prefill( const auto batch_size = static_cast(B.unwrap()); constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); - RuntimeCheck(!use_req_ring || compress_ratio == 4); RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); // `swa_page_size` >= `ring_size` >= `compress_ratio` 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 - // survives any accept length. Zero without speculation -- nothing rolls back, - // and the ring is then exactly one window wide. Otherwise the ring bounds it: - // a write at `w` aliases onto `w - ring_size`, and the earliest position a - // future compression still needs is `prefix_len - window_size + 2` (the next - // batch commits >= 1 token, and `run_prefill` launches the compress kernel - // before the write kernel, so a batch's own compressions read the pre-write - // ring). Padding past the extend range is harmless: the loops only span - // `[prefix_len, seq_len)`. + // Write pad: trailing tokens kept resident so a verify batch's committed tail survives + // any accept length. Zero without speculation -- nothing rolls back, and the ring is + // then exactly one window wide. Otherwise the ring bounds it: a write at `w` aliases + // onto `w - ring_size`, and the earliest position a future compression still needs is + // `prefix_len - window_size + 2` (the next batch commits >= 1 token, and `run_prefill` + // launches the compress kernel before the write kernel, so a batch's own compressions + // read the pre-write ring). Padding past the extend range is harmless: the loops only + // span `[prefix_len, seq_len)`. const auto mtp_pad = ring_size > window_size ? ring_size - window_size + 2 : 0; const auto device = device_.unwrap(); const auto stream = LaunchKernel::resolve_device(device); if (cpu_or_gpu.unwrap().device_type == kDLGPU) { - // GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata - // directly on device, padding to num_q_tokens with invalid; kernel_1 then - // finalizes the SWA-translated read/write locations. Used for MTP / - // cuda-graph capture where a host sync would be expensive. + // GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata directly + // on device, padding to num_q_tokens with invalid; kernel_1 then finalizes the + // SWA-translated read/write locations. Used for MTP / cuda-graph capture where + // a host sync would be expensive. 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 W = ffi::empty({num_q_tokens, sizeof(PlanW)}, kDLUInt8, device); @@ -555,11 +537,9 @@ inline PrefillPlan plan_compress_prefill( .compress_ratio = compress_ratio, .swa_page_size = swa_page_size, .mtp_pad = mtp_pad, - .use_req_ring = use_req_ring, }; 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. const auto params1 = Prefill1Params{ .plan_c = static_cast(C.data_ptr()), .plan_w = static_cast(W.data_ptr()), @@ -575,7 +555,6 @@ inline PrefillPlan plan_compress_prefill( .swa_page_size = swa_page_size, .ring_size = ring_size, .compress_ratio = compress_ratio, - .use_req_ring = use_req_ring, }; const auto block_size_1 = 256; 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); const auto should_write = [=](int32_t position) { 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)) { const int32_t position = prefix_len + j; @@ -652,7 +631,6 @@ inline PrefillPlan plan_compress_prefill( .swa_page_size = swa_page_size, .ring_size = ring_size, .compress_ratio = compress_ratio, - .use_req_ring = use_req_ring, }; const auto block_size = 256; 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 int32_t compress_ratio, const int32_t swa_page_size, - const int32_t ring_size, - const bool use_req_ring) { + const int32_t ring_size) { auto B = SymbolicSize{"batch_size"}; auto device_ = SymbolicDevice{}; device_.set_options(); @@ -690,7 +667,6 @@ inline tvm::ffi::Tensor plan_compress_decode( .with_device(device_) .verify(seq_lens); - RuntimeCheck(!use_req_ring || compress_ratio == 4); const auto batch_size = static_cast(B.unwrap()); const auto device = device_.unwrap(); 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, .ring_size = ring_size, .compress_ratio = compress_ratio, - .use_req_ring = use_req_ring, }; const auto block_size = 256; const auto num_blocks = div_ceil(batch_size, block_size); diff --git a/python/sglang/kernels/ops/attention/dsv4/attn.py b/python/sglang/kernels/ops/attention/dsv4/attn.py index b98dc922d..f973bf916 100644 --- a/python/sglang/kernels/ops/attention/dsv4/attn.py +++ b/python/sglang/kernels/ops/attention/dsv4/attn.py @@ -100,7 +100,6 @@ def create_paged_compress_data_kernel( stride_out_1_1: tl.constexpr, compress_ratio: tl.constexpr, is_overlap: tl.constexpr, - use_req_ring: tl.constexpr, swa_page_size: tl.constexpr, ring_size: tl.constexpr, BLOCK: tl.constexpr, @@ -136,8 +135,6 @@ def create_paged_compress_data_kernel( pos = tl.maximum(pos, 0) if compress_ratio == 128: state_loc = rid * ring_size + (pos % ring_size) - elif use_req_ring: - state_loc = rid * ring_size + (pos % ring_size) else: loc = tl.load( req_to_token_ptr @@ -185,7 +182,6 @@ def triton_create_paged_compress_data( extend_seq_lens: torch.Tensor, req_to_token: torch.Tensor, full_to_swa_index_mapping: torch.Tensor, - use_req_ring: bool = False, block: int = 128, ) -> Tuple[torch.Tensor, torch.Tensor]: 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 compress_ratio=compress_ratio, # 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 ring_size=ring_size, # type: ignore BLOCK=block, # type: ignore diff --git a/python/sglang/kernels/ops/attention/dsv4/compress.py b/python/sglang/kernels/ops/attention/dsv4/compress.py index f126ba175..a9e915d67 100644 --- a/python/sglang/kernels/ops/attention/dsv4/compress.py +++ b/python/sglang/kernels/ops/attention/dsv4/compress.py @@ -162,7 +162,6 @@ class CompressorDecodePlan(NamedTuple): seq_lens: torch.Tensor, swa_page_size: int, ring_size: int, - use_req_ring: bool = False, ) -> CompressorDecodePlan: if _is_xpu: fn = plan_compress_decode @@ -170,7 +169,7 @@ class CompressorDecodePlan(NamedTuple): module = _jit_compress_plan_module() fn = module.plan_decode - args = ( + plan_d = fn( req_pool_indices, req_to_token, full_to_state, @@ -179,7 +178,6 @@ class CompressorDecodePlan(NamedTuple): int(swa_page_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)) @staticmethod @@ -249,7 +247,6 @@ class CompressorPrefillPlan(NamedTuple): ring_size: int, num_q_tokens: int, use_cuda_graph: bool = False, - use_req_ring: bool = False, ) -> CompressorPrefillPlan: is_gpu_input = seq_lens.device.type in ["cuda", "xpu"] pin_buffer = torch.empty( @@ -277,7 +274,7 @@ class CompressorPrefillPlan(NamedTuple): module = _jit_compress_plan_module() fn = module.plan_prefill - args = ( + plan_c, plan_w = fn( req_pool_indices, req_to_token, full_to_state, @@ -288,11 +285,7 @@ class CompressorPrefillPlan(NamedTuple): int(compress_ratio), int(swa_page_size), int(ring_size), - ) - plan_c, plan_w = ( - fn(*args, bool(use_cuda_graph)) - if _is_xpu - else fn(*args, bool(use_req_ring), bool(use_cuda_graph)) + bool(use_cuda_graph), ) return CompressorPrefillPlan( compress_ratio, diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 06b4384b5..03cd52a33 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1768,18 +1768,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): if total_prefix_len is None: 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]) assert req_pool_indices is not None, ( "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) req.kv.kv_committed_len = fill_len diff --git a/python/sglang/srt/layers/attention/dsv4/compress_hip.py b/python/sglang/srt/layers/attention/dsv4/compress_hip.py index 7ae66e27b..225a008a3 100644 --- a/python/sglang/srt/layers/attention/dsv4/compress_hip.py +++ b/python/sglang/srt/layers/attention/dsv4/compress_hip.py @@ -144,9 +144,7 @@ class CompressorHip(_CompressorBase): pre_state_indices = self.compute_state_len_indices( seq_len=prefix_lens[i], ratio=self.ratio ).to(device) - if self.ratio == 128 or ( - self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False) - ): + if self.ratio == 128: state_loc = state_pool.translate_from_req_position_to_state_loc( req_pool_indices[i], pre_state_indices ) @@ -168,9 +166,7 @@ class CompressorHip(_CompressorBase): post_state_len = post_state_indices.size(0) assert post_state_len <= valid_kv_len - if self.ratio == 128 or ( - self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False) - ): + if self.ratio == 128: post_state_loc = state_pool.translate_from_req_position_to_state_loc( req_pool_indices[i], post_state_indices ) @@ -275,9 +271,7 @@ class CompressorHip(_CompressorBase): seq_lens = seq_lens_2d.view(-1) req_pool_indices = req_pool_indices.repeat_interleave(draft_tokens) - if self.ratio == 128 or ( - self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False) - ): + if self.ratio == 128: state_locs = state_pool.translate_from_req_position_to_state_loc( req_pool_indices, seq_lens - 1 ) @@ -292,9 +286,7 @@ class CompressorHip(_CompressorBase): -compress_bulk_len, 0, device=seq_lens.device ) compress_indices.clamp_(min=-1) - if self.ratio == 128 or ( - self.ratio == 4 and getattr(token_to_kv_pool, "_unified_kv", False) - ): + if self.ratio == 128: compress_indices_state = ( state_pool.translate_from_req_position_to_state_loc( req_pool_indices[:, None], compress_indices diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index d5b8c29b8..9b213c1ba 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -26,7 +26,9 @@ from sglang.srt.layers.utils.cp_utils import ( cp_all_gather_rerange_finish, 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.model_executor.forward_context import get_attn_backend from sglang.srt.models.deepseek_v2 import _is_hip @@ -262,7 +264,6 @@ def create_paged_compressor_data( ) -> FusedCompressMetadata: swa_page_size = token_to_kv_pool.swa_page_size 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 def clip_down(positions: torch.Tensor) -> torch.Tensor: @@ -272,8 +273,6 @@ def create_paged_compressor_data( positions = positions.masked_fill(positions < 0, 0) if compress_ratio == 128: 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: loc = req_to_token[req_pool_indices, positions] 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, req_to_token=req_to_token, full_to_swa_index_mapping=token_to_kv_pool.full_to_swa_index_mapping, - use_req_ring=use_req_ring, ) plan_kwargs: dict diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index c35154131..707efb74d 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -441,7 +441,6 @@ def create_paged_compressor_data( swa_page_size = token_to_kv_pool.swa_page_size 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. # As a workaround, we use `.detach()` to get the real tensor. 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, swa_page_size=swa_page_size, ring_size=ring_size, - use_req_ring=use_req_ring, num_q_tokens=num_q_tokens, use_cuda_graph=use_prefill_cuda_graph, ) @@ -481,7 +479,6 @@ def create_paged_compressor_data( seq_lens=seq_lens.to(torch.int64), swa_page_size=swa_page_size, ring_size=ring_size, - use_req_ring=use_req_ring, ) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index e539e91e9..db3d02c29 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -57,7 +57,6 @@ import dataclasses import logging import re import sys -import time from array import array from concurrent.futures import Future 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 ( 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.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( @@ -159,9 +161,6 @@ _MM_HASH_MASK = (1 << 64) - 1 logger = logging.getLogger(__name__) -# Throttle for the unified-KV SWA bottleneck diagnostic (seconds). -_last_swa_bottleneck_log = 0.0 - ReturnHiddenStatesMode = Union[bool, Literal["last"]] @@ -3078,50 +3077,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): shortfalls retract gracefully instead of tripping fail-loud alloc errors.""" num_tokens = self.new_tokens_required_next_decode(selected_indices) - allocator = self.token_to_kv_pool_allocator - ok = allocator.check_decode_capacity( + return self.token_to_kv_pool_allocator.check_decode_capacity( 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]]: """Retract the decoding requests when there is not enough memory.""" diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index a4fbf5e6e..023071a67 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -5,7 +5,10 @@ from array import array from sglang.srt.environ import envs 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 _ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG") @@ -660,16 +663,8 @@ class PrefillAdder: @property 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 ( - allocator.swa_available_size() + self.token_to_kv_pool_allocator.swa_available_size() + self.tree_cache.swa_evictable_size() - self.rem_swa_token_offset ) @@ -716,13 +711,6 @@ class PrefillAdder: where alloc = min(extend, rem_chunk); the min() cap keeps the two terms 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: alloc = min(extend_input_len, self.rem_chunk_tokens) else: @@ -850,9 +838,6 @@ class PrefillAdder: max_new_tokens: int, retracted_stain: bool, 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 extend_input_len = self.ceil_paged_tokens(extend_input_len) @@ -876,17 +861,9 @@ class PrefillAdder: self.rem_input_tokens -= extend_input_len 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 + self.rem_swa_token_offset += self._swa_budget_for_req( + extend_input_len, max_new_tokens ) - if not (_unified and is_chunked_continuation): - self.rem_swa_token_offset += self._swa_budget_for_req( - extend_input_len, max_new_tokens - ) if self.dllm_config is not None: self.rem_dllm_tokens -= extend_input_len @@ -1021,15 +998,9 @@ class PrefillAdder: _rem_tokens = self._get_dllm_remain_tokens() else: _rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens)) - if self.is_hybrid_swa and not getattr( - self.token_to_kv_pool_allocator.get_kvcache(), "_unified_kv", False - ): + if self.is_hybrid_swa: # alloc_extend needs extend_num_tokens + page_size per request, - # 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. + # so reserve one page here to avoid OOM _rem_tokens = min( _rem_tokens, int(self.rem_swa_tokens) - self.page_size ) @@ -1068,7 +1039,6 @@ class PrefillAdder: ), req.retracted_stain, mamba_gap_reserve=self._mamba_gap_budget_for_req(req), - is_chunked_continuation=True, ) # Return if chunked prefill not finished @@ -1278,7 +1248,7 @@ class PrefillAdder: self._swa_new_tokens(req), 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( real_input_tokens, self._swa_new_tokens(req), @@ -1314,7 +1284,7 @@ class PrefillAdder: self._swa_new_tokens(req), 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( real_input_tokens, self._swa_new_tokens(req), diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index fc861c589..ed24b8422 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -3,7 +3,14 @@ from __future__ import annotations import logging from collections import deque 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 @@ -25,7 +32,10 @@ from sglang.srt.observability.scheduler_stage_metrics import ( scheduler_stage_method, ) 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 if TYPE_CHECKING: @@ -142,23 +152,6 @@ class SchedulerInvariantChecker: def _check_swa_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]: 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 if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator): # 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 # 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 = getattr(alloc, "logical_attn_allocator", alloc) sub_allocs = ( [alloc] if getattr(alloc, "free_pages", None) is not None diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index e8f9a6697..e1dd565eb 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -2,7 +2,14 @@ from __future__ import annotations import dataclasses 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 ( UnifiedMambaSWATokenToKVPoolAllocator, @@ -294,16 +301,6 @@ class SchedulerPoolStatsObserver: swa_available_size = allocator.swa_available_size() full_evictable_size = self.tree_cache.full_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_available_size + full_evictable_size ) diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index 4cf251db0..5e109a4c8 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -230,7 +230,6 @@ def alloc_req_slots( req_to_token_pool: ReqToTokenPool, reqs: list[Req], tree_cache: BasePrefixCache | None, - token_to_kv_pool=None, ) -> list[int]: """Allocate request slots from the pool. @@ -261,7 +260,6 @@ def alloc_req_slots( tree_cache.evict_for_alloc( 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) if req_pool_indices is None: raise RuntimeError( @@ -269,13 +267,6 @@ def alloc_req_slots( "Please set a smaller number for `--max-running-requests`. " 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 @@ -320,10 +311,7 @@ def alloc_for_extend( # Allocate req slots (raises RuntimeError if the pool is exhausted) req_pool_indices = alloc_req_slots( - batch.req_to_token_pool, - batch.reqs, - batch.tree_cache, - token_to_kv_pool=batch.token_to_kv_pool_allocator.get_kvcache(), + batch.req_to_token_pool, batch.reqs, batch.tree_cache ) req_pool_indices_cpu = torch.tensor( req_pool_indices, dtype=torch.int64, pin_memory=pin_memory diff --git a/python/sglang/srt/mem_cache/allocator/hisparse.py b/python/sglang/srt/mem_cache/allocator/hisparse.py index 44047c9b5..5647154f7 100644 --- a/python/sglang/srt/mem_cache/allocator/hisparse.py +++ b/python/sglang/srt/mem_cache/allocator/hisparse.py @@ -343,10 +343,6 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): def get_kvcache(self): 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): return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices) diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index c5abd1d75..4b5c61812 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -1,5 +1,3 @@ -import logging - import torch 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.invariants import Bucket, Invariant, IsTrue, expect -logger = logging.getLogger(__name__) - _is_npu = is_npu() if _is_npu: @@ -41,7 +37,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): device: str, kvcache: BaseSWAKVPool, need_sort: bool, - req_to_token_pool=None, ): assert isinstance(kvcache, BaseSWAKVPool) self._size_full = size @@ -109,46 +104,10 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self.swa_free_group = [] 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._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): - 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( self.full_attn_allocator.available_size(), self.swa_attn_allocator.available_size(), @@ -158,12 +117,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): return self.full_attn_allocator.available_size() 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() # Slot-conservation views for the leak invariant. On the non-shared allocator @@ -218,15 +171,11 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): return alloc_full_indices def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool: - full_ok = ( + return ( num_full_pages <= self.full_attn_allocator.available_size() // self.page_size - ) - if self._unified: - # 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 + and num_swa_pages + <= self.swa_attn_allocator.available_size() // self.page_size ) def alloc_extend( @@ -246,20 +195,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): if not self.new_pages_available(num_new_pages, num_new_pages): 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) 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 ): 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) alloc_full_indices = self.full_attn_allocator.alloc_decode( diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 991a8a5fb..8287abbf8 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from contextlib import nullcontext -from typing import List, Literal, NamedTuple, Optional, Sequence, Tuple +from typing import List, Literal, NamedTuple, Optional, Tuple import torch @@ -564,11 +564,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.c4_size = c4_size self.c4_logical_size = c4_logical_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 c128_ring_size = self.get_ring_size(128) if ONLINE_C128: @@ -629,9 +624,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ) 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: self.swa_kv_pool = None @@ -1058,37 +1050,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): assert self.online_c128_mtp_pending_seq_lens is not None 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: """Reset request-scoped C128 state for one req slot.""" for pool in self.compress_state_pools: @@ -1115,13 +1076,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): accept_lens: torch.Tensor, num_draft_tokens: int, ) -> None: - """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. - """ + """Clear offline C128 ring slots written for rejected speculative tokens.""" if ONLINE_C128 or num_draft_tokens <= 1 or req_pool_indices.numel() == 0: return diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index a4077d892..1c8df4911 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -204,7 +204,9 @@ if TYPE_CHECKING: from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( 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): @@ -334,35 +336,6 @@ class KVCacheConfigurator: 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( f"Memory pool end. " 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_running_requests=sizes.max_running_requests, 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, token_to_kv_pool=pools.token_to_kv_pool, token_to_kv_pool_allocator=pools.token_to_kv_pool_allocator, @@ -1033,7 +1006,9 @@ class KVCacheConfigurator: extra_max_context_len: int, pre_alloc_size: int, ) -> ReqToTokenPool: - from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool + from sglang.srt.disaggregation.decode import ( + HybridMambaDecodeReqToTokenPool, + ) req_to_token_pool = HybridMambaDecodeReqToTokenPool( size=max_num_reqs, @@ -1321,7 +1296,9 @@ class KVCacheConfigurator: assert swa_page_size == 256, "In paged swa mode, page_size must be 256." 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 = [ COMPRESS_RATIO_NEXTN_LAYER @@ -1436,7 +1413,9 @@ class KVCacheConfigurator: full_max_total_num_tokens: Optional[int], swa_max_total_num_tokens: Optional[int], ) -> KVCache: - from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool + from sglang.srt.hardware_backend.npu.memory_pool_npu import ( + NPUMHATokenToKVPool, + ) kwargs = {} if self.is_hybrid_swa_compress: @@ -1502,7 +1481,9 @@ class KVCacheConfigurator: def _build_ascend_mla_kv_pool( self, *, max_total_num_tokens: int, is_dsa_model: bool ) -> 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( max_total_num_tokens, @@ -1520,7 +1501,9 @@ class KVCacheConfigurator: return token_to_kv_pool 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( max_total_num_tokens, @@ -1960,11 +1943,12 @@ class KVCacheConfigurator: device=self.device, kvcache=token_to_kv_pool, need_sort=need_sort, - req_to_token_pool=req_to_token_pool, ) else: 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() token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index e21b61f56..3a5e7d3da 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -408,7 +408,9 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): 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( 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 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 - 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). + decode reserves a draft worker, mirroring dflash's cell_size scaling); bias = 0. """ 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_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: # Ring is sized once here, so it must serve the largest adaptive tier. self._assert_ring_serves_draft_tokens( @@ -940,8 +918,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): # bytes_per_full_token: tokens = avail / (bpft * (T+D)/T). draft_layers = 1 target_layers = self.num_layers_total - self._spec_infl = (target_layers + draft_layers) / target_layers - self.bytes_per_full_token *= self._spec_infl + self.bytes_per_full_token *= (target_layers + draft_layers) / target_layers # Online c128 keeps a single in-progress (max, sum, kv) state per index # and assumes a strict forward-only schedule. Speculative decode (MTP) @@ -994,11 +971,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): ) 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 c4_state_dtype_size, c128_state_dtype_size = ( @@ -1022,40 +995,16 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): c4_frac = 1 / (4 * self.c4_shrink_factor) return ( - # Unified_kv: SWA is a fixed per-request ring (see _fixed_swa_bytes), - # 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 - ) + self.swa_ratio * kv_bytes * self.num_layers_total + c4_frac * kv_bytes * self.num_layers_ca4 + 1 / 128 * kv_bytes * self.num_layers_ca128 + 1 / 4 * self.indexer_bytes_per_token * self.num_layers_ca4 - # Unified_kv: the c4 (attn + indexer) compress-state is a ring buffer - # 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 - ) + + self.swa_ratio * c4_state_ratio * c4_state_bytes * self.num_layers_ca4 + c128_state_ratio * c128_state_bytes * self.num_layers_ca128 - + ( - 0.0 - if self._unified - else self.swa_ratio - * c4_state_ratio - * c4_indexer_state_bytes - * self.num_layers_ca4 - ) + + self.swa_ratio + * c4_state_ratio + * c4_indexer_state_bytes + * self.num_layers_ca4 ) def _compute_dsv4_sizes(self, full_token: int, page_size: int) -> _DSV4PoolSizes: @@ -1067,14 +1016,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): swa_max_total_num_tokens=swa_tokens, c4_max_total_num_tokens=full_token // (4 * self.c4_shrink_factor), c128_max_total_num_tokens=full_token // 128, - # Unified_kv sizes the c4 state ring from the fixed SWA ring - # (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 - ), + c4_state_pool_size=swa_tokens // self.swa_page_size * self.c4_ring_size, 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 ) - def _unified_c4_state_pool_size(self, max_running_requests: int) -> int: - """Exact request-scoped C4 ring size for the unified address contract. - - 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.""" + def _get_c128_state_fixed_bytes_for_token_capacity( + self, token_capacity: int + ) -> int: 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 = int(token_capacity / 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) + 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: full = sizes.full_max_total_num_tokens @@ -1193,13 +1089,6 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): config.c128_state_pool_size = num_req_slots else: 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 def calculate_pool_sizes( @@ -1209,34 +1098,25 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): "page_size must be multiple of 128 for compressed attention" ) - max_running_requests_per_worker = self._resolve_max_running_requests_per_worker( - available_bytes - ) - c128_state_fixed_bytes = self._get_c128_state_fixed_bytes( - max_running_requests_per_worker - ) - 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 - ) + if self.requested_max_running_requests_per_worker is not None: + c128_state_fixed_bytes = self._get_c128_state_fixed_bytes( + self.requested_max_running_requests_per_worker + ) + else: + full_token = int(available_bytes / self.bytes_per_full_token) + c128_state_fixed_bytes = ( + self._get_c128_state_fixed_bytes_for_token_capacity(full_token) + ) - available_bytes_for_tokens = max( - available_bytes - - c128_state_fixed_bytes - - swa_ring_fixed_bytes - - c4_state_fixed_bytes, - 0, - ) + available_bytes_for_tokens = max(available_bytes - c128_state_fixed_bytes, 0) full_token = int(available_bytes_for_tokens / self.bytes_per_full_token) sizes = self._compute_dsv4_sizes(full_token, page_size) 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"available_bytes={available_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}" ) return self._to_config(sizes) diff --git a/test/registered/kernels/ops/attention/test_c4_v2.py b/test/registered/kernels/ops/attention/test_c4_v2.py index cd419ef76..f05aaa903 100644 --- a/test/registered/kernels/ops/attention/test_c4_v2.py +++ b/test/registered/kernels/ops/attention/test_c4_v2.py @@ -7,11 +7,7 @@ import pytest import torch import triton -from sglang.kernels.ops.attention.dsv4 import ( - CompressorDecodePlan, - CompressorPrefillPlan, - compress_forward, -) +from sglang.kernels.ops.attention.dsv4 import compress_forward from sglang.srt.utils import get_device from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci 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("seq_len", [4, 8, 32, 256, 1024]) def test_prefill_no_context(mode: str, seq_len: int) -> None: diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index bd60d6a68..d74db97d6 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -9,7 +9,10 @@ from sglang.srt.managers.schedule_policy import ( PrefillAdder, 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.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.utils.common import Range @@ -68,13 +71,6 @@ class TestPrefillAdder(CustomTestCase): allocator.swa_available_size.return_value = swa_available_size allocator.available_size.return_value = available_size 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 def create_running_batch(self, reqs=None) -> MagicMock: diff --git a/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py index 54440519b..6152e59cc 100644 --- a/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py +++ b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py @@ -23,9 +23,6 @@ class _FakeAllocator: self.alloc_calls = [] self.extend_calls = [] - def get_kvcache(self): - return None - def available_size(self): return 1 << 30 diff --git a/test/registered/unit/mem_cache/test_dsv4_c4_state_lifecycle.py b/test/registered/unit/mem_cache/test_dsv4_c4_state_lifecycle.py deleted file mode 100644 index 24ec0e83d..000000000 --- a/test/registered/unit/mem_cache/test_dsv4_c4_state_lifecycle.py +++ /dev/null @@ -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() diff --git a/test/registered/unit/mem_cache/test_hisparse_allocator.py b/test/registered/unit/mem_cache/test_hisparse_allocator.py index ebc173a4c..7ea10c1f4 100644 --- a/test/registered/unit/mem_cache/test_hisparse_allocator.py +++ b/test/registered/unit/mem_cache/test_hisparse_allocator.py @@ -131,7 +131,6 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase): queue = DecodePreallocQueue.__new__(DecodePreallocQueue) queue.req_to_token_pool = req_to_token_pool queue.token_to_kv_pool_allocator = allocator - queue.token_to_kv_pool = None queue.tree_cache = SimpleNamespace( evictable_size=MagicMock(return_value=0), protected_size=MagicMock(return_value=0), diff --git a/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py b/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py index e3de143f2..7cef341de 100644 --- a/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py +++ b/test/registered/unit/mem_cache/test_swa_alloc_extend_page_estimation.py @@ -36,11 +36,6 @@ def _make_self(*, page_size: int, full_available: int, swa_available: int): return SimpleNamespace( 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( available_size=lambda: full_available, alloc_extend=MagicMock(return_value=full_indices),