Skip to content
Closed
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
60 changes: 42 additions & 18 deletions lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@ def __init__(self, hash_page_size: int, big_page_num: int):

self.ref_counter = 0
self.time_id = time_gen.generate_time_id()
# 快照单独记录使用时间,避免匹配更长前缀时,沿途未使用的历史快照也被刷新为热门。
self.state_time_id = self.time_id

self.node_value_len = 0
self.node_prefix_total_len = 0
Expand Down Expand Up @@ -54,7 +56,7 @@ def get_compare_key_for_buffer_idx(self):
assert self.is_big_page_node() is False
# 对于有 buffer_id 的节点的回收处理比较器
assert self.small_page_buffer_idx is not None
return (self.time_id,)
return (self.state_time_id,)

def add_and_return_new_child(
self,
Expand Down Expand Up @@ -172,6 +174,14 @@ def _add_node(self, node: LinearAttPagedTreeNode):
self._evict_tree_set_for_linear_att.add(node)
return

def mark_small_page_state_used(self, node: LinearAttPagedTreeNode):
"""只有实际恢复了该节点的线性状态,才刷新快照 LRU;仅遍历到节点不算使用。"""
assert node.small_page_buffer_idx is not None
# SortedSet 的排序键发生变化时,必须先移除节点,再更新时间并重新加入。
self._evict_tree_set_for_linear_att.discard(node)
node.state_time_id = time_gen.generate_time_id()
self._evict_tree_set_for_linear_att.add(node)

def insert(
self,
key: torch.Tensor,
Expand Down Expand Up @@ -302,6 +312,8 @@ def _insert_helper(
# 将这个buffer id 移交给这个存在的节点。
self._discard_node(child)
child.small_page_buffer_idx = block_linear_idxs[0]
# 原快照已被淘汰,补入的新快照应从当前时间开始参与 LRU 排序。
child.state_time_id = time_gen.generate_time_id()
self._add_node(child)
else:
# 说明节点已经存在了,直接提前移除掉这个节点占用的线性缓存,外部不用处理这个细节了
Expand Down Expand Up @@ -619,22 +631,34 @@ def _evict(self, need_remove_tokens, evict_callback):
num_evicted = 0
while num_evicted < need_remove_tokens:
node: LinearAttPagedTreeNode = self._evict_tree_set.pop(0)
self._discard_node(node)

assert (
node.ref_counter == 0 and len(node.children) == 0 and node is not self.root_node
), "error evict tree node state"
num_evicted += len(node.token_mem_index_value)

if node.is_big_page_node():
assert node.big_page_buffer_idx is not None
self.linear_att_big_page_buffers.free_state_cache([node.big_page_buffer_idx])

evict_callback(node.token_mem_index_value, node.small_page_buffer_idx)
self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value)
parent_node: LinearAttPagedTreeNode = node.parent
parent_node.remove_child(node)

self._add_node(parent_node)
while True:
self._discard_node(node)

assert (
node.ref_counter == 0 and len(node.children) == 0 and node is not self.root_node
), "error evict tree node state"
num_evicted += len(node.token_mem_index_value)

if node.is_big_page_node():
assert node.big_page_buffer_idx is not None
self.linear_att_big_page_buffers.free_state_cache([node.big_page_buffer_idx])

evict_callback(node.token_mem_index_value, node.small_page_buffer_idx)
self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value)
parent_node: LinearAttPagedTreeNode = node.parent
parent_node.remove_child(node)

self._add_node(parent_node)
# 末尾恢复点被删除后,向上连续的无快照叶节点已无法复用,一并回收,
# 即使释放量已达标也继续清理;遇到引用、其他分支或有效快照就停止。
if (
parent_node is self.root_node
or parent_node.ref_counter != 0
or not parent_node.is_leaf()
or parent_node.is_big_page_node()
or parent_node.small_page_buffer_idx is not None
):
break
node = parent_node

return
3 changes: 3 additions & 0 deletions lightllm/server/router/model_infer/infer_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -734,6 +734,7 @@ def _linear_match_radix_cache(self):
req=self,
linear_att_small_page_buffers=g_infer_context.radix_cache.linear_att_small_page_buffers,
)
g_infer_context.radix_cache.mark_small_page_state_used(share_node)
else:
# 如果 大页本质是被启用的,则需要使用小页的匹配结果, 将小页的kv 复制到的新申请的kv位置,同时释放
# 对应的小页对应的节点,递归找到对应最近的大叶节点进行返回,然后赋值到req.shared_node 对象上
Expand Down Expand Up @@ -767,6 +768,8 @@ def _linear_match_radix_cache(self):
req=self,
linear_att_small_page_buffers=g_infer_context.radix_cache.linear_att_small_page_buffers,
)
# 仅在真正使用小页状态时刷新;内存不足而退回大页的分支不刷新。
radix_cache.mark_small_page_state_used(share_node)
self.shared_kv_node = None

big_page_shared_node = radix_cache.deref_to_first_big_page_node(node=share_node)
Expand Down
Loading