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
11 changes: 5 additions & 6 deletions python/sglang/srt/mem_cache/base_prefix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -455,12 +455,11 @@ def rotation_base_of(self, node: Any) -> Optional[int]:
return None

@abstractmethod
def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs):
"""Hand a finished request's KV to the tree: insert what can be keyed
(advancing ``cache_protected_len``), ``free_kv_row`` the rest of
``[cache_protected_len, owned_kv_len)``, ``unpin``. Slicing the row by
token count instead strands the slots up to ``owned_kv_len``; the
caller frees everything past it."""
def insert_req(self, req: Req, *, up_to: int, **kwargs):
"""Hand a finished request's KV up to row position ``up_to`` to the
tree: insert what can be keyed and advance ``cache_protected_len``
past it. The caller then frees ``[cache_protected_len, up_to)`` and
everything after, and unpins; nothing here releases a slot."""

@abstractmethod
def cache_unfinished_req(self, req: Req, **kwargs):
Expand Down
7 changes: 2 additions & 5 deletions python/sglang/srt/mem_cache/chunk_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,11 +76,8 @@ def insert(self, params: InsertParams) -> InsertResult:
# ChunkCache does not support prefix caching, so insert is a no-op
return InsertResult(prefix_len=0)

def cache_finished_req(self, req: Req, *, owned_kv_len: int):
# For decode server: if req.output_ids is empty, we want to free all req.origin_input_ids
# The protected prefix is not this req's to free.
self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
self.unpin(req)
def insert_req(self, req: Req, *, up_to: int):
pass

def cache_unfinished_req(self, req: Req, chunked=False):
kv_indices = self.req_to_token_pool.req_to_token[
Expand Down
11 changes: 5 additions & 6 deletions python/sglang/srt/mem_cache/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,11 +313,10 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
owned_kv_len = req.owned_kv_len()
is_insert = is_insert and not req.skip_radix_cache_insert
if is_insert:
tree_cache.cache_finished_req(req, owned_kv_len=owned_kv_len)
else:
# The protected prefix is not this req's to free.
tree_cache.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
tree_cache.unpin(req)
tree_cache.insert_req(req, up_to=owned_kv_len)
# The protected prefix is not this req's to free.
tree_cache.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
tree_cache.unpin(req)
_release_overallocated_kv_indices(
req, owned_kv_len, req.kv.kv_allocated_len, tree_cache
)
Expand Down Expand Up @@ -360,7 +359,7 @@ def _release_overallocated_kv_indices(

if start_p < end_p:
# start_p is aligned to the allocator's page above, so it never shares a
# page with cache_finished_req's tail free in this group.
# page with the tail free_kv_row in this group.
tree_cache.free_kv_row(req.kv, [(start_p, end_p)])


Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/mem_cache/memory_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -1036,7 +1036,7 @@ def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor):
(``write_pos[src] == 0``). Only ``temporal`` is copied, not the ring, so
an un-flushed source would drop its last ``write_pos`` updates. Callers
comply: COW copies radix checkpoints; ``cache_unfinished_req`` copies an
active slot only during prefill (ring empty); ``cache_finished_req``
active slot only during prefill (ring empty); ``insert_req``
caps the donate to the last flush boundary. The dst cursor is reset to 0
(the copied checkpoint has no pending ring entries).
"""
Expand Down
27 changes: 11 additions & 16 deletions python/sglang/srt/mem_cache/pure_swa_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,22 +57,17 @@ def evict(self, params: EvictParams) -> EvictResult:
num_tokens = max(params.num_tokens, params.swa_num_tokens)
return super().evict(EvictParams(num_tokens=num_tokens))

def cache_finished_req(self, req: Req, *, owned_kv_len: int):
"""Insert only the prefill portion [0, evict_floor); free_kv_row skips
the span _evict_swa already freed during decode."""
if not self.disable:
token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len]
swa_evict_floor = req.kv.swa_evict_floor
key_limit = (
ceil_align(swa_evict_floor, self.page_size)
if swa_evict_floor > 0
else None
)
radix_key, _, _ = self._adopt(req, token_ids, key_limit=key_limit)
req.kv.cache_protected_len = len(radix_key)
# The protected prefix is not this req's to free.
self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
self.unpin(req)
def insert_req(self, req: Req, *, up_to: int):
"""Insert only the prefill portion [0, swa_evict_floor)."""
if self.disable:
return
token_ids = (req.origin_input_ids + req.output_ids)[:up_to]
swa_evict_floor = req.kv.swa_evict_floor
key_limit = (
ceil_align(swa_evict_floor, self.page_size) if swa_evict_floor > 0 else None
)
radix_key, _, _ = self._insert_cache(req, token_ids, key_limit=key_limit)
req.kv.cache_protected_len = len(radix_key)

def cache_unfinished_req(self, req: Req, chunked=False):
"""During chunked prefill, swa_evicted_seqlen is 0 and no SWA eviction
Expand Down
28 changes: 13 additions & 15 deletions python/sglang/srt/mem_cache/radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -433,7 +433,7 @@ def insert(self, params: InsertParams) -> InsertResult:
)
return InsertResult(prefix_len=prefix_len, last_device_node=last_node)

def _adopt(
def _insert_cache(
self,
req: Req,
token_ids,
Expand Down Expand Up @@ -492,22 +492,21 @@ def _adopt(
)
return radix_key, kv_indices, result.prefix_len

def cache_finished_req(self, req: Req, *, owned_kv_len: int):
"""Cache request when it finishes."""
if not self.disable:
token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len]
radix_key, _, _ = self._adopt(req, token_ids, split_prompt=True)
req.kv.cache_protected_len = len(radix_key)
# The protected prefix is not this req's to free.
self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
self.unpin(req)
def insert_req(self, req: Req, *, up_to: int):
if self.disable:
return
token_ids = (req.origin_input_ids + req.output_ids)[:up_to]
radix_key, _, _ = self._insert_cache(req, token_ids, split_prompt=True)
req.kv.cache_protected_len = len(radix_key)

def cache_unfinished_req(self, req: Req, chunked=False):
"""Cache request when it is unfinished."""
if self.disable:
return

radix_key, kv_indices, _ = self._adopt(req, req.get_fill_ids(), chunked=chunked)
radix_key, kv_indices, _ = self._insert_cache(
req, req.get_fill_ids(), chunked=chunked
)

# The prefix indices could be updated, reuse it
match_result = self.match_prefix(MatchPrefixParams(key=radix_key))
Expand All @@ -524,10 +523,9 @@ def cache_unfinished_req(self, req: Req, chunked=False):
new_indices[req.kv.cache_protected_len :],
)

# The cache_protected_len is not always equal to len(req.prefix_indices)
# since for page_size > 1, the partial part is added to req.prefix_indices, but that part of kv indices is not added to the tree.
# It should be freed in the next cache_unfinished_req and final cache_finished_req to avoid memory leak.
# So we introduce this `cache_protected_len` field to make sure the partial part can be freed correctly.
# With page_size > 1 the partial page sits in req.prefix_indices but not
# in the tree; cache_protected_len marks the tree-owned part so the next
# cache_unfinished_req or release_kv_cache frees the rest.
req.kv.cache_protected_len = len(new_indices)

self.dec_lock_ref(req.last_node)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -382,18 +382,16 @@ def _allocate_and_load(
return fetched_slots, new_node

# ------------------------------------------------------------------
# cache_finished_req (STORE)
# insert_req (STORE)
# ------------------------------------------------------------------

def on_release(self, req: Req, *, inserted: bool) -> None:
if not inserted:
self._load_markers.pop(req.cache_request_handle, None)

def cache_finished_req( # type: ignore[override]
self, req: Req, *, owned_kv_len: int
) -> None:
"""Base cache_finished_req then fire an async FlexKV store."""
super().cache_finished_req(req, owned_kv_len=owned_kv_len)
def insert_req(self, req: Req, *, up_to: int) -> None: # type: ignore[override]
"""Base insert_req then fire an async FlexKV store."""
super().insert_req(req, up_to=up_to)

# Compute the committed prefix.
topk = get_spec().speculative_eagle_topk
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -259,11 +259,11 @@ def on_release(self, req: Req, *, inserted: bool) -> None:
if not inserted:
self.release_aborted_request(req.cache_request_handle)

def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs) -> None:
self._publish_external_loaded_prefix(req, token_ids_len=owned_kv_len)
super().cache_finished_req(req, owned_kv_len=owned_kv_len, **kwargs)
def insert_req(self, req: Req, *, up_to: int, **kwargs) -> None:
self._publish_external_loaded_prefix(req, token_ids_len=up_to)
super().insert_req(req, up_to=up_to, **kwargs)
self._retire_loaded_flow(req.rid)
token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len]
token_ids = (req.origin_input_ids + req.output_ids)[:up_to]
self._submit_store(req, token_ids)
self._request_session_finish(req.rid)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -228,16 +228,16 @@ receipt proves were taken. The eventual full release must pass

