Skip to content
Open
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
5 changes: 5 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -691,6 +691,11 @@ class Envs:
# - Source builds with a missing or unusable Rust toolchain.
# This also applies when Rust is explicitly selected.
SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("rust")
# Python TreeCore only: keep a persistent lazy-deletion heap over the Full
# component's evictable leaves instead of rebuilding it on every eviction
# call. False re-keys every leaf at each eviction (legacy O(#leaves) cost,
# identical eviction order) through the same code path.
SGLANG_UNIFIED_RADIX_LAZY_EVICTION_HEAP = EnvBool(True)
# Once decode passes the sliding window, drop the SWA part of the prefill's
# tree lock; its SWA KV becomes evictable rather than freed.
SGLANG_OPT_RELEASE_PREFILL_SWA = EnvBoolWithAlias(
Expand Down
80 changes: 35 additions & 45 deletions python/sglang/srt/mem_cache/unified_cache/components/full.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
from __future__ import annotations

import heapq
from typing import TYPE_CHECKING, Callable, Optional, Sequence

import torch
Expand Down Expand Up @@ -67,6 +66,7 @@ def _dec_session_coverage(self, session_id: str, leaf: UnifiedTreeNode) -> None:
cd = node.component_data[self.component_type]
assert cd.session_ref > 0
cd.session_ref -= 1
self.tree_core._touch_full_eviction_key(node)
node = node.parent

def _advance_session_coverage(
Expand All @@ -83,6 +83,7 @@ def _advance_session_coverage(
and node is not self.tree_core.root_node
):
node.component_data[self.component_type].session_ref += 1
self.tree_core._touch_full_eviction_key(node)
node = node.parent

def _recede_session_coverage(
Expand All @@ -101,6 +102,7 @@ def _recede_session_coverage(
cd = node.component_data[self.component_type]
assert cd.session_ref > 0
cd.session_ref -= 1
self.tree_core._touch_full_eviction_key(node)
node = node.parent

def create_match_validator(
Expand Down Expand Up @@ -200,14 +202,9 @@ def _session_ref_eviction_strategy(self, node: UnifiedTreeNode):
return ref > 0, ref, self.tree_core.eviction_strategy.get_priority(node)

def _evict_device_start(self, request_cnt: int) -> None:
self._ensure_eviction_strategy()
self._evict_device_request_cnt = request_cnt
self._evict_device_last_node = None
self._evict_device_heap = [
(self.session_ref_eviction_strategy(n), n)
for n in self.tree_core.evictable_device_leaves
]
heapq.heapify(self._evict_device_heap)
self.tree_core.full_device_heap.begin_walk()

def _evict_device_next_node(
self,
Expand All @@ -216,28 +213,23 @@ def _evict_device_next_node(
host_frees: dict[ComponentType, list[torch.Tensor]],
) -> Optional[NodeId]:
ct = self.component_type
heap = self.tree_core.full_device_heap
lv = self._evict_device_last_node
if (
lv is not None
and lv.parent is not None
and lv.parent in self.tree_core.evictable_device_leaves
):
heapq.heappush(
self._evict_device_heap,
(self.session_ref_eviction_strategy(lv.parent), lv.parent),
)
if lv is not None and lv.parent is not None:
# The evicted leaf's parent may have become a device leaf: admit
# it to the current walk, keyed now (the legacy explicit push).
heap.promote(lv.parent)
self._evict_device_last_node = None
while tracker[ct] < self._evict_device_request_cnt and self._evict_device_heap:
_, x = heapq.heappop(self._evict_device_heap)
if x not in self.tree_core.evictable_device_leaves:
continue
self._evict_device_last_node = x
return x.id
if tracker[ct] < self._evict_device_request_cnt:
x = heap.pop_next()
if x is not None:
self._evict_device_last_node = x
return x.id
return None

def _evict_device_end(self) -> None:
self._evict_device_heap = []
self._evict_device_last_node = None
self.tree_core.full_device_heap.end_walk()

def drive_host_eviction(
self,
Expand All @@ -247,26 +239,19 @@ def drive_host_eviction(
host_frees: dict[ComponentType, list[torch.Tensor]],
) -> None:
"""Evict host leaves to free KV host pool space."""
self._ensure_eviction_strategy()
heap = [
(self.session_ref_eviction_strategy(n), n)
for n in self.tree_core.evictable_host_leaves
]
heapq.heapify(heap)
ct = self.component_type
while tracker[ct] < num_tokens and heap:
_, x = heapq.heappop(heap)
if x not in self.tree_core.evictable_host_leaves:
continue
self.tree_core._evict_host_leaf(x, tracker, device_frees, host_frees)
if (
x.parent is not None
and x.parent in self.tree_core.evictable_host_leaves
):
heapq.heappush(
heap,
(self.session_ref_eviction_strategy(x.parent), x.parent),
)
heap = self.tree_core.full_host_heap
heap.begin_walk()
try:
while tracker[ct] < num_tokens:
x = heap.pop_next()
if x is None:
break
self.tree_core._evict_host_leaf(x, tracker, device_frees, host_frees)
if x.parent is not None:
heap.promote(x.parent)
finally:
heap.end_walk()

def acquire_component_lock(
self,
Expand Down Expand Up @@ -301,18 +286,23 @@ def acquire_component_lock(

# Lock the device-on segment up to root
delta = 0
tree_core = self.tree_core
evictable_device_leaves = tree_core.evictable_device_leaves
device_live = tree_core._full_device_live
while cur is not root:
cd = cur.component_data[ct]
assert cd.value is not None, (
f"FULL invariant broken: evicted ancestor {cur.id} above device-on segment"
)
if cd.lock_ref == 0:
key_len = len(cd.value)
self.tree_core.component_evictable_size_[ct] -= key_len
self.tree_core.component_protected_size_[ct] += key_len
tree_core.component_evictable_size_[ct] -= key_len
tree_core.component_protected_size_[ct] += key_len
delta += key_len
cd.lock_ref += 1
self.tree_core.evictable_device_leaves.discard(cur)
evictable_device_leaves.discard(cur)
if cur in device_live:
tree_core.full_device_heap.forget(cur)
cur = cur.parent
result.delta = delta
return result
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -241,10 +241,12 @@ def commit_insert_component_data(
node.id, self.component_type, params.mamba_value
)
node.last_access_time = get_and_increase_time_counter()
self.tree_core._touch_full_eviction_key(node)
self._emit_excess_path_states_eviction(node, cache_actions)
return
self.tree_core.lru_lists[self.component_type].reset_node_mru(node)
node.last_access_time = get_and_increase_time_counter()
self.tree_core._touch_full_eviction_key(node)
result.mamba_exist = True

def _emit_excess_path_states_eviction(
Expand Down
Loading
Loading