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
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
ComponentData,
ComponentType,
EvictLayer,
LRURefreshPhase,
TreeComponent,
get_and_increase_time_counter,
next_component_uuid,
Expand All @@ -20,6 +21,7 @@
"EvictLayer",
"FullComponent",
"CacheTransferPhase",
"LRURefreshPhase",
"MambaComponent",
"SWAComponent",
"TreeComponent",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
CacheTransferPhase,
ComponentType,
EvictLayer,
LRURefreshPhase,
TreeComponent,
next_component_uuid,
)
Expand Down Expand Up @@ -60,6 +61,31 @@ def _translate_full_to_swa(self, full_indices: torch.Tensor) -> torch.Tensor:
full_indices
)

def refresh_lru(
self,
phase: LRURefreshPhase,
node: UnifiedTreeNode,
root_node: UnifiedTreeNode,
) -> None:
match phase:
case LRURefreshPhase.WALKDOWN:
# Walk-down would refresh every visited ancestor to MRU,
# but most are outside the active sliding window and must
# stay evictable. Window-bounded refresh runs at
# MATCH_END / INSERT_END instead.
return
case LRURefreshPhase.MATCH_END | LRURefreshPhase.INSERT_END:
self.cache.lru_lists[
self.component_type
].reset_node_and_window_ancestors_mru(
node,
root_node,
self.sliding_window_size + self.cache.page_size,
self.node_has_component_data,
)
case _:
raise ValueError(f"Unknown LRURefreshPhase: {phase}")

def _restore_device_value(self, node: UnifiedTreeNode, value: torch.Tensor) -> None:
ct = self.component_type
node.component_data[ct].value = value
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,13 @@ class CacheTransferPhase(str, Enum):
PREFETCH = "prefetch" # Storage→H


class LRURefreshPhase(str, Enum):

WALKDOWN = "walkdown" # touching a node while walking through the tree
MATCH_END = "match_end" # end of a successful prefix match
INSERT_END = "insert_end" # after a new/updated leaf is committed


def get_and_increase_time_counter() -> float64:
global _LAST_ACCESS_TIME_COUNTER_FLOAT
ret = _LAST_ACCESS_TIME_COUNTER_FLOAT
Expand Down Expand Up @@ -115,6 +122,29 @@ def value_len(self, node: UnifiedTreeNode) -> int:
value = node.component_data[self.component_type].value
return len(value) if value is not None else 0

def refresh_lru(
self,
phase: LRURefreshPhase,
node: UnifiedTreeNode,
root_node: UnifiedTreeNode,
) -> None:
ct = self.component_type
match phase:
case LRURefreshPhase.WALKDOWN:
if node.component_data[ct].value is None:
return
self.cache.lru_lists[ct].reset_node_mru(node)
case LRURefreshPhase.MATCH_END:
self.cache.lru_lists[ct].reset_node_and_parents_mru(
node, root_node, self.node_has_component_data
)
case LRURefreshPhase.INSERT_END:
# WALKDOWN already refreshed every node on the insert path
# (including the new leaf), so there is nothing more to do.
return
case _:
raise ValueError(f"Unknown LRURefreshPhase: {phase}")

@abstractmethod
def create_match_validator(
self, match_device_only: bool = False
Expand Down
37 changes: 33 additions & 4 deletions python/sglang/srt/mem_cache/unified_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
ComponentType,
EvictLayer,
FullComponent,
LRURefreshPhase,
MambaComponent,
SWAComponent,
TreeComponent,
Expand Down Expand Up @@ -186,6 +187,24 @@ def reset_node_and_parents_mru(
prev_node = node
node = node.parent

def reset_node_and_window_ancestors_mru(
self,
node: UnifiedTreeNode,
root_node: UnifiedTreeNode,
window_size: int,
should_include,
):
prev_node = self.head
accumulated = 0
while node != root_node and accumulated < window_size:
if should_include(node):
assert node.id in self.cache
self._remove_node(node)
self._add_node_after(prev_node, node)
prev_node = node
accumulated += len(node.key)
node = node.parent

def in_list(self, node: Optional[UnifiedTreeNode]):
return node is not None and node.id in self.cache

Expand Down Expand Up @@ -800,9 +819,7 @@ def _match_post_processor(
for comp in self._components_tuple:
if comp.component_type == BASE_COMPONENT_TYPE:
continue # Full uses last_access_time, not LRU
self.lru_lists[comp.component_type].reset_node_and_parents_mru(
node_update, self.root_node, comp.node_has_component_data
)
comp.refresh_lru(LRURefreshPhase.MATCH_END, node_update, self.root_node)

cur_time = get_and_increase_time_counter()
while node_update:
Expand Down Expand Up @@ -878,7 +895,10 @@ def _split_node(
def _touch_node(self, node: UnifiedTreeNode):
node.last_access_time = get_and_increase_time_counter()
if node != self.root_node:
self._for_each_component_lru(node, UnifiedLRUList.reset_node_mru)
for comp in self._components_tuple:
if comp.component_type == BASE_COMPONENT_TYPE:
continue
comp.refresh_lru(LRURefreshPhase.WALKDOWN, node, self.root_node)

def _add_new_node(
self,
Expand Down Expand Up @@ -1014,6 +1034,15 @@ def _insert_helper(
params=params,
result=result,
)

if target_node is not self.root_node:
for component in self._components_tuple:
if component.component_type == BASE_COMPONENT_TYPE:
continue
component.refresh_lru(
LRURefreshPhase.INSERT_END, target_node, self.root_node
)

if is_new_leaf:
self._inc_hit_count(target_node, params.chunked)
return result
Expand Down
Loading
Loading