From 37bab41079c29cff584ec2326d52b4197d423cda Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 7 Aug 2026 09:29:41 +0000 Subject: [PATCH] fix(pd): balance cache-aware prefill dispatch --- .../httpserver_for_pd_master/manager.py | 14 +-- .../pd_selector/cache_aware.py | 90 ++++++++----------- .../pd_selector/pd_selector.py | 3 - lightllm/server/pd_io_struct.py | 2 +- .../test_pd_master_multi_choice.py | 72 ++++++++++++++- .../test_pd_master_cached_tokens.py | 5 +- unit_tests/server/test_pd_cache_aware.py | 46 ++++++++++ unit_tests/server/test_pd_master_mode.py | 8 +- 8 files changed, 171 insertions(+), 69 deletions(-) create mode 100644 unit_tests/server/test_pd_cache_aware.py diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 9b632d744..18cfa0d51 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -211,9 +211,12 @@ async def _generate_one( block_group_request_id = origin_request_id p_node = None d_node = None + prefill_load_released = False 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 = [] @@ -252,6 +255,9 @@ 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 yield origin_request_id, request_output, metadata, finish_status await self.remove_req(group_request_id=block_group_request_id) @@ -271,6 +277,8 @@ 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)) await self.remove_req(block_group_request_id) return @@ -731,12 +739,6 @@ def register_pd(self, pd_info_json, websocket): if pd_client.mode == "prefill": self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] self.prefill_nodes.append(pd_client) - # dispatched_prompt_chars is the cumulative counter used by the - # CacheAware policy to balance request dispatch across prefill nodes. - # Reset all counters together when a node registers or reconnects so - # stale history does not bias CacheAware toward the zero-valued node. - for prefill_node in self.prefill_nodes: - prefill_node.dispatched_prompt_chars = 0 elif pd_client.mode == "decode": self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] self.decode_nodes.append(pd_client) 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 b7b2d525e..ad7348c86 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,7 @@ - 用前缀树(见 PromptCacheTree)记录「历史 prompt -> 处理它的 worker」; - 树中的 prefill_node 对应 worker.client_ip_port; - prompt 会按 sample_stride 抽稀后再插入/匹配,降低树的深度与内存; - - 用 worker.dispatched_prompt_chars(累计派发的 prompt 字符数)做粗粒度均衡。 + - 用 worker.dispatched_prompt_chars(当前尚未产出首 token 的 prompt 字符数)做负载均衡。 选点流程见 CacheAwarePolicy.select_worker。 """ @@ -35,8 +35,8 @@ class CacheAwareConfig: # 前缀匹配成功率阈值:matched_char_count / input_char_count 超过该值才路由到命中节点。 cache_threshold: float = 0.5 - # 派发量不均衡判定:max > min * balance_rel_threshold 时强制选派发量最少的节点。 - balance_rel_threshold: float = 1.2 + # cache 命中节点的在途量超过最空闲节点该倍数时,优先选择最空闲节点。 + balance_rel_threshold: float = 1.8 # 前缀树允许的最大节点数(不含 root)。 max_node_count: int = 1_000_000 # 每次 LRU 驱逐的叶节点数量。 @@ -66,16 +66,6 @@ def __init__(self, config: Optional[CacheAwareConfig] = None) -> None: recursion_limit=self.config.recursion_limit, ) - def _select_worker_min_dispatched( - self, - workers: List[PD_Client_Obj], - request_text: str, - ) -> PD_Client_Obj: - """派发量优先兜底:选择累计 dispatched_prompt_chars 最小的 worker,并写入前缀树。""" - min_dispatched_worker = min(workers, key=lambda worker: worker.dispatched_prompt_chars) - self.prompt_cache_tree.insert(request_text, min_dispatched_worker.client_ip_port) - return min_dispatched_worker - def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """ 为一次请求选择 prefill worker。 @@ -89,38 +79,19 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti 决策顺序: 1) workers 为空 -> 返回 None; - 2) 若 max(dispatched) > min(dispatched) * balance_rel_threshold, - 认为派发不均衡,直接选派发量最少的节点; - 3) 否则对 request_text 做前缀匹配,计算 + 2) 对 request_text 做前缀匹配,计算 match_rate = matched_char_count / input_char_count; - 4) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 路由到该节点并更新树; - 5) 未命中阈值或 prefill_node 不在当前 workers 中 -> 回退到派发量最少选择。 + 3) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 得到 cache 命中节点; + 4) cache 命中节点负载未严重高于最空闲节点 -> 选择 cache 命中节点; + 5) 未命中或负载严重失衡 -> 选择最空闲节点; + 6) 将当前 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 亲和 ---- - dispatched_chars = [worker.dispatched_prompt_chars for worker in workers] - min_dispatched = min(dispatched_chars) if dispatched_chars else 0 - max_dispatched = max(dispatched_chars) if dispatched_chars else 0 - - is_imbalanced = max_dispatched > (min_dispatched * self.config.balance_rel_threshold) - - logger.info( - f"CacheAwarePolicy: min_dispatched={min_dispatched}, max_dispatched={max_dispatched}, " - f"balance_rel_threshold={self.config.balance_rel_threshold:.4f}, " - f"is_imbalanced={is_imbalanced}" - ) - - if is_imbalanced: - return self._select_worker_min_dispatched( - workers=workers, - request_text=request_text, - ) - - # ---- 2. 前缀匹配:估计当前请求与历史请求的 cache 复用潜力 ---- + # ---- 1. 前缀匹配:估计当前请求与历史请求的 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 @@ -131,26 +102,39 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti f"prefill_node={result.prefill_node}" ) - selected_worker: Optional[PD_Client_Obj] = None + cache_worker: Optional[PD_Client_Obj] = None if match_rate > self.config.cache_threshold and result.prefill_node is not None: - # 树中的 prefill_node 是 client_ip_port,需要映射回当前在线 worker 对象。 for worker in workers: if worker.client_ip_port == result.prefill_node: - selected_worker = worker + cache_worker = worker break + # ---- 2. 负载配平:cache 节点过载时选择当前在途量最少的节点 ---- + 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 + if cache_worker is not None: + cache_projected_load = cache_worker.dispatched_prompt_chars + request_load + cache_worker_is_overloaded = cache_projected_load > least_projected_load * self.config.balance_rel_threshold + + if cache_worker is None or cache_worker_is_overloaded: + selected_worker = least_loaded_worker + else: + selected_worker = cache_worker + logger.info( - f"CacheAwarePolicy: selected_worker=" - f"{selected_worker.client_ip_port if selected_worker else None}, " - f"match_rate={match_rate:.4f}, cache_threshold={self.config.cache_threshold:.4f}" + f"CacheAwarePolicy: cache_worker={cache_worker.client_ip_port if cache_worker else None}, " + f"cache_worker_load={cache_worker.dispatched_prompt_chars if cache_worker else None}, " + f"cache_projected_load={cache_projected_load}, " + f"least_loaded_worker={least_loaded_worker.client_ip_port}, " + f"least_loaded_worker_load={least_loaded_worker.dispatched_prompt_chars}, " + f"least_projected_load={least_projected_load}, " + f"balance_rel_threshold={self.config.balance_rel_threshold:.4f}, " + f"cache_worker_is_overloaded={cache_worker_is_overloaded}, " + f"selected_worker={selected_worker.client_ip_port}" ) - # ---- 3. 命中则更新树;未命中则派发量兜底并写入树 ---- - if selected_worker is not None: - self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port) - return selected_worker - else: - return self._select_worker_min_dispatched( - workers=workers, - request_text=request_text, - ) + self.prompt_cache_tree.insert(request_text, 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 d1315ae82..8e095099a 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 @@ -87,9 +87,6 @@ def select_p_d_node( p_node = self.policy.select_worker(self.prefill_nodes, request_text=prompt) d_node = self._importance_sampling(self.decode_nodes) - # 累计派发字符数,供后续 cache-aware 做派发均衡判断。 - p_node.dispatched_prompt_chars += len(prompt) - logger.info( f"LoadBalancedCacheAwareSelector: selected p_node={p_node.client_ip_port}, " f"d_node={d_node.client_ip_port}" diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index d0f32419d..e469adf94 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -56,7 +56,7 @@ class PD_Client_Obj: start_args: object # 节点的启动参数信息,用于做匹配性的校验,防止运行过程中出现问题。 websocket: WebSocket = None # 用于通信的 websocket 连接对象 run_status: _PD_Client_RunStatus = field(default_factory=_PD_Client_RunStatus) - # cache-aware 选点用:累计派发到该节点的 prompt 字符数(只增不减,非实时负载)。 + # cache-aware 选点用:当前派发到该节点且尚未产出首 token 的 prompt 字符数。 dispatched_prompt_chars: int = 0 def __post_init__(self): 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 36945e652..5bad24bdd 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() + p_node = MagicMock(dispatched_prompt_chars=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]) @@ -93,6 +93,7 @@ async def wait_to_token_package( call("lightllm_request_count"), call("lightllm_request_success"), ] + assert p_node.dispatched_prompt_chars == 0 asyncio.run(asyncio.wait_for(run(), timeout=2)) @@ -203,3 +204,72 @@ async def choice(): assert choice_closed.is_set() asyncio.run(asyncio.wait_for(run(), timeout=2)) + + +def test_pd_master_releases_prefill_load_when_generation_fails(): + async def run(): + manager = _manager() + manager._split_max_new_tokens = MagicMock(return_value=[4]) + manager.id_gen.generate_id.return_value = 808 + manager.remove_req = AsyncMock() + manager.abort = AsyncMock() + p_node = MagicMock(dispatched_prompt_chars=0) + d_node = MagicMock() + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + + async def failing_wait_to_token_package(*_args, **_kwargs): + raise RuntimeError("generation failed") + yield None + + manager._wait_to_token_package = failing_wait_to_token_package + + with pytest.raises(RuntimeError, match="generation failed"): + async for _ in manager._generate_one( + "prompt", + SamplingParams(), + MagicMock(), + MagicMock(), + 0, + 800, + ): + pass + + assert p_node.dispatched_prompt_chars == 0 + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + +def test_pd_master_releases_prefill_load_when_stream_is_closed(): + async def run(): + manager = _manager() + manager._split_max_new_tokens = MagicMock(return_value=[4]) + manager.id_gen.generate_id.return_value = 808 + manager.remove_req = AsyncMock() + manager.abort = AsyncMock() + other_request_load = 17 + p_node = MagicMock(dispatched_prompt_chars=other_request_load) + d_node = MagicMock() + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + + async def wait_to_token_package(*_args, **_kwargs): + yield 808, "first", {"prompt_tokens": 1}, FinishStatus() + await asyncio.sleep(10) + + manager._wait_to_token_package = wait_to_token_package + generator = manager._generate_one( + "prompt", + SamplingParams(), + MagicMock(), + MagicMock(), + 0, + 800, + ) + + assert (await generator.__anext__())[1] == "first" + assert p_node.dispatched_prompt_chars == other_request_load + await generator.aclose() + + assert p_node.dispatched_prompt_chars == other_request_load + 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 1e65ec990..9fa60d26e 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,8 @@ 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.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(1, 1)) + p_node = SimpleNamespace(dispatched_prompt_chars=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 @@ -61,6 +62,7 @@ async def run(): def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): mgr = _make_manager(monkeypatch) sp = SamplingParams() + sp.n = 1 sp.max_new_tokens = 3 sp.best_of = 1 sp.group_request_id = 0 @@ -71,6 +73,7 @@ def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): def test_multi_block_keeps_first_block_hit(monkeypatch): mgr = _make_manager(monkeypatch) sp = SamplingParams() + sp.n = 1 sp.max_new_tokens = 5 sp.best_of = 1 sp.group_request_id = 0 diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py new file mode 100644 index 000000000..665311767 --- /dev/null +++ b/unit_tests/server/test_pd_cache_aware.py @@ -0,0 +1,46 @@ +from types import SimpleNamespace + +from lightllm.server.httpserver_for_pd_master.pd_selector.cache_aware import CacheAwarePolicy + + +def _worker(address: str, dispatched_prompt_chars: int = 0): + return SimpleNamespace( + client_ip_port=address, + dispatched_prompt_chars=dispatched_prompt_chars, + ) + + +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) + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) + + selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt) + + assert selected_worker is cache_worker + + +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) + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) + + selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt) + + assert selected_worker is least_loaded_worker + + +def test_cache_aware_keeps_cache_worker_when_both_workers_are_idle(): + policy = CacheAwarePolicy() + cache_worker = _worker("10.0.0.1:8000") + other_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([other_worker, cache_worker], prompt) + + assert selected_worker is cache_worker diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 14e4017a9..777d41412 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -184,7 +184,7 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True -def test_prefill_registration_resets_all_dispatched_prompt_chars(): +def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) @@ -208,10 +208,10 @@ def register_prefill(node_id, client_ip_port): register_prefill(2, "10.0.0.2:8000") - assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [0, 0] + assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [1234, 0] -def test_prefill_reconnection_resets_all_dispatched_prompt_chars(): +def test_prefill_reconnection_preserves_other_nodes_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) @@ -235,7 +235,7 @@ def pd_info(node_id, client_ip_port): 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] == [0, 0] + assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [100, 0] def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch):