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
17 changes: 14 additions & 3 deletions lightllm/server/multi_level_kv_cache/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
import threading
import concurrent.futures
import setproctitle
from queue import Queue
from queue import Empty, Queue
from typing import List
from lightllm.server.core.objs import ShmReqManager, Req, StartArgs
from lightllm.server.core.objs.io_objs import GroupReqIndexes
Expand Down Expand Up @@ -45,6 +45,8 @@ def __init__(
# 控制进行 cpu cache 页面匹配的时间,超过时间则不再匹配,直接转发。
self.cpu_cache_time_out = 0.5
self.recv_queue = Queue(maxsize=1024)
# Workers publish completions; only recv_loop accesses the PUSH socket.
self.send_to_router_queue = Queue()
self.cpu_cache_thread = threading.Thread(target=self.cpu_cache_hanle_loop, daemon=True)
self.cpu_cache_thread.start()

Expand Down Expand Up @@ -144,7 +146,7 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes
# 超时时,放弃进行 cache page 的匹配。
current_time = time.time()
if current_time - start_time >= self.cpu_cache_time_out:
self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL)
self.send_to_router_queue.put(group_req_indexes)
logger.warning(
f"cache matching time out {current_time - start_time}s, "
f"group_req_id: {group_req_indexes.group_req_id}"
Expand Down Expand Up @@ -211,14 +213,23 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes
for req in reqs:
self.shm_req_manager.put_back_req_obj(req)

self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL)
self.send_to_router_queue.put(group_req_indexes)
return

def _send_finished_group_reqs(self):
while True:
try:
group_req_indexes = self.send_to_router_queue.get_nowait()
except Empty:
return
self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL)

def recv_loop(self):
try:
recv_max_count = 128

while True:
self._send_finished_group_reqs()
recv_objs = []
try:
# 一次最多从 zmq 中取 recv_max_count 个请求,防止 zmq 队列中请求数量过多导致阻塞了主循环。
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
import pickle
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from queue import Queue
from types import SimpleNamespace
from unittest.mock import Mock

import pytest

from lightllm.server.multi_level_kv_cache.manager import MultiLevelKVCacheManager


@pytest.mark.parametrize("expired", [False, True])
def test_concurrent_workers_leave_socket_sends_to_owner_thread(expired):
manager = MultiLevelKVCacheManager.__new__(MultiLevelKVCacheManager)
manager.cpu_cache_time_out = 0.5
manager.send_to_router_queue = Queue()
manager.send_to_router = Mock()
manager.shm_req_manager = Mock()
groups = [SimpleNamespace(group_req_id=i, shm_req_indexes=[]) for i in range(100)]
start_time = time.time() - 10 if expired else time.time()

with ThreadPoolExecutor(max_workers=8) as executor:
list(executor.map(lambda group: manager._handle_group_req_multi_cache_match(group, start_time), groups))

manager.send_to_router.send_pyobj.assert_not_called()
owner_thread = threading.get_ident()
sent_groups = []

def send(group, protocol):
assert threading.get_ident() == owner_thread
assert protocol == pickle.HIGHEST_PROTOCOL
sent_groups.append(group.group_req_id)

manager.send_to_router.send_pyobj.side_effect = send
manager._send_finished_group_reqs()
assert sorted(sent_groups) == list(range(100))
assert manager.send_to_router_queue.empty()
manager._send_finished_group_reqs()
assert len(sent_groups) == 100
Loading