From 79518c0f7ed9911877703f42011a40c33ef4ec66 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 12 Aug 2026 02:45:04 +0000 Subject: [PATCH 1/5] fix(pd): account split prefill requests --- .../httpserver_for_pd_master/manager.py | 30 +++++--- lightllm/server/pd_io_struct.py | 2 + .../test_pd_master_multi_choice.py | 74 ++++++++++++++++++- .../test_pd_master_cached_tokens.py | 2 +- unit_tests/server/test_pd_master_mode.py | 5 ++ 5 files changed, 96 insertions(+), 17 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 18cfa0d51..89b8ad19b 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -211,19 +211,15 @@ async def _generate_one( block_group_request_id = origin_request_id p_node = None d_node = None - prefill_load_released = False + pending_prefill_load_chars = None try: p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) - # 记录当前 P 节点的在途 prompt 负载,首 token 生成后即可释放。 - p_node.dispatched_prompt_chars += len(prompt) - - history_gen_token_strs = [] - if not p_node or not d_node: logger.error(f"{origin_request_id}: No p_node or d_node found") raise Exception(f"{origin_request_id}: No p_node or d_node found") + history_gen_token_strs = [] origin_prompt_cache_len = None for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list): @@ -233,11 +229,17 @@ async def _generate_one( logger.info(f"pd log gen sub req id {block_group_request_id} for main req id {origin_request_id}") sampling_params.max_new_tokens = block_max_new_tokens + # 分段请求始终复用循环外选定的 P 节点;这里只按每段实际发送的 + # prompt 更新该节点的在途 prefill 负载,不会重新选点。 + block_prompt = prompt + "".join(history_gen_token_strs) + pending_prefill_load_chars = len(block_prompt) + p_node.dispatched_prompt_chars += pending_prefill_load_chars + p_node.dispatched_req_num += 1 results_generator = self._wait_to_token_package( p_node, d_node, start_time, - prompt + "".join(history_gen_token_strs), + block_prompt, sampling_params, multimodal_params, request, @@ -255,9 +257,12 @@ async def _generate_one( if iter_index == 0 and origin_prompt_cache_len is None: origin_prompt_cache_len = metadata.get("prompt_cache_len", 0) metadata["prompt_cache_len"] = origin_prompt_cache_len or 0 - if not prefill_load_released: - p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt)) - prefill_load_released = True + if pending_prefill_load_chars is not None: + p_node.dispatched_prompt_chars = max( + 0, p_node.dispatched_prompt_chars - pending_prefill_load_chars + ) + p_node.dispatched_req_num = max(0, p_node.dispatched_req_num - 1) + pending_prefill_load_chars = None yield origin_request_id, request_output, metadata, finish_status await self.remove_req(group_request_id=block_group_request_id) @@ -277,8 +282,9 @@ async def _generate_one( raise e finally: - if p_node is not None and not prefill_load_released: - p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt)) + if p_node is not None and pending_prefill_load_chars is not None: + p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - pending_prefill_load_chars) + p_node.dispatched_req_num = max(0, p_node.dispatched_req_num - 1) await self.remove_req(block_group_request_id) return diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index e469adf94..c019e90ac 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -58,6 +58,8 @@ class PD_Client_Obj: run_status: _PD_Client_RunStatus = field(default_factory=_PD_Client_RunStatus) # cache-aware 选点用:当前派发到该节点且尚未产出首 token 的 prompt 字符数。 dispatched_prompt_chars: int = 0 + # 当前派发到该节点且尚未产出首 token 的请求数。 + dispatched_req_num: int = 0 def __post_init__(self): if self.mode not in ["prefill", "decode"]: diff --git a/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index 5bad24bdd..55d9eb98e 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -36,7 +36,7 @@ async def run(): started = set() all_started = asyncio.Event() captured_params = [] - p_node = MagicMock(dispatched_prompt_chars=0) + p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) manager._split_max_new_tokens = MagicMock(return_value=[4]) @@ -94,6 +94,7 @@ async def wait_to_token_package( call("lightllm_request_success"), ] assert p_node.dispatched_prompt_chars == 0 + assert p_node.dispatched_req_num == 0 asyncio.run(asyncio.wait_for(run(), timeout=2)) @@ -213,7 +214,7 @@ async def run(): manager.id_gen.generate_id.return_value = 808 manager.remove_req = AsyncMock() manager.abort = AsyncMock() - p_node = MagicMock(dispatched_prompt_chars=0) + p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) @@ -224,8 +225,9 @@ async def failing_wait_to_token_package(*_args, **_kwargs): manager._wait_to_token_package = failing_wait_to_token_package with pytest.raises(RuntimeError, match="generation failed"): + # 空 prompt 的字符负载为 0,但已派发请求数仍必须在异常路径释放。 async for _ in manager._generate_one( - "prompt", + "", SamplingParams(), MagicMock(), MagicMock(), @@ -235,6 +237,64 @@ async def failing_wait_to_token_package(*_args, **_kwargs): pass assert p_node.dispatched_prompt_chars == 0 + assert p_node.dispatched_req_num == 0 + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + +def test_pd_master_accounts_each_split_prefill_on_the_same_node(): + async def run(): + manager = _manager() + manager._split_max_new_tokens = MagicMock(return_value=[1, 1]) + manager.id_gen.generate_id.side_effect = [808, 816] + manager.remove_req = AsyncMock() + manager.abort = AsyncMock() + other_request_load = 17 + other_request_count = 3 + p_node = MagicMock( + dispatched_prompt_chars=other_request_load, + dispatched_req_num=other_request_count, + ) + d_node = MagicMock() + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + dispatched_nodes = [] + dispatched_prompts = [] + dispatched_loads = [] + dispatched_req_counts = [] + + async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args): + dispatched_nodes.append(selected_p_node) + dispatched_prompts.append(block_prompt) + dispatched_loads.append(selected_p_node.dispatched_prompt_chars) + dispatched_req_counts.append(selected_p_node.dispatched_req_num) + yield ( + sampling_params.group_request_id, + "x", + {"prompt_tokens": 1}, + FinishStatus(FinishStatus.FINISHED_LENGTH), + ) + + manager._wait_to_token_package = wait_to_token_package + + results = [] + async for result in manager._generate_one( + "prompt", + SamplingParams(), + MagicMock(), + MagicMock(), + 0, + 800, + ): + results.append(result) + + manager.select_p_d_node.assert_awaited_once() + assert dispatched_nodes == [p_node, p_node] + assert dispatched_prompts == ["prompt", "promptx"] + assert dispatched_loads == [other_request_load + len("prompt"), other_request_load + len("promptx")] + assert dispatched_req_counts == [other_request_count + 1, other_request_count + 1] + assert p_node.dispatched_prompt_chars == other_request_load + assert p_node.dispatched_req_num == other_request_count + assert len(results) == 2 asyncio.run(asyncio.wait_for(run(), timeout=2)) @@ -247,7 +307,11 @@ async def run(): manager.remove_req = AsyncMock() manager.abort = AsyncMock() other_request_load = 17 - p_node = MagicMock(dispatched_prompt_chars=other_request_load) + other_request_count = 3 + p_node = MagicMock( + dispatched_prompt_chars=other_request_load, + dispatched_req_num=other_request_count, + ) d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) @@ -267,9 +331,11 @@ async def wait_to_token_package(*_args, **_kwargs): assert (await generator.__anext__())[1] == "first" assert p_node.dispatched_prompt_chars == other_request_load + assert p_node.dispatched_req_num == other_request_count await generator.aclose() assert p_node.dispatched_prompt_chars == other_request_load + assert p_node.dispatched_req_num == other_request_count manager.abort.assert_awaited_once() asyncio.run(asyncio.wait_for(run(), timeout=2)) diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index 9fa60d26e..09bf04672 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -26,7 +26,7 @@ def gen_id(): mgr.metric_client = SimpleNamespace(counter_inc=lambda *a, **k: None, histogram_observe=lambda *a, **k: None) mgr.tokens = lambda *a, **k: 10 mgr._log_req_header = lambda *a, **k: asyncio.sleep(0) - p_node = SimpleNamespace(dispatched_prompt_chars=0) + p_node = SimpleNamespace(dispatched_prompt_chars=0, dispatched_req_num=0) mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(p_node, 1)) mgr.remove_req = lambda *a, **k: asyncio.sleep(0) return mgr diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 777d41412..edb758a94 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -205,10 +205,12 @@ def register_prefill(node_id, client_ip_port): register_prefill(1, "10.0.0.1:8000") manager.prefill_nodes[0].dispatched_prompt_chars = 1234 + manager.prefill_nodes[0].dispatched_req_num = 12 register_prefill(2, "10.0.0.2:8000") assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [1234, 0] + assert [node.dispatched_req_num for node in manager.prefill_nodes] == [12, 0] def test_prefill_reconnection_preserves_other_nodes_inflight_prompt_chars(): @@ -231,11 +233,14 @@ def pd_info(node_id, client_ip_port): manager.register_pd(pd_info(2, "10.0.0.2:8000"), websocket=object()) manager.prefill_nodes[0].dispatched_prompt_chars = 100 manager.prefill_nodes[1].dispatched_prompt_chars = 200 + manager.prefill_nodes[0].dispatched_req_num = 1 + manager.prefill_nodes[1].dispatched_req_num = 2 manager.register_pd(pd_info(3, "10.0.0.2:8000"), websocket=object()) assert [node.client_ip_port for node in manager.prefill_nodes] == ["10.0.0.1:8000", "10.0.0.2:8000"] assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [100, 0] + assert [node.dispatched_req_num for node in manager.prefill_nodes] == [1, 0] def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): From 418bbae6a1b109378be12df23676bc6125e6f7b1 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 12 Aug 2026 06:13:13 +0000 Subject: [PATCH 2/5] feat(pd): prioritize idle prefill nodes --- .../pd_selector/cache_aware.py | 73 ++++++++++++++----- unit_tests/server/test_pd_cache_aware.py | 36 +++++++-- 2 files changed, 87 insertions(+), 22 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index ad7348c86..20444b669 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -10,7 +10,8 @@ - 用前缀树(见 PromptCacheTree)记录「历史 prompt -> 处理它的 worker」; - 树中的 prefill_node 对应 worker.client_ip_port; - prompt 会按 sample_stride 抽稀后再插入/匹配,降低树的深度与内存; - - 用 worker.dispatched_prompt_chars(当前尚未产出首 token 的 prompt 字符数)做负载均衡。 + - 优先使用 dispatched_req_num 为 0 的空闲节点,避免 GPU 闲置; + - 所有节点都忙时,用 worker.dispatched_prompt_chars 做负载均衡。 选点流程见 CacheAwarePolicy.select_worker。 """ @@ -79,19 +80,35 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti 决策顺序: 1) workers 为空 -> 返回 None; - 2) 对 request_text 做前缀匹配,计算 + 2) 存在 dispatched_req_num 为 0 的空闲节点 -> 强制从空闲节点中选择; + 多个节点空闲时优先选择 cache 命中节点; + 3) 对 request_text 做前缀匹配,计算 match_rate = matched_char_count / input_char_count; - 3) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 得到 cache 命中节点; - 4) cache 命中节点负载未严重高于最空闲节点 -> 选择 cache 命中节点; - 5) 未命中或负载严重失衡 -> 选择最空闲节点; - 6) 将当前 prompt 与最终选中的节点写入前缀树。 + 4) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 得到 cache 命中节点; + 5) cache 命中节点负载未严重高于最空闲节点 -> 选择 cache 命中节点; + 6) 未命中或负载严重失衡 -> 选择最空闲节点; + 7) 将当前 prompt 与最终选中的节点写入前缀树。 """ if not workers: return None if len(request_text) <= 1: raise ValueError(f"request_text length must be > 1, got {len(request_text)}") - # ---- 1. 前缀匹配:估计当前请求与历史请求的 cache 复用潜力 ---- + # ---- 1. 空闲优先:避免有可用 GPU 闲置 ---- + idle_worker = self._select_idle_worker(workers, request_text) + if idle_worker is not None: + self.prompt_cache_tree.insert(request_text, idle_worker.client_ip_port) + return idle_worker + + # ---- 2. 所有节点都忙时,在 cache 亲和与负载均衡之间权衡 ---- + cache_worker = self._get_cache_worker(workers, request_text) + selected_worker = self._select_worker_by_cache_and_load(workers, cache_worker, len(request_text)) + + self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port) + return selected_worker + + def _get_cache_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: + """在指定候选节点中返回达到匹配阈值的 cache 节点。""" result = self.prompt_cache_tree.prefix_match(request_text) match_rate = 0.0 if result.input_char_count == 0 else result.matched_char_count / result.input_char_count @@ -102,16 +119,40 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti f"prefill_node={result.prefill_node}" ) - cache_worker: Optional[PD_Client_Obj] = None - if match_rate > self.config.cache_threshold and result.prefill_node is not None: - for worker in workers: - if worker.client_ip_port == result.prefill_node: - cache_worker = worker - break + if match_rate <= self.config.cache_threshold or result.prefill_node is None: + return None + + for worker in workers: + if worker.client_ip_port == result.prefill_node: + return worker + return None + + def _select_idle_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: + """优先选择空闲节点;多个空闲节点之间优先复用 cache。""" + idle_workers = [worker for worker in workers if worker.dispatched_req_num == 0] + if not idle_workers: + return None - # ---- 2. 负载配平:cache 节点过载时选择当前在途量最少的节点 ---- + cache_worker = self._get_cache_worker(idle_workers, request_text) if len(idle_workers) > 1 else None + selected_worker = cache_worker or min( + idle_workers, + key=lambda worker: (worker.dispatched_prompt_chars, worker.client_ip_port), + ) + logger.info( + f"CacheAwarePolicy: select idle worker, idle_worker_num={len(idle_workers)}, " + f"cache_worker={cache_worker.client_ip_port if cache_worker else None}, " + f"selected_worker={selected_worker.client_ip_port}" + ) + return selected_worker + + def _select_worker_by_cache_and_load( + self, + workers: List[PD_Client_Obj], + cache_worker: Optional[PD_Client_Obj], + request_load: int, + ) -> PD_Client_Obj: + """所有节点都忙时,在 cache 亲和与 prompt 负载之间选择节点。""" least_loaded_worker = min(workers, key=lambda worker: worker.dispatched_prompt_chars) - request_load = len(request_text) least_projected_load = least_loaded_worker.dispatched_prompt_chars + request_load cache_projected_load = None cache_worker_is_overloaded = False @@ -135,6 +176,4 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti f"cache_worker_is_overloaded={cache_worker_is_overloaded}, " f"selected_worker={selected_worker.client_ip_port}" ) - - self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port) return selected_worker diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 665311767..23b6de590 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -3,17 +3,18 @@ from lightllm.server.httpserver_for_pd_master.pd_selector.cache_aware import CacheAwarePolicy -def _worker(address: str, dispatched_prompt_chars: int = 0): +def _worker(address: str, dispatched_prompt_chars: int = 0, dispatched_req_num: int = 0): return SimpleNamespace( client_ip_port=address, dispatched_prompt_chars=dispatched_prompt_chars, + dispatched_req_num=dispatched_req_num, ) def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy = CacheAwarePolicy() - cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110) - least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100) + cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) + least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100, dispatched_req_num=2) prompt = "shared prefix " * 100 policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) @@ -24,8 +25,8 @@ def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): def test_cache_aware_uses_least_loaded_worker_when_cache_worker_is_overloaded(): policy = CacheAwarePolicy() - cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=2000) - least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100) + cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=2000, dispatched_req_num=2) + least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100, dispatched_req_num=2) prompt = "shared prefix " * 100 policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) @@ -44,3 +45,28 @@ def test_cache_aware_keeps_cache_worker_when_both_workers_are_idle(): selected_worker = policy.select_worker([other_worker, cache_worker], prompt) assert selected_worker is cache_worker + + +def test_cache_aware_forces_idle_worker_over_busy_cache_worker(): + policy = CacheAwarePolicy() + cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=100, dispatched_req_num=1) + idle_worker = _worker("10.0.0.2:8000") + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) + + selected_worker = policy.select_worker([cache_worker, idle_worker], prompt) + + assert selected_worker is idle_worker + + +def test_cache_aware_matches_cache_only_within_idle_workers(): + policy = CacheAwarePolicy() + busy_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=100, dispatched_req_num=1) + cache_idle_worker = _worker("10.0.0.2:8000") + other_idle_worker = _worker("10.0.0.3:8000") + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_idle_worker.client_ip_port) + + selected_worker = policy.select_worker([other_idle_worker, busy_worker, cache_idle_worker], prompt) + + assert selected_worker is cache_idle_worker From 5b27b5345f0de69759b70d36f179516faf062e5f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 12 Aug 2026 07:23:56 +0000 Subject: [PATCH 3/5] fix(pd): disable linear attention big pages on decode --- lightllm/server/api_start.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 33cc1fa71..7524b77fe 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -256,6 +256,12 @@ def _launch_subprocesses(args: StartArgs): per_dp_cache_size = max(1, math.ceil(args.running_max_req_size / dp_size_in_node) * 2) args.linear_att_cache_size = min(default_cache_size, per_dp_cache_size) + if args.run_mode == "decode": + # PD Decode 节点只接收 prompt 末尾位置的 linear attention state,不具备 + # 中间大页边界对应的 state。因此 Decode 节点必须使用默认值关闭大页功能, + # 避免请求释放时将不完整的大页 state 写入 radix cache 并触发断言。 + args.linear_att_page_block_num = 10000000 + if args.enable_cpu_cache and is_linear_att_mixed_model(args.model_dir): args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for linear hybrid att model") From 94f4f63bccbbefb11c18356c5c981e31a228c13a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 12 Aug 2026 07:40:31 +0000 Subject: [PATCH 4/5] tune(pd): tighten cache-aware load threshold --- .../server/httpserver_for_pd_master/pd_selector/cache_aware.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index 20444b669..ce186b9f0 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -37,7 +37,7 @@ class CacheAwareConfig: # 前缀匹配成功率阈值:matched_char_count / input_char_count 超过该值才路由到命中节点。 cache_threshold: float = 0.5 # cache 命中节点的在途量超过最空闲节点该倍数时,优先选择最空闲节点。 - balance_rel_threshold: float = 1.8 + balance_rel_threshold: float = 1.5 # 前缀树允许的最大节点数(不含 root)。 max_node_count: int = 1_000_000 # 每次 LRU 驱逐的叶节点数量。 From 0a01acd5833bede9f7699e454d38e12c72559825 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 12 Aug 2026 08:51:05 +0000 Subject: [PATCH 5/5] feat(pd): dynamically tune cache-aware balance threshold --- .../httpserver_for_pd_master/manager.py | 2 + .../pd_selector/cache_aware.py | 50 +++++++++++++++ .../pd_selector/pd_selector.py | 7 ++ .../test_pd_master_multi_choice.py | 1 + .../test_pd_master_cached_tokens.py | 12 +++- unit_tests/server/test_pd_cache_aware.py | 64 ++++++++++++++++++- 6 files changed, 132 insertions(+), 4 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 89b8ad19b..96d1361e6 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -256,6 +256,8 @@ async def _generate_one( metadata["prompt_tokens"] = prompt_tokens if iter_index == 0 and origin_prompt_cache_len is None: origin_prompt_cache_len = metadata.get("prompt_cache_len", 0) + prompt_cache_hit_rate = origin_prompt_cache_len / max(prompt_tokens, 1) + self.pd_manager.selector.record_prompt_cache_hit_rate(prompt_cache_hit_rate) metadata["prompt_cache_len"] = origin_prompt_cache_len or 0 if pending_prefill_load_chars is not None: p_node.dispatched_prompt_chars = max( diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index ce186b9f0..adf911bc2 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -10,6 +10,7 @@ - 用前缀树(见 PromptCacheTree)记录「历史 prompt -> 处理它的 worker」; - 树中的 prefill_node 对应 worker.client_ip_port; - prompt 会按 sample_stride 抽稀后再插入/匹配,降低树的深度与内存; + - 根据推理侧返回的平均 prompt cache 命中率,动态调整 cache 亲和与负载均衡的权重; - 优先使用 dispatched_req_num 为 0 的空闲节点,避免 GPU 闲置; - 所有节点都忙时,用 worker.dispatched_prompt_chars 做负载均衡。 @@ -38,6 +39,13 @@ class CacheAwareConfig: cache_threshold: float = 0.5 # cache 命中节点的在途量超过最空闲节点该倍数时,优先选择最空闲节点。 balance_rel_threshold: float = 1.5 + # 动态调整 balance_rel_threshold 时允许的上下限。 + min_balance_rel_threshold: float = 1.0 + max_balance_rel_threshold: float = 2.0 + # 每轮用于统计平均 prompt cache 命中率的请求数。 + cache_hit_rate_window_size: int = 1000 + # 相邻统计周期命中率变化时,负载均衡阈值的调整步长。 + balance_rel_threshold_step: float = 0.05 # 前缀树允许的最大节点数(不含 root)。 max_node_count: int = 1_000_000 # 每次 LRU 驱逐的叶节点数量。 @@ -48,6 +56,41 @@ class CacheAwareConfig: recursion_limit: int = 4000 +class BalanceRelThresholdController: + """根据最近请求的 prompt cache 命中率动态调整负载均衡阈值。""" + + def __init__(self) -> None: + self._cache_hit_rates = [] + self._last_average_cache_hit_rate = None + + def append(self, cache_hit_rate: float) -> None: + """追加一次真实 prompt cache 命中率。""" + cache_hit_rate = min(max(cache_hit_rate, 0.0), 1.0) + self._cache_hit_rates.append(cache_hit_rate) + + def update_config(self, config: CacheAwareConfig) -> None: + """每收集一个统计窗口,根据命中率趋势调整负载均衡阈值。""" + if len(self._cache_hit_rates) < config.cache_hit_rate_window_size: + return + + average_cache_hit_rate = ( + sum(self._cache_hit_rates[-config.cache_hit_rate_window_size :]) / config.cache_hit_rate_window_size + ) + self._cache_hit_rates.clear() + + if self._last_average_cache_hit_rate is not None: + if average_cache_hit_rate > self._last_average_cache_hit_rate: + config.balance_rel_threshold += config.balance_rel_threshold_step + elif average_cache_hit_rate < self._last_average_cache_hit_rate: + config.balance_rel_threshold -= config.balance_rel_threshold_step + config.balance_rel_threshold = min( + max(config.balance_rel_threshold, config.min_balance_rel_threshold), + config.max_balance_rel_threshold, + ) + + self._last_average_cache_hit_rate = average_cache_hit_rate + + class CacheAwarePolicy: """ 维护 prompt 前缀树,并据此为请求选择 prefill worker。 @@ -66,6 +109,7 @@ def __init__(self, config: Optional[CacheAwareConfig] = None) -> None: evict_node_batch=self.config.evict_node_batch, recursion_limit=self.config.recursion_limit, ) + self.balance_rel_threshold_controller = BalanceRelThresholdController() def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """ @@ -107,6 +151,11 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port) return selected_worker + def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: + """记录推理侧上报的真实 cache 命中率,并更新动态负载阈值。""" + self.balance_rel_threshold_controller.append(cache_hit_rate) + self.balance_rel_threshold_controller.update_config(self.config) + def _get_cache_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """在指定候选节点中返回达到匹配阈值的 cache 节点。""" result = self.prompt_cache_tree.prefix_match(request_text) @@ -141,6 +190,7 @@ def _select_idle_worker(self, workers: List[PD_Client_Obj], request_text: str) - logger.info( f"CacheAwarePolicy: select idle worker, idle_worker_num={len(idle_workers)}, " f"cache_worker={cache_worker.client_ip_port if cache_worker else None}, " + f"balance_rel_threshold={self.config.balance_rel_threshold:.4f}, " f"selected_worker={selected_worker.client_ip_port}" ) return selected_worker diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index 8e095099a..5474806b7 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py @@ -26,6 +26,10 @@ def select_p_d_node( ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: raise NotImplementedError("Subclass must implement this method") + def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: + """记录推理侧返回的 prompt cache 命中率;非 cache-aware 策略无需处理。""" + return + class RandomSelector(PDSelector): """随机选择器""" @@ -93,3 +97,6 @@ def select_p_d_node( ) return p_node, d_node + + def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: + self.policy.record_prompt_cache_hit_rate(cache_hit_rate) diff --git a/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index 55d9eb98e..7d0b32ded 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -12,6 +12,7 @@ def _manager() -> HttpServerManagerForPDMaster: manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_manager = MagicMock() manager.id_gen = MagicMock() manager.id_gen.generate_id.return_value = 800 manager.metric_client = MagicMock() diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index 09bf04672..a3dc268a9 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -26,6 +26,10 @@ def gen_id(): mgr.metric_client = SimpleNamespace(counter_inc=lambda *a, **k: None, histogram_observe=lambda *a, **k: None) mgr.tokens = lambda *a, **k: 10 mgr._log_req_header = lambda *a, **k: asyncio.sleep(0) + mgr.recorded_cache_hit_rates = [] + mgr.pd_manager = SimpleNamespace( + selector=SimpleNamespace(record_prompt_cache_hit_rate=mgr.recorded_cache_hit_rates.append) + ) p_node = SimpleNamespace(dispatched_prompt_chars=0, dispatched_req_num=0) mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(p_node, 1)) mgr.remove_req = lambda *a, **k: asyncio.sleep(0) @@ -38,10 +42,10 @@ def _collect(mgr, sampling_params, monkeypatch, split): async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request): sub_req_id = sp.group_request_id hit = sp.max_new_tokens * 10 - yield sub_req_id, "x", {"prompt_tokens": 10, "prompt_cache_len": hit}, FinishStatus() + yield sub_req_id, "x", {"prompt_tokens": 100, "prompt_cache_len": hit}, FinishStatus() for _ in range(2): - yield sub_req_id, "y", {"prompt_tokens": 10, "prompt_cache_len": 0}, FinishStatus() - yield sub_req_id, "z", {"prompt_tokens": 10, "prompt_cache_len": 0}, FinishStatus(FinishStatus.FINISHED_STOP) + yield sub_req_id, "y", {"prompt_tokens": 100, "prompt_cache_len": 0}, FinishStatus() + yield sub_req_id, "z", {"prompt_tokens": 100, "prompt_cache_len": 0}, FinishStatus(FinishStatus.FINISHED_STOP) monkeypatch.setattr(mgr, "_wait_to_token_package", fake_wait) @@ -68,6 +72,7 @@ def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): sp.group_request_id = 0 cached = _collect(mgr, sp, monkeypatch, split=[3]) assert cached and all(c == 30 for c in cached), cached + assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] def test_multi_block_keeps_first_block_hit(monkeypatch): @@ -79,3 +84,4 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): sp.group_request_id = 0 cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, cached + assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 23b6de590..6bdde7857 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -1,6 +1,12 @@ from types import SimpleNamespace -from lightllm.server.httpserver_for_pd_master.pd_selector.cache_aware import CacheAwarePolicy +import pytest + +from lightllm.server.httpserver_for_pd_master.pd_selector.cache_aware import ( + BalanceRelThresholdController, + CacheAwareConfig, + CacheAwarePolicy, +) def _worker(address: str, dispatched_prompt_chars: int = 0, dispatched_req_num: int = 0): @@ -11,6 +17,62 @@ def _worker(address: str, dispatched_prompt_chars: int = 0, dispatched_req_num: ) +def test_balance_threshold_controller_adjusts_threshold_each_window(): + config = CacheAwareConfig(cache_hit_rate_window_size=3, balance_rel_threshold_step=0.1) + controller = BalanceRelThresholdController() + + for cache_hit_rate, expected_threshold in ( + (0.5, 1.5), + (0.6, 1.6), + (0.4, 1.5), + (0.4, 1.5), + ): + for _ in range(config.cache_hit_rate_window_size): + controller.append(cache_hit_rate) + controller.update_config(config) + assert config.balance_rel_threshold == pytest.approx(expected_threshold) + + +@pytest.mark.parametrize( + ("initial_threshold", "first_hit_rate", "second_hit_rate", "expected_threshold"), + ( + (1.95, 0.5, 0.6, 2.0), + (1.05, 0.5, 0.4, 1.0), + ), +) +def test_balance_threshold_controller_limits_threshold_range( + initial_threshold, + first_hit_rate, + second_hit_rate, + expected_threshold, +): + config = CacheAwareConfig( + balance_rel_threshold=initial_threshold, + cache_hit_rate_window_size=1, + balance_rel_threshold_step=0.1, + ) + controller = BalanceRelThresholdController() + + controller.append(first_hit_rate) + controller.update_config(config) + controller.append(second_hit_rate) + controller.update_config(config) + + assert config.balance_rel_threshold == pytest.approx(expected_threshold) + + +def test_cache_aware_updates_threshold_from_inference_cache_hit_rate(): + policy = CacheAwarePolicy(CacheAwareConfig(cache_hit_rate_window_size=2)) + + for _ in range(2): + policy.record_prompt_cache_hit_rate(0.25) + assert policy.config.balance_rel_threshold == 1.5 + for _ in range(2): + policy.record_prompt_cache_hit_rate(0.75) + + assert policy.config.balance_rel_threshold == pytest.approx(1.55) + + def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy = CacheAwarePolicy() cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2)