From a63e39454973d759e894eeee7156c1475fa206c4 Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 11 Sep 2026 16:17:43 +0800 Subject: [PATCH 1/2] fix(cache): reuse historical CPU tail pages for hybrid attention --- lightllm/server/core/objs/req.py | 4 +- .../server/multi_level_kv_cache/manager.py | 27 ++++ .../mode_backend/multi_level_kv_cache.py | 7 +- .../test_hybrid_tail_match.py | 122 ++++++++++++++++++ .../mode_backend/test_multi_level_kv_cache.py | 56 ++++++++ 5 files changed, 213 insertions(+), 3 deletions(-) create mode 100644 unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b5a5bf685f..12b4297820 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -152,6 +152,8 @@ class Req(ctypes.Structure): ("token_hash_page_len_list", TokenPageLenList), # 用于保存查找匹配到的可以被复用的cpu cache 页面信息。 ("cpu_cache_match_page_indexes", CpuCachePageList), + # 历史碎页的实际前缀终点;0 表示沿用原始页边界。与 offload 的 hash/长度列表分开保存。 + ("cpu_cache_match_tail_len", ctypes.c_int), ] def get_str(self): @@ -188,6 +190,7 @@ def init( self.candetoken_out_len = 0 self.prompt_cache_len = 0 self.cpu_prompt_cache_len = 0 + self.cpu_cache_match_tail_len = 0 self.disk_prompt_cache_len = 0 self.finish_token_index = -1 self.can_released_mark = False @@ -528,5 +531,4 @@ def get_decode_need_tokens(self): return need_tokens def get_first_router_need_tokens(self): - return min(self.input_len + self.shm_cur_output_len, self.chunked_prefill_size) diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index ef5b7369c9..a442bf74e4 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -19,6 +19,7 @@ from lightllm.utils.process_check import start_parent_check_thread from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args +from lightllm.utils.config_utils import is_hybrid_att_model logger = init_logger(__name__) @@ -29,6 +30,7 @@ def __init__( args: StartArgs, ): self.args: StartArgs = args + self.is_hybrid_att_model = is_hybrid_att_model(args.model_dir) ports = get_shm_port_args() context = zmq.Context(2) self.zmq_recv_socket = context.socket(zmq.PULL) @@ -137,6 +139,28 @@ def _disk_cache_match(self, token_hash_list: List[int], all_pages: List[int]) -> self.cpu_cache_client.lock.release() return all_pages, len(new_page_indexes) + def _match_hybrid_att_tail(self, req: Req, pages: List[int]): + """在首个缺失页内回退到最长的历史碎页,命中后不再向后拼接页面。""" + if not self.is_hybrid_att_model: + return + page_lens = req.token_hash_page_len_list.get_all() + if len(pages) == len(page_lens): + return + page_start = len(pages) * self.args.cpu_cache_token_page_size + page_end = page_lens[len(pages)] + hash_size = self.args.linear_att_hash_page_size + hashes = req.hybrid_token_hash_list.get_all() + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + for end in range(page_end - hash_size, page_start, -hash_size): + page_index, _ = self.cpu_cache_client.query_one_page(hashes[end // hash_size - 1]) + if page_index is not None: + pages.append(page_index) + req.cpu_cache_match_tail_len = end + return + finally: + self.cpu_cache_client.lock.release() + def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes, start_time: float): """ match cpu cache and disk cache pages @@ -203,6 +227,9 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes logger.exception(f"calculate disk prompt cache len has exception {str(e)}") raise e + # 优先保留完整 CPU/disk 页的命中机会,再查首个缺失页内的历史碎页。 + self._match_hybrid_att_tail(req, finded_page_indexes) + while not self.cpu_cache_client.check_allpages_ready(finded_page_indexes): time.sleep(0.01) diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py index 2a366f0256..13447189ea 100644 --- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py @@ -75,8 +75,12 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): continue page_len_list = req.shm_req.token_hash_page_len_list.get_all() - page_len_start_list = [0] + page_len_list assert len(page_list) <= len(page_len_list) + # 只调整加载视图,不能把后续 offload 的新尾页边界替换成历史边界。 + page_len_list = page_len_list[: len(page_list)] + if page_list and req.shm_req.cpu_cache_match_tail_len: + page_len_list[-1] = req.shm_req.cpu_cache_match_tail_len + page_len_start_list = [0] + page_len_list if page_list: match_tokens = page_len_list[len(page_list) - 1] @@ -226,7 +230,6 @@ def _start_kv_cache_offload_task( assert len(token_hash_list) == len(page_len_list) if self.backend.is_master_in_dp: - find_index = bisect.bisect_right(page_len_list, req.cur_kv_len) move_block_size = find_index diff --git a/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py b/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py new file mode 100644 index 0000000000..8c0cd5270b --- /dev/null +++ b/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py @@ -0,0 +1,122 @@ +import time +from collections import Counter +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from lightllm.server.core.objs import Req +from lightllm.server.core.objs import req as req_impl +from lightllm.server.core.objs.io_objs import GroupReqIndexes +from lightllm.server.multi_level_kv_cache.manager import MultiLevelKVCacheManager +from lightllm.utils.kv_cache_utils import compute_token_list_hash + + +@pytest.fixture +def cache_match(monkeypatch): + args = SimpleNamespace( + cpu_cache_token_page_size=1024, + linear_att_hash_page_size=256, + linear_att_page_block_num=4, + diverse_mode=False, + ) + monkeypatch.setattr(req_impl, "get_env_start_args", lambda: args) + manager = MultiLevelKVCacheManager.__new__(MultiLevelKVCacheManager) + manager.args = args + manager.is_hybrid_att_model = True + manager.only_cpu_cache_enable = True + manager.cpu_cache_time_out = 10 + manager.send_to_router = Mock() + manager.shm_req_manager = Mock() + ready_pages = {} + refs = Counter() + + def query(hash_key): + page = ready_pages.get(hash_key) + if page is not None: + refs[page] += 1 + return page, page is not None + + manager.cpu_cache_client = Mock() + manager.cpu_cache_client.query_one_page.side_effect = query + manager.cpu_cache_client.check_allpages_ready.return_value = True + + def make_request(length, stored_ends): + request = Req() + request.input_len = length + request.sample_params.prompt_logprobs = -1 + hashes = compute_token_list_hash(list(range(length)), args.linear_att_hash_page_size) + request.hybrid_token_hash_list.fill(hashes) + page_hashes, page_lens = request._calcu_hybrid_cpu_cache_page_len_list() + request.token_hash_list.fill(page_hashes) + request.token_hash_page_len_list.fill(page_lens) + for page, end in enumerate(stored_ends): + ready_pages[hashes[end // args.linear_att_hash_page_size - 1]] = page + manager.shm_req_manager.get_req_obj_by_index.return_value = request + return request + + def match(request): + group = GroupReqIndexes(0, None, [0], time.time()) + manager._handle_group_req_multi_cache_match(group, time.time()) + manager.send_to_router.send_pyobj.assert_called_once() + return request.cpu_cache_match_page_indexes.get_all() + + return manager, make_request, match, refs + + +@pytest.mark.parametrize( + "length, stored_ends, expected_pages, tail_end", + [ + (1025, [512, 768], [1], 768), + (769, [512], [0], 512), + (2049, [1024, 1280, 1792], [0, 2], 1792), + (2049, [512, 2048], [0], 512), + (2049, [1024, 2048, 1792], [0, 1], 0), + (769, [768, 512], [0], 0), + (2049, [1024], [0], 0), + (257, [], [], 0), + (256, [], [], 0), + ], +) +def test_match_longest_ready_tail_after_full_pages(cache_match, length, stored_ends, expected_pages, tail_end): + manager, make_request, match, refs = cache_match + request = make_request(length, stored_ends) + page_hashes = request.token_hash_list.get_all() + page_lens = request.token_hash_page_len_list.get_all() + + assert match(request) == expected_pages + assert request.cpu_cache_match_tail_len == tail_end + assert refs == Counter(expected_pages) + assert request.token_hash_list.get_all() == page_hashes + assert request.token_hash_page_len_list.get_all() == page_lens + lock = manager.cpu_cache_client.lock + assert lock.acquire_sleep1ms.call_count == lock.release.call_count + + +def test_full_attention_does_not_match_hybrid_tail(cache_match): + manager, make_request, match, refs = cache_match + manager.is_hybrid_att_model = False + request = make_request(1025, [768]) + assert match(request) == [] + assert not refs + + +def test_disk_pages_take_precedence_over_cpu_tail(cache_match): + manager, make_request, match, refs = cache_match + request = make_request(3073, [1024, 1792, 2560]) + manager.only_cpu_cache_enable = False + manager._disk_cache_match = Mock(return_value=([0, 99], 1)) + + assert match(request) == [0, 99, 2] + assert request.cpu_cache_match_tail_len == 2560 + assert request.disk_prompt_cache_len == 1024 + assert refs == Counter([0, 2]) + + +def test_prompt_logprobs_skips_tail_matching(cache_match): + manager, make_request, match, refs = cache_match + request = make_request(1025, [768]) + request.sample_params.prompt_logprobs = 0 + assert match(request) == [] + assert not refs + manager.cpu_cache_client.query_one_page.assert_not_called() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py index 96925b6cb2..cf6df87ab5 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py @@ -1,5 +1,6 @@ from types import SimpleNamespace from collections import deque +from unittest.mock import Mock import pytest import torch @@ -12,6 +13,7 @@ from lightllm.server.router.model_infer.mode_backend.multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.server.router.model_infer.mode_backend import multi_level_kv_cache as multi_level_kv_cache_impl from lightllm.server.router.model_infer.infer_batch import InferenceContext +from lightllm.server.core.objs import Req def test_cache_tiers_reassignment_is_rejected(): @@ -157,3 +159,57 @@ def test_non_gpu_linear_cache_tiers_release_pending_state_pages(): assert freed_big_pages == [8, 9] assert req.tail_small_page_buffer_id is None assert req.hybrid_len_to_big_page_id == {} + + +@pytest.mark.parametrize("cached_gpu_len", [0, 1024, 1280]) +@pytest.mark.parametrize("available_tokens", [0, 4096]) +def test_cpu_tail_load_uses_actual_endpoint_and_releases_pages(monkeypatch, cached_gpu_len, available_tokens): + shm_req = Req() + shm_req.input_len = 2049 + shm_req.token_hash_page_len_list.fill([1024, 2048]) + shm_req.cpu_cache_match_page_indexes.fill([10, 11]) + shm_req.cpu_cache_match_tail_len = 1792 + req = SimpleNamespace( + shm_req=shm_req, + cur_kv_len=cached_gpu_len, + req_idx=0, + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), + ) + original_indexes = torch.arange(2048, dtype=torch.int32) + 1000 + req_manager = SimpleNamespace( + req_to_token_indexs=original_indexes.clone().unsqueeze(0), + mem_manager=SimpleNamespace(alloc=lambda need_size: torch.arange(need_size, dtype=torch.int32) + 3000), + ) + context = SimpleNamespace(get_can_alloc_token_num=lambda: available_tokens, req_manager=req_manager) + monkeypatch.setattr(multi_level_kv_cache_impl, "g_infer_context", context) + # Exercise the loading plan on CPU; the transfer kernel itself is unchanged. + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, **kwargs: self) + monkeypatch.setattr(torch.cuda, "current_stream", Mock()) + monkeypatch.setattr(multi_level_kv_cache_impl.dist, "barrier", Mock()) + operator = Mock() + module = MultiLevelKvCacheModule.__new__(MultiLevelKvCacheModule) + module.backend = SimpleNamespace( + is_master_in_dp=True, + radix_cache=None, + model=SimpleNamespace(mem_manager=SimpleNamespace(operator=operator), req_manager=req_manager), + ) + module.need_sync_compute_stream = lambda: False + module.cpu_cache_client = Mock() + module.init_sync_group = object() + + module.load_cpu_cache_to_reqs([req]) + + assert shm_req.token_hash_page_len_list.get_all() == [1024, 2048] + module.cpu_cache_client.deref_pages.assert_called_once_with(page_list=[10, 11]) + module.cpu_cache_client.lock.release.assert_called_once() + if available_tokens: + assert req.cur_kv_len == shm_req.shm_cur_kv_len == 1792 + assert shm_req.cpu_prompt_cache_len == 1792 - cached_gpu_len + kwargs = operator.load_cpu_cache_to_gpu.call_args.kwargs + start = 0 if cached_gpu_len < 1024 else 1024 + expected = torch.cat([original_indexes[start:cached_gpu_len], torch.arange(1792 - cached_gpu_len) + 3000]) + assert torch.equal(kwargs["mem_indexes"], expected) + assert kwargs["page_indexes"].tolist() == ([10, 11] if start == 0 else [11]) + else: + assert req.cur_kv_len == cached_gpu_len + operator.load_cpu_cache_to_gpu.assert_not_called() From 1172066bdff29dd73021a22dec112c1d2a47b9a4 Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 11 Sep 2026 16:21:16 +0800 Subject: [PATCH 2/2] test: remove CPU tail cache regression tests from PR --- .../test_hybrid_tail_match.py | 122 ------------------ .../mode_backend/test_multi_level_kv_cache.py | 56 -------- 2 files changed, 178 deletions(-) delete mode 100644 unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py diff --git a/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py b/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py deleted file mode 100644 index 8c0cd5270b..0000000000 --- a/unit_tests/server/multi_level_kv_cache/test_hybrid_tail_match.py +++ /dev/null @@ -1,122 +0,0 @@ -import time -from collections import Counter -from types import SimpleNamespace -from unittest.mock import Mock - -import pytest - -from lightllm.server.core.objs import Req -from lightllm.server.core.objs import req as req_impl -from lightllm.server.core.objs.io_objs import GroupReqIndexes -from lightllm.server.multi_level_kv_cache.manager import MultiLevelKVCacheManager -from lightllm.utils.kv_cache_utils import compute_token_list_hash - - -@pytest.fixture -def cache_match(monkeypatch): - args = SimpleNamespace( - cpu_cache_token_page_size=1024, - linear_att_hash_page_size=256, - linear_att_page_block_num=4, - diverse_mode=False, - ) - monkeypatch.setattr(req_impl, "get_env_start_args", lambda: args) - manager = MultiLevelKVCacheManager.__new__(MultiLevelKVCacheManager) - manager.args = args - manager.is_hybrid_att_model = True - manager.only_cpu_cache_enable = True - manager.cpu_cache_time_out = 10 - manager.send_to_router = Mock() - manager.shm_req_manager = Mock() - ready_pages = {} - refs = Counter() - - def query(hash_key): - page = ready_pages.get(hash_key) - if page is not None: - refs[page] += 1 - return page, page is not None - - manager.cpu_cache_client = Mock() - manager.cpu_cache_client.query_one_page.side_effect = query - manager.cpu_cache_client.check_allpages_ready.return_value = True - - def make_request(length, stored_ends): - request = Req() - request.input_len = length - request.sample_params.prompt_logprobs = -1 - hashes = compute_token_list_hash(list(range(length)), args.linear_att_hash_page_size) - request.hybrid_token_hash_list.fill(hashes) - page_hashes, page_lens = request._calcu_hybrid_cpu_cache_page_len_list() - request.token_hash_list.fill(page_hashes) - request.token_hash_page_len_list.fill(page_lens) - for page, end in enumerate(stored_ends): - ready_pages[hashes[end // args.linear_att_hash_page_size - 1]] = page - manager.shm_req_manager.get_req_obj_by_index.return_value = request - return request - - def match(request): - group = GroupReqIndexes(0, None, [0], time.time()) - manager._handle_group_req_multi_cache_match(group, time.time()) - manager.send_to_router.send_pyobj.assert_called_once() - return request.cpu_cache_match_page_indexes.get_all() - - return manager, make_request, match, refs - - -@pytest.mark.parametrize( - "length, stored_ends, expected_pages, tail_end", - [ - (1025, [512, 768], [1], 768), - (769, [512], [0], 512), - (2049, [1024, 1280, 1792], [0, 2], 1792), - (2049, [512, 2048], [0], 512), - (2049, [1024, 2048, 1792], [0, 1], 0), - (769, [768, 512], [0], 0), - (2049, [1024], [0], 0), - (257, [], [], 0), - (256, [], [], 0), - ], -) -def test_match_longest_ready_tail_after_full_pages(cache_match, length, stored_ends, expected_pages, tail_end): - manager, make_request, match, refs = cache_match - request = make_request(length, stored_ends) - page_hashes = request.token_hash_list.get_all() - page_lens = request.token_hash_page_len_list.get_all() - - assert match(request) == expected_pages - assert request.cpu_cache_match_tail_len == tail_end - assert refs == Counter(expected_pages) - assert request.token_hash_list.get_all() == page_hashes - assert request.token_hash_page_len_list.get_all() == page_lens - lock = manager.cpu_cache_client.lock - assert lock.acquire_sleep1ms.call_count == lock.release.call_count - - -def test_full_attention_does_not_match_hybrid_tail(cache_match): - manager, make_request, match, refs = cache_match - manager.is_hybrid_att_model = False - request = make_request(1025, [768]) - assert match(request) == [] - assert not refs - - -def test_disk_pages_take_precedence_over_cpu_tail(cache_match): - manager, make_request, match, refs = cache_match - request = make_request(3073, [1024, 1792, 2560]) - manager.only_cpu_cache_enable = False - manager._disk_cache_match = Mock(return_value=([0, 99], 1)) - - assert match(request) == [0, 99, 2] - assert request.cpu_cache_match_tail_len == 2560 - assert request.disk_prompt_cache_len == 1024 - assert refs == Counter([0, 2]) - - -def test_prompt_logprobs_skips_tail_matching(cache_match): - manager, make_request, match, refs = cache_match - request = make_request(1025, [768]) - request.sample_params.prompt_logprobs = 0 - assert match(request) == [] - assert not refs - manager.cpu_cache_client.query_one_page.assert_not_called() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py index cf6df87ab5..96925b6cb2 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py @@ -1,6 +1,5 @@ from types import SimpleNamespace from collections import deque -from unittest.mock import Mock import pytest import torch @@ -13,7 +12,6 @@ from lightllm.server.router.model_infer.mode_backend.multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.server.router.model_infer.mode_backend import multi_level_kv_cache as multi_level_kv_cache_impl from lightllm.server.router.model_infer.infer_batch import InferenceContext -from lightllm.server.core.objs import Req def test_cache_tiers_reassignment_is_rejected(): @@ -159,57 +157,3 @@ def test_non_gpu_linear_cache_tiers_release_pending_state_pages(): assert freed_big_pages == [8, 9] assert req.tail_small_page_buffer_id is None assert req.hybrid_len_to_big_page_id == {} - - -@pytest.mark.parametrize("cached_gpu_len", [0, 1024, 1280]) -@pytest.mark.parametrize("available_tokens", [0, 4096]) -def test_cpu_tail_load_uses_actual_endpoint_and_releases_pages(monkeypatch, cached_gpu_len, available_tokens): - shm_req = Req() - shm_req.input_len = 2049 - shm_req.token_hash_page_len_list.fill([1024, 2048]) - shm_req.cpu_cache_match_page_indexes.fill([10, 11]) - shm_req.cpu_cache_match_tail_len = 1792 - req = SimpleNamespace( - shm_req=shm_req, - cur_kv_len=cached_gpu_len, - req_idx=0, - sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), - ) - original_indexes = torch.arange(2048, dtype=torch.int32) + 1000 - req_manager = SimpleNamespace( - req_to_token_indexs=original_indexes.clone().unsqueeze(0), - mem_manager=SimpleNamespace(alloc=lambda need_size: torch.arange(need_size, dtype=torch.int32) + 3000), - ) - context = SimpleNamespace(get_can_alloc_token_num=lambda: available_tokens, req_manager=req_manager) - monkeypatch.setattr(multi_level_kv_cache_impl, "g_infer_context", context) - # Exercise the loading plan on CPU; the transfer kernel itself is unchanged. - monkeypatch.setattr(torch.Tensor, "cuda", lambda self, **kwargs: self) - monkeypatch.setattr(torch.cuda, "current_stream", Mock()) - monkeypatch.setattr(multi_level_kv_cache_impl.dist, "barrier", Mock()) - operator = Mock() - module = MultiLevelKvCacheModule.__new__(MultiLevelKvCacheModule) - module.backend = SimpleNamespace( - is_master_in_dp=True, - radix_cache=None, - model=SimpleNamespace(mem_manager=SimpleNamespace(operator=operator), req_manager=req_manager), - ) - module.need_sync_compute_stream = lambda: False - module.cpu_cache_client = Mock() - module.init_sync_group = object() - - module.load_cpu_cache_to_reqs([req]) - - assert shm_req.token_hash_page_len_list.get_all() == [1024, 2048] - module.cpu_cache_client.deref_pages.assert_called_once_with(page_list=[10, 11]) - module.cpu_cache_client.lock.release.assert_called_once() - if available_tokens: - assert req.cur_kv_len == shm_req.shm_cur_kv_len == 1792 - assert shm_req.cpu_prompt_cache_len == 1792 - cached_gpu_len - kwargs = operator.load_cpu_cache_to_gpu.call_args.kwargs - start = 0 if cached_gpu_len < 1024 else 1024 - expected = torch.cat([original_indexes[start:cached_gpu_len], torch.arange(1792 - cached_gpu_len) + 3000]) - assert torch.equal(kwargs["mem_indexes"], expected) - assert kwargs["page_indexes"].tolist() == ([10, 11] if start == 0 else [11]) - else: - assert req.cur_kv_len == cached_gpu_len - operator.load_cpu_cache_to_gpu.assert_not_called()