Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 0 additions & 43 deletions python/sglang/srt/mem_cache/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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():
Expand All @@ -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


Expand Down
32 changes: 11 additions & 21 deletions python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
76 changes: 0 additions & 76 deletions test/registered/unit/mem_cache/test_dsv4_c128_req_state.py

This file was deleted.

Loading