---

### `cache_finished_req(req: Req, *, owned_kv_len: int)`
### `insert_req(req: Req, *, up_to: int)`

Cache a completed request's KV data into the tree.

| Aspect | Detail |
|--------|--------|
| **Purpose** | After a request finishes, insert its token/KV data into the tree for future reuse |
| **Inputs** | `req` — the finished request; `owned_kv_len` — end of the request-owned KV range; slots past it are freed by `release_kv_cache`. A request that leaves without inserting goes through `release_kv_cache(is_insert=False)` instead, which frees the row and calls `on_release(req, inserted=False)` for component cleanup |
| **Inputs** | `req` — the finished request; `up_to` — end of the request-owned KV range. `release_kv_cache` frees `[cache_protected_len, up_to)` and everything past it and unpins; with `is_insert=False` it skips the insert and calls `on_release(req, inserted=False)` for component cleanup |
| **Output** | `None` |
| **Mutation** | Calls component hooks → `insert` → `dec_lock_ref` → component cleanup. Frees unaligned tail KV indices. |
| **Mutation** | Calls component hooks → `insert` → component cleanup; advances `cache_protected_len` past the inserted key. Frees nothing and drops no lock: `release_kv_cache` does both afterwards. |
| **Complexity** | **O(K + D·C)** — insert O(K + D·C) + lock release O(D). Simplifies to **O(K)**. |

**Algorithm detail:**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional

from sglang.srt.mem_cache.unified_cache.unified_tree_core import NodeId

if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.unified_cache.components.base import (
Expand Down Expand Up @@ -51,8 +53,9 @@ def session_id_for_req(self, req: Req) -> Optional[str]:
session_id = req.session.session_id
return session_id

def register_session_ref(self, req: Req) -> None:
"""Register a non-streaming request's reusable leaves with each component."""
def register_session_ref(self, req: Req, leaf: NodeId) -> None:
"""Register the leaf a finished request's insert ended on with each
component; the lock anchor ``req.last_node`` is a different node."""
if not self.enable_session_radix_cache:
return

Expand All @@ -69,14 +72,13 @@ def register_session_ref(self, req: Req) -> None:
logger.warning("register_session_ref called for stale request; Skip it.")
return

assert req.last_node is not None
last_node = self.tree_core.node_by_id(req.last_node)
if last_node is self.tree_core.root_node:
node = self.tree_core.node_by_id(leaf)
if node is self.tree_core.root_node:
return

for component in self.components:
leaf = component.resolve_session_leaf(req, last_node)
component.register_session_leaf(session_id, leaf)
component_leaf = component.resolve_session_leaf(req, node)
component.register_session_leaf(session_id, component_leaf)

def _remember_closed_session(self, session_id: str) -> None:
self._closed_session_ids[session_id] = None
Expand Down
29 changes: 12 additions & 17 deletions python/sglang/srt/mem_cache/unified_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -975,18 +975,15 @@ def on_release(self, req: Req, *, inserted: bool) -> None:
for comp in self._components_tuple:
comp.cleanup_after_caching_req(req, is_finished=True)

@rank_consensus(same_params=["req.rid", "owned_kv_len"])
def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs) -> None:
@rank_consensus(same_params=["req.rid", "up_to"])
def insert_req(self, req: Req, *, up_to: int, **kwargs) -> None:
if self.disable:
self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)])
for comp in self._components_tuple:
comp.cleanup_after_caching_req(req, is_finished=True)
return

