From a797ca14083fdfdb83a1fa3ea5c63b7228932728 Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Sat, 27 Jun 2026 23:15:11 +0800 Subject: [PATCH] Remove C128 req slot allocation reset --- python/sglang/srt/mem_cache/common.py | 43 ----------- .../srt/mem_cache/deepseek_v4_memory_pool.py | 32 +++----- .../mem_cache/test_dsv4_c128_req_state.py | 76 ------------------- 3 files changed, 11 insertions(+), 140 deletions(-) delete mode 100644 test/registered/unit/mem_cache/test_dsv4_c128_req_state.py diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 58a092acdf00..0a8c83aa6c1f 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -49,47 +49,6 @@ logger = logging.getLogger(__name__) -def _clear_c128_req_state_for_new_slots( - tree_cache: BasePrefixCache | None, - reqs: list[Req], - already_allocated: list[bool], -) -> None: - # C128_STATE is request-scoped runtime state indexed by req_pool_idx. When a - # req slot is reused for a new request, stale partial C128 state from the - # previous request must be cleared before the new request starts using it. - # This is intentionally tied to req slot allocation instead of radix-cache - # hits: radix/HiCache hits are page-aligned and reuse complete C128 KV, while - # chunked continuation reuses an existing req slot and must keep its live - # partial state. - if tree_cache is None: - return - - allocator = getattr(tree_cache, "token_to_kv_pool_allocator", None) - if allocator is None: - return - - get_kvcache = getattr(allocator, "get_kvcache", None) - if get_kvcache is None: - return - - kv_pool = get_kvcache() - clear_c128_req_states = getattr(kv_pool, "clear_c128_req_states", None) - clear_c128_req_state = getattr(kv_pool, "clear_c128_req_state", None) - if clear_c128_req_states is None and clear_c128_req_state is None: - return - - new_req_pool_indices = [ - int(req.req_pool_idx) - for req, reused in zip(reqs, already_allocated) - if not reused and req.req_pool_idx is not None - ] - if clear_c128_req_states is not None: - clear_c128_req_states(new_req_pool_indices) - else: - for req_pool_idx in new_req_pool_indices: - clear_c128_req_state(req_pool_idx) - - def kv_to_page_indices(kv_indices: np.ndarray, page_size: int): # The page is guaranteed to be full except the last page. if page_size == 1: @@ -455,7 +414,6 @@ def alloc_req_slots( ) -> list[int]: """Allocate request slots from the pool.""" num_reqs = len(reqs) - already_allocated = [req.req_pool_idx is not None for req in reqs] if isinstance(req_to_token_pool, HybridReqToTokenPool): mamba_available_size = req_to_token_pool.mamba_allocator.available_size() if tree_cache.supports_mamba(): @@ -480,7 +438,6 @@ def alloc_req_slots( f"{req_to_token_pool.available_size()=}, " f"{num_reqs=}, " ) - _clear_c128_req_state_for_new_slots(tree_cache, reqs, already_allocated) return req_pool_indices 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 9840f2be74ec..f691f564419c 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -946,34 +946,24 @@ def get_online_c128_mtp_pending_seq_lens(self) -> torch.Tensor: assert self.online_c128_mtp_pending_seq_lens is not None return self.online_c128_mtp_pending_seq_lens - def clear_c128_req_states(self, req_pool_indices: List[int]) -> None: - """Reset request-scoped C128 states for newly allocated req slots.""" - if not req_pool_indices: - return - + 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: if pool is None or pool.ratio != 128: continue state = pool.kv_score_buffer.kv_score - req_indices = torch.tensor( - req_pool_indices, dtype=torch.long, device=state.device - ) if ONLINE_C128: - head_dim = state.shape[-1] // 3 - state[req_indices, :head_dim] = float("-inf") - state[req_indices, head_dim:] = 0 + row = state[req_pool_idx] + head_dim = row.shape[-1] // 3 + row[:head_dim].fill_(float("-inf")) + row[head_dim:].zero_() else: - offsets = torch.arange(pool.ring_size, device=state.device) - row_indices = ( - req_indices[:, None] * pool.ring_size + offsets[None, :] - ).reshape(-1) - half = state.shape[-1] // 2 - state[row_indices, :half] = 0 - state[row_indices, half:] = float("-inf") - - def clear_c128_req_state(self, req_pool_idx: int) -> None: - self.clear_c128_req_states([req_pool_idx]) + start = req_pool_idx * pool.ring_size + rows = state[start : start + pool.ring_size] + half = rows.shape[-1] // 2 + rows[:, :half].zero_() + rows[:, half:].fill_(float("-inf")) def clear_unaccepted_c128_draft_states( self, diff --git a/test/registered/unit/mem_cache/test_dsv4_c128_req_state.py b/test/registered/unit/mem_cache/test_dsv4_c128_req_state.py deleted file mode 100644 index 322b16eb6064..000000000000 --- a/test/registered/unit/mem_cache/test_dsv4_c128_req_state.py +++ /dev/null @@ -1,76 +0,0 @@ -import unittest - -from sglang.srt.mem_cache.common import alloc_req_slots -from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=5, suite="base-b-test-cpu") - - -class FakeKVPool: - def __init__(self): - self.cleared = [] - - def clear_c128_req_states(self, req_pool_indices): - self.cleared.append(list(req_pool_indices)) - - -class FakeAllocator: - def __init__(self, kv_pool): - self.kv_pool = kv_pool - - def get_kvcache(self): - return self.kv_pool - - -class FakeTreeCache: - def __init__(self, kv_pool): - self.token_to_kv_pool_allocator = FakeAllocator(kv_pool) - - -class FakeReq: - def __init__(self, req_pool_idx=None): - self.req_pool_idx = req_pool_idx - self.inflight_middle_chunks = 0 - self.kv_committed_len = 0 - - -class TestDSV4C128ReqState(unittest.TestCase): - def _make_pool(self): - return ReqToTokenPool( - size=4, - max_context_len=16, - device="cpu", - enable_memory_saver=False, - ) - - def test_new_req_slots_clear_c128_state(self): - req_to_token_pool = self._make_pool() - kv_pool = FakeKVPool() - reqs = [FakeReq(), FakeReq()] - - req_pool_indices = alloc_req_slots( - req_to_token_pool, reqs, FakeTreeCache(kv_pool) - ) - - self.assertEqual(req_pool_indices, [1, 2]) - self.assertEqual(kv_pool.cleared, [[1, 2]]) - - def test_reused_req_slot_skips_c128_state_clear(self): - req_to_token_pool = self._make_pool() - req_to_token_pool.free_slots.remove(3) - kv_pool = FakeKVPool() - reused_req = FakeReq(req_pool_idx=3) - reused_req.kv_committed_len = 1 - new_req = FakeReq() - - req_pool_indices = alloc_req_slots( - req_to_token_pool, [reused_req, new_req], FakeTreeCache(kv_pool) - ) - - self.assertEqual(req_pool_indices, [3, 1]) - self.assertEqual(kv_pool.cleared, [[1]]) - - -if __name__ == "__main__": - unittest.main()