token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len]
kv_indices = self.req_to_token_pool.req_to_token[
req.kv.req_pool_idx, :owned_kv_len
]
token_ids = (req.origin_input_ids + req.output_ids)[:up_to]
kv_indices = self.req_to_token_pool.req_to_token[req.kv.req_pool_idx, :up_to]

insert_params = InsertParams(
prev_prefix_len=req.kv.cache_protected_len,
Expand Down Expand Up @@ -1060,8 +1057,9 @@ def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs) -> None:
)
)

# Everything past the inserted key goes back, the protected prefix
# never does. After a rotation decline nothing was inserted.
# Everything past the inserted key goes back to the caller, the
# protected prefix never does. After a rotation decline nothing was
# inserted.
free_from = (
min(req.kv.cache_protected_len, len(kv_indices))
if result.rotation_tail_declined
Expand All @@ -1074,12 +1072,7 @@ def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs) -> None:
f"{free_from=} {len(kv_indices)=} {req.kv.cache_protected_len=}"
)
free_from = tail_free_start
self.free_kv_row(req.kv, [(free_from, len(kv_indices_full))])

self.unpin(req)

if result is not None and result.last_device_node is not None:
req.last_node = result.last_device_node
req.kv.cache_protected_len = free_from

# cleanup
for comp in self._components_tuple:
Expand All @@ -1093,7 +1086,9 @@ def cache_finished_req(self, req: Req, *, owned_kv_len: int, **kwargs) -> None:
if req.finished_reason is not None and not isinstance(
req.finished_reason, FINISH_ABORT
):
self.session_refs.register_session_ref(req)
self.session_refs.register_session_ref(
req, leaf=result.last_device_node
)

@rank_consensus(same_params=["req.rid", "chunked"])
def advance_unpublished_req(self, req: Req, chunked: bool = False) -> None:
Expand Down Expand Up @@ -1217,7 +1212,7 @@ def cache_unfinished_req(self, req: Req, chunked: bool = False, **kwargs) -> Non
# gather contract forbids -- keep the request entirely on its own
# pages: no dedup free, no rebind, no protection change. The insert
# declined before its walk, so nothing was freed underneath us. The
# final cache_finished_req releases everything past the protected
# release_kv_cache releases everything past the protected
# prefix.
req.prefix_indices = kv_indices_orig.to(dtype=torch.int64, copy=True)
for comp in self._components_tuple:
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/session/session_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -423,7 +423,7 @@ def _close(self, session_id: str):
# An in-flight request is still decoding on this session's KV
# memory. Freeing now would corrupt the scheduler. Mark the
# session for deferred cleanup: the request keeps its session
# reference so cache_finished_req takes the streaming path,
# reference so release_kv_cache takes the streaming path,
# and we schedule release_session for after it completes.
session.close_on_finish = True
logger.info(
Expand Down
6 changes: 3 additions & 3 deletions python/sglang/srt/session/streaming_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,7 +174,7 @@ def find_active_slot(self, req: Req) -> Optional[SessionSlot]:
"""Returns an active slot for this req, or None.

Side effect: if req is pre-aborted (to_finish set, e.g. input too
long), detach it from the session so cache_finished_req treats it
long), detach it from the session so release_kv_cache treats it
as a normal req. The slot stays intact for the next request.
"""
if not _is_streaming(req):
Expand Down Expand Up @@ -359,8 +359,8 @@ def claim_kv_row(self, req: Req) -> bool:
def on_release(self, req: Req, *, inserted: bool) -> None:
self.inner.on_release(req, inserted=inserted)

def cache_finished_req(self, req: Req, **kwargs):
self.inner.cache_finished_req(req, **kwargs)
def insert_req(self, req: Req, **kwargs):
self.inner.insert_req(req, **kwargs)

def cache_unfinished_req(self, req: Req, **kwargs):
if self.try_cache_unfinished_req(req, **kwargs):
Expand Down
8 changes: 8 additions & 0 deletions python/sglang/test/mem_cache_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
"""Shared helpers for mem_cache unit tests."""


def finish_req(cache, req, up_to):
"""insert_req, then what release_kv_cache does after it: free the rest, drop the lock."""
cache.insert_req(req, up_to=up_to)
cache.free_kv_row(req.kv, [(req.kv.cache_protected_len, up_to)])
cache.unpin(req)
Loading
Loading