From 7830122210a339598d90f56e7db53bb35fc7a836 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Tue, 1 Sep 2026 10:39:07 +0800 Subject: [PATCH 01/17] feat: dynamically split PD decode requests --- deploy_dspark_1p1d.sh | 1 - .../httpserver_for_pd_master/manager.py | 49 ++++--- .../model_infer/mode_backend/base_backend.py | 19 ++- .../pd/decode_node_impl/decode_impl.py | 37 +++++ .../pd/decode_node_impl/decode_impl_for_dp.py | 6 + lightllm/server/router/req_queue/__init__.py | 12 +- .../router/req_queue/chunked_prefill/impl.py | 11 +- .../req_queue/chunked_prefill/impl_for_pd.py | 11 +- lightllm/utils/envs_utils.py | 5 - .../test_pd_master_multi_choice.py | 55 +++++--- .../test_pd_master_cached_tokens.py | 48 +++++-- .../mode_backend/test_pd_dynamic_split.py | 127 ++++++++++++++++++ 12 files changed, 304 insertions(+), 77 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py diff --git a/deploy_dspark_1p1d.sh b/deploy_dspark_1p1d.sh index 5a8e1996d3..0dcb04432d 100755 --- a/deploy_dspark_1p1d.sh +++ b/deploy_dspark_1p1d.sh @@ -28,7 +28,6 @@ export LIGHTLLM_TRITON_AUTOTUNE_LEVEL="${LIGHTLLM_TRITON_AUTOTUNE_LEVEL:-1}" export LIGHTLLM_ANTHROPIC_ENABLE_PDF_PARSING="${LIGHTLLM_ANTHROPIC_ENABLE_PDF_PARSING:-1}" export LIGHTLLM_LOG_LEVEL="${LIGHTLLM_LOG_LEVEL:-debug}" export PYTHONUNBUFFERED=1 -export LIGHTLLM_PD_SPLIT_MAX_NEW_TOKENS="${LIGHTLLM_PD_SPLIT_MAX_NEW_TOKENS:-4096}" # Keep same-host PD WebSocket traffic away from HTTP/WebSocket proxies. export NO_PROXY="${NO_PROXY:+${NO_PROXY},}${PD_MASTER_IP},localhost" diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 80c114e116..d4b4f24cfe 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -23,7 +23,7 @@ from lightllm.utils.statics_utils import MovingAverage from lightllm.server.httpserver.manager import AsyncQueue from lightllm.utils.error_utils import ClientDisconnected, ServerBusyError -from lightllm.utils.envs_utils import get_pd_high_priority_request_timeout_seconds, get_pd_split_max_new_tokens +from lightllm.utils.envs_utils import get_pd_high_priority_request_timeout_seconds from lightllm.utils.shm_port_args import get_shm_port_args from .pd_selector import create_selector @@ -215,13 +215,6 @@ async def _generate_one( start_time: float, origin_request_id: int, ): - # 先将请求根据max_new_tokens 参数进行分块操作,主要是 pd 分离场景中, - # 只能使用保守调度,但是如果用户都设置一个很大的 max_new_tokens 值,会 - # 导致极大显存预留,照成系统的吞吐能力下降,所以我们将请求分割成几段进行 - # 推理,只要保证分块合理,实际分段推理是极少发生的情况,系统吞吐就不会受 - # 到影响。 - max_new_tokens_list = self._split_max_new_tokens(max_new_tokens=origin_sampling_params.max_new_tokens) - block_group_request_id = origin_request_id p_node = None d_node = None @@ -237,17 +230,22 @@ async def _generate_one( history_gen_token_strs = [] origin_prompt_cache_len = None + remaining_max_new_tokens = origin_sampling_params.max_new_tokens + segment_index = 0 - for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list): + # Decode 节点容量不足时会将已运行请求以 length 结束。 + # PD master 把这个内部 length 当作动态分段边界,用剩余 token + # 限额在同一组 P/D 节点上继续。 + while remaining_max_new_tokens > 0: sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) block_group_request_id = self.id_gen.generate_id() sampling_params.group_request_id = block_group_request_id 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 + sampling_params.max_new_tokens = remaining_max_new_tokens # 预计输入 cache 命中率高于 0.8 时,将首段也标记为 PD 高优先级请求, # 使其优先进入 Router 调度队列,尽快复用已命中的 KV cache。第二段及后续 # 分段仍统一使用高优先级,避免因临时资源紧张导致分段续跑失败。 - sampling_params.pd_high_priority_request = iter_index > 0 or estimated_cache_hit_rate > 0.8 + sampling_params.pd_high_priority_request = segment_index > 0 or estimated_cache_hit_rate > 0.8 # 为高优先级请求下发较长的有限等待时间;P/D 节点仅在自身开启 # 本地限流时使用该值,未开启限流时仍保持无限等待。 if sampling_params.pd_high_priority_request: @@ -270,17 +268,23 @@ async def _generate_one( multimodal_params, request, ) - is_last_block = iter_index == len(max_new_tokens_list) - 1 prompt_tokens = sys.maxsize # 因为分段的原因 - async for sub_req_id, request_output, metadata, finish_status in results_generator: + segment_output_tokens = 0 + segment_finish_status = FinishStatus() + async for sub_req_id, request_output, metadata, raw_finish_status in results_generator: # pd 分离模式下,返回的 metadata 可能序号信息可能存在不准确性。 assert sub_req_id == block_group_request_id - if finish_status.is_finished_length() and not is_last_block: - finish_status = FinishStatus() # 转换为NoFinished + segment_output_tokens = metadata["count_output_tokens"] + + segment_finish_status = raw_finish_status + finish_status = raw_finish_status + if raw_finish_status.is_finished_length() and segment_output_tokens < remaining_max_new_tokens: + # Decode 节点因容量不足产生的内部分段,不暴露给用户。 + finish_status = FinishStatus() history_gen_token_strs.append(request_output) prompt_tokens = min(prompt_tokens, metadata["prompt_tokens"]) metadata["prompt_tokens"] = prompt_tokens - if iter_index == 0 and origin_prompt_cache_len is None: + if 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) @@ -298,7 +302,10 @@ async def _generate_one( yield origin_request_id, request_output, metadata, finish_status await self.remove_req(group_request_id=block_group_request_id) - if finish_status.is_finished(): + if segment_finish_status.is_finished_length(): + remaining_max_new_tokens -= segment_output_tokens + segment_index += 1 + else: break except (ClientDisconnected, BaseException) as e: @@ -718,14 +725,6 @@ async def handle_loop(self): logger.exception(str(e)) return - def _split_max_new_tokens(self, max_new_tokens: int) -> List[int]: - block_max_new_tokens = get_pd_split_max_new_tokens() - ans_list = [block_max_new_tokens for _ in range(max_new_tokens // block_max_new_tokens)] - left_token = max_new_tokens - (max_new_tokens // block_max_new_tokens) * block_max_new_tokens - if left_token > 0: - ans_list.append(left_token) - return ans_list - class ReqStatus: def __init__(self, req_id, p_node, d_node) -> None: diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index c64cf906bf..1a3b265e14 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -740,8 +740,7 @@ def _get_classed_reqs( decode_reqs.append(req_obj) can_alloc_token_num -= token_num else: - if wait_pause_count < pause_max_req_num: - req_obj.wait_pause = True + if wait_pause_count < pause_max_req_num and self._handle_decode_alloc_failure(req_obj): wait_pause_count += 1 else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 @@ -784,7 +783,7 @@ def _get_classed_reqs( true_finished_reqs = finished_reqs g_infer_context.filter_reqs(finished_reqs=true_finished_reqs) - g_infer_context.pause_reqs(wait_pause_reqs, is_master_in_dp=self.is_master_in_dp) + self._pause_reqs(wait_pause_reqs) if recover_paused: g_infer_context.recover_paused_reqs( @@ -806,6 +805,20 @@ def _get_classed_reqs( return prefill_reqs, decode_reqs + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + """Handle a decode request that cannot allocate its next token. + + Returns whether this request consumed one slot of the per-iteration + shortage handling limit. Backends can customize the eventual pause + action through ``_pause_reqs``. + """ + req_obj.wait_pause = True + return True + + def _pause_reqs(self, pause_reqs: List[InferReq]): + g_infer_context.pause_reqs(pause_reqs, is_master_in_dp=self.is_master_in_dp) + return + # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 472049442f..07ebdd96a7 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -119,6 +119,43 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: ans_list.append(req_obj) return ans_list + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + """Send an actually running request through the overlap-safe pause path.""" + if not self.support_overlap: + return super()._handle_decode_alloc_failure(req_obj) + if req_obj.cur_output_len <= req_obj.shm_req.shm_cur_output_len: + return False + + req_obj.wait_pause = True + logger.info( + f"wait to yield running pd decode req_id={req_obj.req_id} " + f"at output_len={req_obj.cur_output_len} because token memory is insufficient" + ) + return True + + def _pause_reqs(self, pause_reqs: List[InferReq]): + """Finish PD Decode segments instead of preserving them for recovery.""" + for req_obj in pause_reqs: + req_obj.wait_pause = False + if req_obj.finish_status.is_finished(): + continue + + req_obj.finish_status.set_status(FinishStatus.FINISHED_LENGTH) + if self.is_master_in_dp: + shm_req = req_obj.shm_req + shm_req.shm_cur_output_len = req_obj.cur_output_len + shm_req.finish_token_index = shm_req.input_len + req_obj.cur_output_len - 1 + shm_req.finish_status = req_obj.finish_status + shm_req.candetoken_out_len = req_obj.cur_output_len + + logger.info( + f"yield pd decode req_id={req_obj.req_id} " + f"at output_len={req_obj.cur_output_len} because token memory is insufficient" + ) + + g_infer_context.filter_reqs(finished_reqs=pause_reqs) + return + def _decode_node_gen_trans_tasks(self, req_obj: InferReq): """ decode node 生成所有的传输任务对象。 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py index 87af300003..b2f652e032 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py @@ -27,6 +27,12 @@ def _post_init_reqs(self, uninit_reqs: List[InferReq]): def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: return PDDecodeNode._filter_not_ready_reqs(self, req_ids=req_ids) + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + return PDDecodeNode._handle_decode_alloc_failure(self, req_obj=req_obj) + + def _pause_reqs(self, pause_reqs: List[InferReq]): + return PDDecodeNode._pause_reqs(self, pause_reqs=pause_reqs) + def _decode_node_gen_trans_tasks(self, req_obj: InferReq): return PDDecodeNode._decode_node_gen_trans_tasks(self, req_obj=req_obj) diff --git a/lightllm/server/router/req_queue/__init__.py b/lightllm/server/router/req_queue/__init__.py index c0de01db28..5ca60eeeb0 100644 --- a/lightllm/server/router/req_queue/__init__.py +++ b/lightllm/server/router/req_queue/__init__.py @@ -1,19 +1,23 @@ from .chunked_prefill.impl import ChunkedPrefillQueue from .chunked_prefill.beam_impl import ChunkedBeamContinuesBatchQueue -from .chunked_prefill.impl_for_pd import PDQueue +from .chunked_prefill.impl_for_pd import PDDQueue, PDPQueue from .dp_base_queue import DpQueue def _get_req_queue_class(args, router, dp_size_in_node: int): + # model_rpc 会优先为 PD 节点选择 PD backend,queue 也必须保持 + # 同样的优先级,避免其他模式开关绕过 Decode 的激进调度。 + if args.run_mode == "decode": + return PDDQueue + if args.run_mode == "prefill": + return PDPQueue + if args.diverse_mode: return ChunkedBeamContinuesBatchQueue if args.output_constraint_mode != "none": return ChunkedPrefillQueue if args.first_token_constraint_mode: return ChunkedPrefillQueue - if args.run_mode in ["prefill", "decode"]: - return PDQueue - if args.disable_chunked_prefill: # 虽然也使用chuncked prefill queue 但是由于 args.chunked_prefill_size = args.max_req_total_len # 所以调度的实际行为类似过去的 continues batch 调度,所以将两种调度的实现统一为一种实现,减少代码重复。 diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py index f8fe510989..160c086b83 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py @@ -20,6 +20,12 @@ def _init_cache_list(self, current_batch: Batch, is_busy): self.cache_len_list = [] return + def _can_add_first_router_tokens(self, first_router_need_tokens: int) -> bool: + return ( + self.args.short_prefill_token_threshold is not None + or first_router_need_tokens <= self.batch_max_tokens + ) + # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens): self.cache_len_list.append( @@ -39,10 +45,7 @@ def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens new_batch_first_router_need_tokens += req.get_first_router_need_tokens() # 长短请求模式由 Infer 控制单轮 prefill token 上限。 - ok_prefill = ( - self.args.short_prefill_token_threshold is not None - or new_batch_first_router_need_tokens <= self.batch_max_tokens - ) + ok_prefill = self._can_add_first_router_tokens(new_batch_first_router_need_tokens) if ok_token_num and ok_req_num and ok_prefill: self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index f6e33144ef..bb0a4b90ce 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py @@ -3,9 +3,10 @@ from typing import Tuple from ...batch import Batch, Req from lightllm.server.router.req_queue.base_queue import BaseQueue +from lightllm.server.router.req_queue.chunked_prefill.impl import ChunkedPrefillQueue -class PDQueue(BaseQueue): +class PDPQueue(BaseQueue): def __init__(self, args, router, dp_index, dp_size_in_node) -> None: super().__init__(args, router, dp_index, dp_size_in_node) @@ -93,3 +94,11 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): estimated_peak_token_num = self._caclu_batch_estimated_peak_token_num(current_batch) return (estimated_peak_token_num, estimated_peak_token_num / self.max_total_tokens) + + +class PDDQueue(ChunkedPrefillQueue): + def is_busy(self): + return False + + def _can_add_first_router_tokens(self, first_router_need_tokens: int) -> bool: + return True diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index fb30f47938..ee1d0bda5a 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -306,11 +306,6 @@ def _get_mtp_draft_backbone_layer_num(draft_model_dir: str) -> int: return int(layer_num) -@lru_cache(maxsize=None) -def get_pd_split_max_new_tokens() -> int: - return int(os.getenv("LIGHTLLM_PD_SPLIT_MAX_NEW_TOKENS", 2048)) - - @lru_cache(maxsize=None) def get_pd_node_shm_req_alloc_timeout_seconds() -> int: """PD 节点申请 ``shm_req`` 对象的最长等待时间,单位为秒。""" 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 98e0014602..5d04494390 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -41,7 +41,6 @@ async def run(): 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, 0.0)) - manager._split_max_new_tokens = MagicMock(return_value=[4]) manager.remove_req = AsyncMock() async def wait_to_token_package( @@ -63,7 +62,7 @@ async def wait_to_token_package( yield ( choice_sampling_params.group_request_id, f"internal-{choice_sampling_params.group_request_id}-{token_index}", - {"prompt_tokens": 2}, + {"prompt_tokens": 2, "count_output_tokens": token_index + 1}, finish_status, ) @@ -250,7 +249,6 @@ async def choice(): 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() @@ -263,12 +261,14 @@ async def failing_wait_to_token_package(*_args, **_kwargs): yield None manager._wait_to_token_package = failing_wait_to_token_package + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 4 with pytest.raises(RuntimeError, match="generation failed"): # 空 prompt 的字符负载为 0,但已派发请求数仍必须在异常路径释放。 async for _ in manager._generate_one( "", - SamplingParams(), + sampling_params, MagicMock(), MagicMock(), 0, @@ -283,10 +283,9 @@ async def failing_wait_to_token_package(*_args, **_kwargs): asyncio.run(asyncio.wait_for(run(), timeout=2)) -def test_pd_master_accounts_each_split_prefill_on_the_same_node(): +def test_pd_master_dynamic_split_reuses_nodes_with_remaining_length(): 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() @@ -299,46 +298,59 @@ async def run(): d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node, 0.0)) dispatched_nodes = [] + dispatched_d_nodes = [] dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] high_priority_request_flags = [] + dispatched_max_new_tokens = [] - async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args): + async def wait_to_token_package( + selected_p_node, selected_d_node, _start_time, block_prompt, sampling_params, *_args + ): dispatched_nodes.append(selected_p_node) + dispatched_d_nodes.append(selected_d_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) high_priority_request_flags.append(sampling_params.pd_high_priority_request) + dispatched_max_new_tokens.append(sampling_params.max_new_tokens) yield ( sampling_params.group_request_id, "x", - {"prompt_tokens": 1}, + {"prompt_tokens": 1, "count_output_tokens": 1}, FinishStatus(FinishStatus.FINISHED_LENGTH), ) manager._wait_to_token_package = wait_to_token_package results = [] + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 2 + multimodal_params = MagicMock() async for result in manager._generate_one( "prompt", - SamplingParams(), - MagicMock(), + sampling_params, + multimodal_params, MagicMock(), 0, 800, ): results.append(result) - manager.select_p_d_node.assert_awaited_once() + manager.select_p_d_node.assert_awaited_once_with("prompt", sampling_params, multimodal_params) assert dispatched_nodes == [p_node, p_node] + assert dispatched_d_nodes == [d_node, d_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 high_priority_request_flags == [False, True] + assert dispatched_max_new_tokens == [2, 1] assert p_node.dispatched_prompt_chars == other_request_load assert p_node.dispatched_req_num == other_request_count assert len(results) == 2 + assert not results[0][3].is_finished() + assert results[1][3].is_finished_length() asyncio.run(asyncio.wait_for(run(), timeout=2)) @@ -353,7 +365,6 @@ def test_pd_master_promotes_request_with_high_estimated_cache_hit_rate( ): async def run(): manager = _manager() - manager._split_max_new_tokens = MagicMock(return_value=[1]) manager.id_gen.generate_id.return_value = 808 manager.remove_req = AsyncMock() manager.abort = AsyncMock() @@ -367,15 +378,17 @@ async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling yield ( sampling_params.group_request_id, "x", - {"prompt_tokens": 1}, + {"prompt_tokens": 1, "count_output_tokens": 1}, FinishStatus(FinishStatus.FINISHED_STOP), ) manager._wait_to_token_package = wait_to_token_package + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 1 async for _ in manager._generate_one( "prompt", - SamplingParams(), + sampling_params, MagicMock(), MagicMock(), 0, @@ -392,7 +405,6 @@ def test_pd_master_sets_high_priority_timeout(): async def run(): manager = _manager() manager.pd_high_priority_request_time_out_seconds = 90 - manager._split_max_new_tokens = MagicMock(return_value=[1]) manager.id_gen.generate_id.return_value = 808 manager.remove_req = AsyncMock() manager.abort = AsyncMock() @@ -409,15 +421,17 @@ async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling yield ( sampling_params.group_request_id, "x", - {"prompt_tokens": 1}, + {"prompt_tokens": 1, "count_output_tokens": 1}, FinishStatus(FinishStatus.FINISHED_STOP), ) manager._wait_to_token_package = wait_to_token_package + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 1 async for _ in manager._generate_one( "prompt", - SamplingParams(), + sampling_params, MagicMock(), MagicMock(), 0, @@ -433,7 +447,6 @@ async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling 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() @@ -447,13 +460,15 @@ async def run(): manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node, 0.0)) async def wait_to_token_package(*_args, **_kwargs): - yield 808, "first", {"prompt_tokens": 1}, FinishStatus() + yield 808, "first", {"prompt_tokens": 1, "count_output_tokens": 1}, FinishStatus() await asyncio.sleep(10) manager._wait_to_token_package = wait_to_token_package + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 4 generator = manager._generate_one( "prompt", - SamplingParams(), + sampling_params, MagicMock(), MagicMock(), 0, 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 d39ff1cc73..416894aac5 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -15,6 +15,8 @@ def _make_manager(monkeypatch): ) monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) + mgr.args = SimpleNamespace(disable_pd_master_decode_capacity_limit=True) + mgr.pd_high_priority_request_time_out_seconds = 60 mgr.running_request_count = 0 counter = [0] @@ -40,16 +42,27 @@ def gen_id(): return mgr -def _collect(mgr, sampling_params, monkeypatch, split): - mgr._split_max_new_tokens = lambda *a, **k: list(split) +def _collect(mgr, sampling_params, monkeypatch, segments): + segment_iter = iter(segments) 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": 100, "prompt_cache_len": hit}, FinishStatus() - for _ in range(2): - 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) + token_count, final_status = next(segment_iter) + hit = sampling_params.max_new_tokens * 10 + for token_index in range(1, token_count + 1): + finish_status = FinishStatus() + if token_index == token_count: + finish_status = FinishStatus(final_status) + yield ( + sub_req_id, + "x", + { + "prompt_tokens": 100, + "prompt_cache_len": hit if token_index == 1 else 0, + "count_output_tokens": token_index, + }, + finish_status, + ) monkeypatch.setattr(mgr, "_wait_to_token_package", fake_wait) @@ -74,30 +87,37 @@ def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): sp.max_new_tokens = 3 sp.best_of = 1 sp.group_request_id = 0 - cached = _collect(mgr, sp, monkeypatch, split=[3]) + cached = _collect(mgr, sp, monkeypatch, segments=[(3, FinishStatus.FINISHED_STOP)]) assert cached and all(c == 30 for c in cached), cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] assert len(mgr.inserted_prompt_caches) == 1 assert mgr.inserted_prompt_caches[0][0] == "hello" -def test_multi_block_keeps_first_block_hit(monkeypatch): +def test_dynamic_split_keeps_first_segment_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 - cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) - assert cached[-1] == 30, cached - assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] + cached = _collect( + mgr, + sp, + monkeypatch, + segments=[ + (3, FinishStatus.FINISHED_LENGTH), + (1, FinishStatus.FINISHED_STOP), + ], + ) + assert cached and all(c == 50 for c in cached), cached + assert mgr.recorded_cache_hit_rates == [pytest.approx(0.5)] assert len(mgr.inserted_prompt_caches) == 1 assert mgr.inserted_prompt_caches[0][0] == "hello" def test_error_result_records_hit_rate_without_inserting_prompt_cache(monkeypatch): mgr = _make_manager(monkeypatch) - mgr._split_max_new_tokens = lambda *a, **k: [1] sampling_params = SamplingParams() sampling_params.n = 1 sampling_params.best_of = 1 @@ -107,7 +127,7 @@ async def failed_wait(_p_node, _d_node, _start_time, _prompt, sp, *_args): yield ( sp.group_request_id, "", - {"prompt_tokens": 100, "prompt_cache_len": 20}, + {"prompt_tokens": 100, "prompt_cache_len": 20, "count_output_tokens": 0}, FinishStatus(FinishStatus.FINISHED_ERROR), ) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py new file mode 100644 index 0000000000..61ca166deb --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -0,0 +1,127 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock + +from lightllm.server.core.objs import FinishStatus +from lightllm.server.router.model_infer.mode_backend.pd.decode_node_impl import ( + decode_impl as pd_decode_impl, +) +from lightllm.server.router.req_queue import _get_req_queue_class +from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd import ( + PDDQueue, + PDPQueue, +) + + +def _make_infer_req(cur_output_len: int, shm_output_len: int): + return SimpleNamespace( + req_id=1, + cur_output_len=cur_output_len, + shm_req=SimpleNamespace(shm_cur_output_len=shm_output_len), + wait_pause=False, + ) + + +def test_pd_decode_capacity_yield_only_pauses_running_overlap_output(): + backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) + backend.support_overlap = True + + running_req = _make_infer_req(cur_output_len=5, shm_output_len=4) + not_running_req = _make_infer_req(cur_output_len=4, shm_output_len=4) + + assert backend._handle_decode_alloc_failure(running_req) + assert running_req.wait_pause + + assert not backend._handle_decode_alloc_failure(not_running_req) + assert not not_running_req.wait_pause + + +def test_pd_decode_capacity_yield_preserves_pause_without_overlap(): + backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) + backend.support_overlap = False + req = _make_infer_req(cur_output_len=5, shm_output_len=5) + req.wait_pause = False + + assert backend._handle_decode_alloc_failure(req) + assert req.wait_pause + + +def test_pd_decode_pause_strategy_finishes_at_committed_mtp_tail(monkeypatch): + backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) + backend.is_master_in_dp = True + filter_reqs = MagicMock() + monkeypatch.setattr( + pd_decode_impl, + "g_infer_context", + SimpleNamespace(filter_reqs=filter_reqs), + ) + + shm_req = SimpleNamespace( + input_len=10, + shm_cur_output_len=3, + finish_token_index=-1, + finish_status=FinishStatus(), + candetoken_out_len=3, + ) + req = SimpleNamespace( + req_id=1, + cur_output_len=3, + finish_status=FinishStatus(), + shm_req=shm_req, + wait_pause=True, + ) + + backend._pause_reqs([req]) + + assert not req.wait_pause + assert req.finish_status.is_finished_length() + assert shm_req.finish_token_index == 12 + assert shm_req.finish_status.is_finished_length() + assert shm_req.candetoken_out_len == 3 + filter_reqs.assert_called_once_with(finished_reqs=[req]) + + +def test_pd_queue_admission_always_uses_aggressive_tuple_estimate(): + queue = PDDQueue.__new__(PDDQueue) + queue.dp_index = 0 + queue.max_total_tokens = 100 + queue.running_max_req_size = 8 + queue.batch_max_tokens = 1 + queue.cache_len_list = [(20, 10)] + queue.router = SimpleNamespace( + router_statics=SimpleNamespace(ema_req_out_len=128), + shared_token_load=SimpleNamespace( + set_estimated_peak_token_count=MagicMock(), + set_dynamic_max_load=MagicMock(), + ), + ) + + req = SimpleNamespace( + get_tuple_tokens=MagicMock(return_value=(20, 10)), + get_first_router_need_tokens=MagicMock(return_value=1024), + ) + admitted, first_router_tokens = queue._can_add_new_req(req, queue.is_busy(), 0) + + assert admitted + assert first_router_tokens == 1024 + req.get_tuple_tokens.assert_called_once_with(False, 128) + queue.router.shared_token_load.set_estimated_peak_token_count.assert_called_once_with( + 60, 0 + ) + + +def test_aggressive_pd_queue_is_only_selected_for_decode_nodes(): + base_args = { + "diverse_mode": False, + "token_healing_mode": False, + "output_constraint_mode": "none", + "first_token_constraint_mode": False, + "disable_chunked_prefill": False, + } + + prefill_args = SimpleNamespace(**base_args, run_mode="prefill") + decode_args = SimpleNamespace(**base_args, run_mode="decode") + + assert ( + _get_req_queue_class(prefill_args, router=None, dp_size_in_node=1) is PDPQueue + ) + assert _get_req_queue_class(decode_args, router=None, dp_size_in_node=1) is PDDQueue From e94865db8912c893bce37489ea8a0c7989d6912a Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:59:55 +0800 Subject: [PATCH 02/17] refactor: end PD chunks through max token limit --- .../model_infer/mode_backend/base_backend.py | 9 +-- .../pd/decode_node_impl/decode_impl.py | 30 ++------- .../pd/decode_node_impl/decode_impl_for_dp.py | 3 - .../mode_backend/test_pd_dynamic_split.py | 64 +++++++++++-------- 4 files changed, 42 insertions(+), 64 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 1a3b265e14..8da4efbf07 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -783,7 +783,7 @@ def _get_classed_reqs( true_finished_reqs = finished_reqs g_infer_context.filter_reqs(finished_reqs=true_finished_reqs) - self._pause_reqs(wait_pause_reqs) + g_infer_context.pause_reqs(wait_pause_reqs, is_master_in_dp=self.is_master_in_dp) if recover_paused: g_infer_context.recover_paused_reqs( @@ -809,16 +809,11 @@ def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: """Handle a decode request that cannot allocate its next token. Returns whether this request consumed one slot of the per-iteration - shortage handling limit. Backends can customize the eventual pause - action through ``_pause_reqs``. + shortage handling limit. """ req_obj.wait_pause = True return True - def _pause_reqs(self, pause_reqs: List[InferReq]): - g_infer_context.pause_reqs(pause_reqs, is_master_in_dp=self.is_master_in_dp) - return - # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 07ebdd96a7..8412ec6e2e 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -120,42 +120,20 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: return ans_list def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: - """Send an actually running request through the overlap-safe pause path.""" + """End an actually running request on its last overlapped output token.""" if not self.support_overlap: return super()._handle_decode_alloc_failure(req_obj) if req_obj.cur_output_len <= req_obj.shm_req.shm_cur_output_len: return False - req_obj.wait_pause = True + sampling_params = req_obj.sampling_param.shm_param + sampling_params.max_new_tokens = min(sampling_params.max_new_tokens, req_obj.cur_output_len) logger.info( - f"wait to yield running pd decode req_id={req_obj.req_id} " + f"yield running pd decode req_id={req_obj.req_id} " f"at output_len={req_obj.cur_output_len} because token memory is insufficient" ) return True - def _pause_reqs(self, pause_reqs: List[InferReq]): - """Finish PD Decode segments instead of preserving them for recovery.""" - for req_obj in pause_reqs: - req_obj.wait_pause = False - if req_obj.finish_status.is_finished(): - continue - - req_obj.finish_status.set_status(FinishStatus.FINISHED_LENGTH) - if self.is_master_in_dp: - shm_req = req_obj.shm_req - shm_req.shm_cur_output_len = req_obj.cur_output_len - shm_req.finish_token_index = shm_req.input_len + req_obj.cur_output_len - 1 - shm_req.finish_status = req_obj.finish_status - shm_req.candetoken_out_len = req_obj.cur_output_len - - logger.info( - f"yield pd decode req_id={req_obj.req_id} " - f"at output_len={req_obj.cur_output_len} because token memory is insufficient" - ) - - g_infer_context.filter_reqs(finished_reqs=pause_reqs) - return - def _decode_node_gen_trans_tasks(self, req_obj: InferReq): """ decode node 生成所有的传输任务对象。 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py index b2f652e032..08dde3de90 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py @@ -30,9 +30,6 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: return PDDecodeNode._handle_decode_alloc_failure(self, req_obj=req_obj) - def _pause_reqs(self, pause_reqs: List[InferReq]): - return PDDecodeNode._pause_reqs(self, pause_reqs=pause_reqs) - def _decode_node_gen_trans_tasks(self, req_obj: InferReq): return PDDecodeNode._decode_node_gen_trans_tasks(self, req_obj=req_obj) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 61ca166deb..09ae766fc3 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -17,11 +17,14 @@ def _make_infer_req(cur_output_len: int, shm_output_len: int): req_id=1, cur_output_len=cur_output_len, shm_req=SimpleNamespace(shm_cur_output_len=shm_output_len), + sampling_param=SimpleNamespace( + shm_param=SimpleNamespace(max_new_tokens=65535), + ), wait_pause=False, ) -def test_pd_decode_capacity_yield_only_pauses_running_overlap_output(): +def test_pd_decode_capacity_yield_only_limits_running_overlap_output(): backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) backend.support_overlap = True @@ -29,9 +32,11 @@ def test_pd_decode_capacity_yield_only_pauses_running_overlap_output(): not_running_req = _make_infer_req(cur_output_len=4, shm_output_len=4) assert backend._handle_decode_alloc_failure(running_req) - assert running_req.wait_pause + assert running_req.sampling_param.shm_param.max_new_tokens == 5 + assert not running_req.wait_pause assert not backend._handle_decode_alloc_failure(not_running_req) + assert not_running_req.sampling_param.shm_param.max_new_tokens == 65535 assert not not_running_req.wait_pause @@ -45,39 +50,42 @@ def test_pd_decode_capacity_yield_preserves_pause_without_overlap(): assert req.wait_pause -def test_pd_decode_pause_strategy_finishes_at_committed_mtp_tail(monkeypatch): +def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) - backend.is_master_in_dp = True - filter_reqs = MagicMock() - monkeypatch.setattr( - pd_decode_impl, - "g_infer_context", - SimpleNamespace(filter_reqs=filter_reqs), - ) + backend.support_overlap = True + req = _make_infer_req(cur_output_len=3, shm_output_len=1) + req.stop_sequences = [] + req.shm_req.input_len = 10 + req.shm_req.shm_prompt_ids = SimpleNamespace(arr=[0] * 13) + req.sampling_param.shm_param.ignore_eos = False + req.finish_status = FinishStatus() + req._stop_sequences_matched = MagicMock(return_value=False) - shm_req = SimpleNamespace( - input_len=10, - shm_cur_output_len=3, - finish_token_index=-1, - finish_status=FinishStatus(), - candetoken_out_len=3, - ) + assert backend._handle_decode_alloc_failure(req) + assert req.sampling_param.shm_param.max_new_tokens == 3 + + pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=2) + assert not req.finish_status.is_finished() + + pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=3) + assert req.finish_status.is_finished_length() + + +def test_pd_decode_capacity_limit_never_extends_original_length(): + backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) + backend.support_overlap = True req = SimpleNamespace( req_id=1, - cur_output_len=3, + cur_output_len=5, + sampling_param=SimpleNamespace( + shm_param=SimpleNamespace(max_new_tokens=3), + ), + shm_req=SimpleNamespace(shm_cur_output_len=1), finish_status=FinishStatus(), - shm_req=shm_req, - wait_pause=True, ) - backend._pause_reqs([req]) - - assert not req.wait_pause - assert req.finish_status.is_finished_length() - assert shm_req.finish_token_index == 12 - assert shm_req.finish_status.is_finished_length() - assert shm_req.candetoken_out_len == 3 - filter_reqs.assert_called_once_with(finished_reqs=[req]) + assert backend._handle_decode_alloc_failure(req) + assert req.sampling_param.shm_param.max_new_tokens == 3 def test_pd_queue_admission_always_uses_aggressive_tuple_estimate(): From 1d0e7fc8ea5c0d17dfa7ad8bf8dbaf2652d1644d Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:06:22 +0800 Subject: [PATCH 03/17] refactor: clarify resource shortage handling --- .../model_infer/mode_backend/base_backend.py | 20 +++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 8da4efbf07..da428349d3 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -694,12 +694,10 @@ def _get_classed_reqs( prefill_reqs = [] decode_reqs = [] - # 一次性最多暂停请求的数量, 防止盲目暂停大量请求 - # 因为部分请求释放占用的token容量后,就会使推理可以正常进行。 - # 如果因为一次推理容量不足,就以当前token容量的判断暂停了大量 - # 请求,其逻辑是不适合的。 - pause_max_req_num = 2 - wait_pause_count = 0 + # 单轮最多处理少量因 token 容量不足而无法继续的请求。 + # 普通 Decode 会暂停请求,PD Decode 会缩短当前分段;避免单轮过度处理。 + resource_shortage_max_req_num = 2 + resource_shortage_count = 0 prefill_tokens = 0 can_alloc_token_num = g_infer_context.get_can_alloc_token_num() @@ -740,8 +738,10 @@ def _get_classed_reqs( decode_reqs.append(req_obj) can_alloc_token_num -= token_num else: - if wait_pause_count < pause_max_req_num and self._handle_decode_alloc_failure(req_obj): - wait_pause_count += 1 + if resource_shortage_count < resource_shortage_max_req_num: + handled = self._handle_decode_alloc_failure(req_obj) + if handled: + resource_shortage_count += 1 else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 # 的推理代码中将master请求的状态复制到slave请求中去, 所以这里 slave 状态的请求,不 @@ -757,9 +757,9 @@ def _get_classed_reqs( prefill_reqs.append(req_obj) can_alloc_token_num -= token_num else: - if wait_pause_count < pause_max_req_num: + if resource_shortage_count < resource_shortage_max_req_num: req_obj.wait_pause = True - wait_pause_count += 1 + resource_shortage_count += 1 # 先由控制器确定请求需要写入的缓存层级,再按是否包含 CPU cache 决定是否发起 offload。 cache_controller = g_infer_context.cache_placement_controller From 474edb0f21eb4aa2e3f0fb4fb2e80189073defdb Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:15:19 +0800 Subject: [PATCH 04/17] refactor: keep existing pause counter names --- .../router/model_infer/mode_backend/base_backend.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index da428349d3..98387700f9 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -696,8 +696,8 @@ def _get_classed_reqs( # 单轮最多处理少量因 token 容量不足而无法继续的请求。 # 普通 Decode 会暂停请求,PD Decode 会缩短当前分段;避免单轮过度处理。 - resource_shortage_max_req_num = 2 - resource_shortage_count = 0 + pause_max_req_num = 2 + wait_pause_count = 0 prefill_tokens = 0 can_alloc_token_num = g_infer_context.get_can_alloc_token_num() @@ -738,10 +738,10 @@ def _get_classed_reqs( decode_reqs.append(req_obj) can_alloc_token_num -= token_num else: - if resource_shortage_count < resource_shortage_max_req_num: + if wait_pause_count < pause_max_req_num: handled = self._handle_decode_alloc_failure(req_obj) if handled: - resource_shortage_count += 1 + wait_pause_count += 1 else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 # 的推理代码中将master请求的状态复制到slave请求中去, 所以这里 slave 状态的请求,不 @@ -757,9 +757,9 @@ def _get_classed_reqs( prefill_reqs.append(req_obj) can_alloc_token_num -= token_num else: - if resource_shortage_count < resource_shortage_max_req_num: + if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True - resource_shortage_count += 1 + wait_pause_count += 1 # 先由控制器确定请求需要写入的缓存层级,再按是否包含 CPU cache 决定是否发起 offload。 cache_controller = g_infer_context.cache_placement_controller From 06d40a6cacec26a39a0eb2947c9db1a4a17ea981 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:16:14 +0800 Subject: [PATCH 05/17] refactor: return decode shortage count --- .../router/model_infer/mode_backend/base_backend.py | 12 +++++------- .../mode_backend/pd/decode_node_impl/decode_impl.py | 6 +++--- .../pd/decode_node_impl/decode_impl_for_dp.py | 2 +- .../mode_backend/test_pd_dynamic_split.py | 10 +++++----- 4 files changed, 14 insertions(+), 16 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 98387700f9..13c8ed145f 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -739,9 +739,7 @@ def _get_classed_reqs( can_alloc_token_num -= token_num else: if wait_pause_count < pause_max_req_num: - handled = self._handle_decode_alloc_failure(req_obj) - if handled: - wait_pause_count += 1 + wait_pause_count += self._handle_decode_alloc_failure(req_obj) else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 # 的推理代码中将master请求的状态复制到slave请求中去, 所以这里 slave 状态的请求,不 @@ -805,14 +803,14 @@ def _get_classed_reqs( return prefill_reqs, decode_reqs - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: """Handle a decode request that cannot allocate its next token. - Returns whether this request consumed one slot of the per-iteration - shortage handling limit. + Returns the number of slots consumed from the per-iteration shortage + handling limit. """ req_obj.wait_pause = True - return True + return 1 # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 8412ec6e2e..ccecf8b09e 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -119,12 +119,12 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: ans_list.append(req_obj) return ans_list - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: """End an actually running request on its last overlapped output token.""" if not self.support_overlap: return super()._handle_decode_alloc_failure(req_obj) if req_obj.cur_output_len <= req_obj.shm_req.shm_cur_output_len: - return False + return 0 sampling_params = req_obj.sampling_param.shm_param sampling_params.max_new_tokens = min(sampling_params.max_new_tokens, req_obj.cur_output_len) @@ -132,7 +132,7 @@ def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: f"yield running pd decode req_id={req_obj.req_id} " f"at output_len={req_obj.cur_output_len} because token memory is insufficient" ) - return True + return 1 def _decode_node_gen_trans_tasks(self, req_obj: InferReq): """ diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py index 08dde3de90..769b2e1ef7 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py @@ -27,7 +27,7 @@ def _post_init_reqs(self, uninit_reqs: List[InferReq]): def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: return PDDecodeNode._filter_not_ready_reqs(self, req_ids=req_ids) - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> bool: + def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: return PDDecodeNode._handle_decode_alloc_failure(self, req_obj=req_obj) def _decode_node_gen_trans_tasks(self, req_obj: InferReq): diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 09ae766fc3..a8ebd09db5 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -31,11 +31,11 @@ def test_pd_decode_capacity_yield_only_limits_running_overlap_output(): running_req = _make_infer_req(cur_output_len=5, shm_output_len=4) not_running_req = _make_infer_req(cur_output_len=4, shm_output_len=4) - assert backend._handle_decode_alloc_failure(running_req) + assert backend._handle_decode_alloc_failure(running_req) == 1 assert running_req.sampling_param.shm_param.max_new_tokens == 5 assert not running_req.wait_pause - assert not backend._handle_decode_alloc_failure(not_running_req) + assert backend._handle_decode_alloc_failure(not_running_req) == 0 assert not_running_req.sampling_param.shm_param.max_new_tokens == 65535 assert not not_running_req.wait_pause @@ -46,7 +46,7 @@ def test_pd_decode_capacity_yield_preserves_pause_without_overlap(): req = _make_infer_req(cur_output_len=5, shm_output_len=5) req.wait_pause = False - assert backend._handle_decode_alloc_failure(req) + assert backend._handle_decode_alloc_failure(req) == 1 assert req.wait_pause @@ -61,7 +61,7 @@ def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): req.finish_status = FinishStatus() req._stop_sequences_matched = MagicMock(return_value=False) - assert backend._handle_decode_alloc_failure(req) + assert backend._handle_decode_alloc_failure(req) == 1 assert req.sampling_param.shm_param.max_new_tokens == 3 pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=2) @@ -84,7 +84,7 @@ def test_pd_decode_capacity_limit_never_extends_original_length(): finish_status=FinishStatus(), ) - assert backend._handle_decode_alloc_failure(req) + assert backend._handle_decode_alloc_failure(req) == 1 assert req.sampling_param.shm_param.max_new_tokens == 3 From d38917423347ceef1dfc1f9e33017aa1d839d40c Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:20:01 +0800 Subject: [PATCH 06/17] refactor: assume overlap in PD decode yielding --- .../mode_backend/pd/decode_node_impl/decode_impl.py | 2 -- .../model_infer/mode_backend/test_pd_dynamic_split.py | 10 ---------- 2 files changed, 12 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index ccecf8b09e..489485b135 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -121,8 +121,6 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: """End an actually running request on its last overlapped output token.""" - if not self.support_overlap: - return super()._handle_decode_alloc_failure(req_obj) if req_obj.cur_output_len <= req_obj.shm_req.shm_cur_output_len: return 0 diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index a8ebd09db5..ea6fb1f90f 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -40,16 +40,6 @@ def test_pd_decode_capacity_yield_only_limits_running_overlap_output(): assert not not_running_req.wait_pause -def test_pd_decode_capacity_yield_preserves_pause_without_overlap(): - backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) - backend.support_overlap = False - req = _make_infer_req(cur_output_len=5, shm_output_len=5) - req.wait_pause = False - - assert backend._handle_decode_alloc_failure(req) == 1 - assert req.wait_pause - - def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) backend.support_overlap = True From 4c12d0fd7add3cfd1cf3c6d066e3ca6c949d55cc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 12:27:47 +0000 Subject: [PATCH 07/17] fix(pd-master): count segment output tokens locally --- .../httpserver_for_pd_master/manager.py | 12 ++--- .../test_pd_master_multi_choice.py | 44 +++++++++++++++++++ 2 files changed, 50 insertions(+), 6 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index d4b4f24cfe..565f43fa61 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -272,9 +272,9 @@ async def _generate_one( segment_output_tokens = 0 segment_finish_status = FinishStatus() async for sub_req_id, request_output, metadata, raw_finish_status in results_generator: - # pd 分离模式下,返回的 metadata 可能序号信息可能存在不准确性。 + # PD 分离模式下 metadata 中的 token 序号可能不准确,按实际产出计数。 assert sub_req_id == block_group_request_id - segment_output_tokens = metadata["count_output_tokens"] + segment_output_tokens += 1 segment_finish_status = raw_finish_status finish_status = raw_finish_status @@ -302,10 +302,10 @@ async def _generate_one( yield origin_request_id, request_output, metadata, finish_status await self.remove_req(group_request_id=block_group_request_id) - if segment_finish_status.is_finished_length(): - remaining_max_new_tokens -= segment_output_tokens - segment_index += 1 - else: + remaining_max_new_tokens -= segment_output_tokens + segment_index += 1 + # 非 length 状态表示请求已正常结束,无需继续生成下一段。 + if not segment_finish_status.is_finished_length(): break except (ClientDisconnected, BaseException) as e: 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 5d04494390..f2d3b3483c 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -355,6 +355,50 @@ async def wait_to_token_package( asyncio.run(asyncio.wait_for(run(), timeout=2)) +def test_pd_master_counts_segment_tokens_without_relying_on_metadata(): + async def run(): + manager = _manager() + manager.id_gen.generate_id.side_effect = [808, 816] + manager.remove_req = AsyncMock() + manager.abort = AsyncMock() + 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, 0.0)) + dispatched_max_new_tokens = [] + + async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling_params, *_args): + dispatched_max_new_tokens.append(sampling_params.max_new_tokens) + token_count = 2 if len(dispatched_max_new_tokens) == 1 else 1 + for token_index in range(token_count): + finish_status = ( + FinishStatus(FinishStatus.FINISHED_LENGTH) if token_index == token_count - 1 else FinishStatus() + ) + yield ( + sampling_params.group_request_id, + "x", + {"prompt_tokens": 1, "count_output_tokens": 100}, + finish_status, + ) + + manager._wait_to_token_package = wait_to_token_package + + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 3 + async for _ in manager._generate_one( + "prompt", + sampling_params, + MagicMock(), + MagicMock(), + 0, + 800, + ): + pass + + assert dispatched_max_new_tokens == [3, 1] + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + @pytest.mark.parametrize( ("estimated_cache_hit_rate", "expected_high_priority"), [(0.8, False), (0.81, True)], From 71d3b64c5b687bdc67904863dc2b617b20048efc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 12:32:33 +0000 Subject: [PATCH 08/17] feat(pd): enable advanced scheduling on decode nodes --- lightllm/server/api_start.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 37fe837ad1..6f8973425f 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -84,10 +84,10 @@ def _launch_subprocesses(args: StartArgs): # 调度参数的自动设置, 人工设置则听人工的 if args.router_token_ratio is None: - if args.run_mode in ["normal"]: + if args.run_mode in ["normal", "decode"]: args.router_token_ratio = 0.85 else: - # pd 分离模式下,不开启高级调度 + # PD 分离模式下,prefill 节点不开启高级调度 args.router_token_ratio = 0.0 # 部分模式还不能支持与高级动态调度算法协同,to do. if args.diverse_mode: From e0f8fcbb71973aaf4a64859a81cb91407d9ca146 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 12:50:26 +0000 Subject: [PATCH 09/17] refactor(pd): use EMA scheduling on decode nodes --- lightllm/server/router/req_queue/__init__.py | 8 +-- .../req_queue/chunked_prefill/impl_for_pd.py | 15 ++--- .../mode_backend/test_pd_dynamic_split.py | 61 ++++++++----------- 3 files changed, 31 insertions(+), 53 deletions(-) diff --git a/lightllm/server/router/req_queue/__init__.py b/lightllm/server/router/req_queue/__init__.py index 5ca60eeeb0..76157aa6cb 100644 --- a/lightllm/server/router/req_queue/__init__.py +++ b/lightllm/server/router/req_queue/__init__.py @@ -1,15 +1,11 @@ from .chunked_prefill.impl import ChunkedPrefillQueue from .chunked_prefill.beam_impl import ChunkedBeamContinuesBatchQueue -from .chunked_prefill.impl_for_pd import PDDQueue, PDPQueue +from .chunked_prefill.impl_for_pd import PDPQueue from .dp_base_queue import DpQueue def _get_req_queue_class(args, router, dp_size_in_node: int): - # model_rpc 会优先为 PD 节点选择 PD backend,queue 也必须保持 - # 同样的优先级,避免其他模式开关绕过 Decode 的激进调度。 - if args.run_mode == "decode": - return PDDQueue - if args.run_mode == "prefill": + if args.run_mode in ["prefill", "decode"]: return PDPQueue if args.diverse_mode: diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index bb0a4b90ce..90bc26cdd1 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py @@ -3,7 +3,6 @@ from typing import Tuple from ...batch import Batch, Req from lightllm.server.router.req_queue.base_queue import BaseQueue -from lightllm.server.router.req_queue.chunked_prefill.impl import ChunkedPrefillQueue class PDPQueue(BaseQueue): @@ -39,7 +38,11 @@ def _caclu_batch_estimated_peak_token_num(self, batch: Batch): req.get_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) ) else: - estimated_peak_token_num += req.input_len + req.sample_params.max_new_tokens + if self.args.run_mode == "decode": + estimated_output_len = self.router.router_statics.ema_req_out_len + else: + estimated_output_len = req.sample_params.max_new_tokens + estimated_peak_token_num += req.input_len + estimated_output_len if decoding_req_list: decoding_req_list.sort(key=lambda x: -x[1]) @@ -94,11 +97,3 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): estimated_peak_token_num = self._caclu_batch_estimated_peak_token_num(current_batch) return (estimated_peak_token_num, estimated_peak_token_num / self.max_total_tokens) - - -class PDDQueue(ChunkedPrefillQueue): - def is_busy(self): - return False - - def _can_add_first_router_tokens(self, first_router_need_tokens: int) -> bool: - return True diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index ea6fb1f90f..bf3b790737 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -6,10 +6,7 @@ decode_impl as pd_decode_impl, ) from lightllm.server.router.req_queue import _get_req_queue_class -from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd import ( - PDDQueue, - PDPQueue, -) +from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd import PDPQueue def _make_infer_req(cur_output_len: int, shm_output_len: int): @@ -78,36 +75,7 @@ def test_pd_decode_capacity_limit_never_extends_original_length(): assert req.sampling_param.shm_param.max_new_tokens == 3 -def test_pd_queue_admission_always_uses_aggressive_tuple_estimate(): - queue = PDDQueue.__new__(PDDQueue) - queue.dp_index = 0 - queue.max_total_tokens = 100 - queue.running_max_req_size = 8 - queue.batch_max_tokens = 1 - queue.cache_len_list = [(20, 10)] - queue.router = SimpleNamespace( - router_statics=SimpleNamespace(ema_req_out_len=128), - shared_token_load=SimpleNamespace( - set_estimated_peak_token_count=MagicMock(), - set_dynamic_max_load=MagicMock(), - ), - ) - - req = SimpleNamespace( - get_tuple_tokens=MagicMock(return_value=(20, 10)), - get_first_router_need_tokens=MagicMock(return_value=1024), - ) - admitted, first_router_tokens = queue._can_add_new_req(req, queue.is_busy(), 0) - - assert admitted - assert first_router_tokens == 1024 - req.get_tuple_tokens.assert_called_once_with(False, 128) - queue.router.shared_token_load.set_estimated_peak_token_count.assert_called_once_with( - 60, 0 - ) - - -def test_aggressive_pd_queue_is_only_selected_for_decode_nodes(): +def test_pd_nodes_use_pd_queue(): base_args = { "diverse_mode": False, "token_healing_mode": False, @@ -119,7 +87,26 @@ def test_aggressive_pd_queue_is_only_selected_for_decode_nodes(): prefill_args = SimpleNamespace(**base_args, run_mode="prefill") decode_args = SimpleNamespace(**base_args, run_mode="decode") - assert ( - _get_req_queue_class(prefill_args, router=None, dp_size_in_node=1) is PDPQueue + assert _get_req_queue_class(prefill_args, router=None, dp_size_in_node=1) is PDPQueue + assert _get_req_queue_class(decode_args, router=None, dp_size_in_node=1) is PDPQueue + + +def test_pd_decode_queue_uses_ema_for_prefill_stage_output_length(): + queue = PDPQueue.__new__(PDPQueue) + queue.args = SimpleNamespace(run_mode="decode") + queue.dp_index = 0 + queue.max_total_tokens = 4096 + queue.router = SimpleNamespace(router_statics=SimpleNamespace(ema_req_out_len=128)) + queue.is_busy = MagicMock(return_value=False) + + req = SimpleNamespace( + input_len=10, + sample_params=SimpleNamespace(suggested_dp_index=0, max_new_tokens=1024), + is_infer_decode=MagicMock(return_value=False), ) - assert _get_req_queue_class(decode_args, router=None, dp_size_in_node=1) is PDDQueue + batch = SimpleNamespace(reqs=[req]) + + assert queue._caclu_batch_estimated_peak_token_num(batch) == 138 + + queue.args.run_mode = "prefill" + assert queue._caclu_batch_estimated_peak_token_num(batch) == 1034 From 99c6379612106198651a952e116701ed026113f7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 12:58:44 +0000 Subject: [PATCH 10/17] refactor(pd): inline decode allocation fallback --- .../model_infer/mode_backend/base_backend.py | 25 ++++--- .../pd/decode_node_impl/decode_impl.py | 13 ---- .../pd/decode_node_impl/decode_impl_for_dp.py | 3 - .../router/req_queue/chunked_prefill/impl.py | 11 ++-- .../mode_backend/test_pd_dynamic_split.py | 65 +++++++++++++------ 5 files changed, 63 insertions(+), 54 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 13c8ed145f..5297e453f2 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -739,7 +739,21 @@ def _get_classed_reqs( can_alloc_token_num -= token_num else: if wait_pause_count < pause_max_req_num: - wait_pause_count += self._handle_decode_alloc_failure(req_obj) + if self.args.run_mode == "decode": + # overlap 已产生新 token 时,收紧当前分段的输出上限,让请求尽快结束。 + if req_obj.cur_output_len > req_obj.shm_req.shm_cur_output_len: + sampling_params = req_obj.sampling_param.shm_param + sampling_params.max_new_tokens = min( + sampling_params.max_new_tokens, req_obj.cur_output_len + ) + wait_pause_count += 1 + self.logger.info( + f"yield running pd decode req_id={req_obj.req_id} " + f"at output_len={req_obj.cur_output_len} because token memory is insufficient" + ) + else: + req_obj.wait_pause = True + wait_pause_count += 1 else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 # 的推理代码中将master请求的状态复制到slave请求中去, 所以这里 slave 状态的请求,不 @@ -803,15 +817,6 @@ def _get_classed_reqs( return prefill_reqs, decode_reqs - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: - """Handle a decode request that cannot allocate its next token. - - Returns the number of slots consumed from the per-iteration shortage - handling limit. - """ - req_obj.wait_pause = True - return 1 - # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 489485b135..472049442f 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -119,19 +119,6 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: ans_list.append(req_obj) return ans_list - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: - """End an actually running request on its last overlapped output token.""" - if req_obj.cur_output_len <= req_obj.shm_req.shm_cur_output_len: - return 0 - - sampling_params = req_obj.sampling_param.shm_param - sampling_params.max_new_tokens = min(sampling_params.max_new_tokens, req_obj.cur_output_len) - logger.info( - f"yield running pd decode req_id={req_obj.req_id} " - f"at output_len={req_obj.cur_output_len} because token memory is insufficient" - ) - return 1 - def _decode_node_gen_trans_tasks(self, req_obj: InferReq): """ decode node 生成所有的传输任务对象。 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py index 769b2e1ef7..87af300003 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl_for_dp.py @@ -27,9 +27,6 @@ def _post_init_reqs(self, uninit_reqs: List[InferReq]): def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: return PDDecodeNode._filter_not_ready_reqs(self, req_ids=req_ids) - def _handle_decode_alloc_failure(self, req_obj: InferReq) -> int: - return PDDecodeNode._handle_decode_alloc_failure(self, req_obj=req_obj) - def _decode_node_gen_trans_tasks(self, req_obj: InferReq): return PDDecodeNode._decode_node_gen_trans_tasks(self, req_obj=req_obj) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py index 160c086b83..f8fe510989 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py @@ -20,12 +20,6 @@ def _init_cache_list(self, current_batch: Batch, is_busy): self.cache_len_list = [] return - def _can_add_first_router_tokens(self, first_router_need_tokens: int) -> bool: - return ( - self.args.short_prefill_token_threshold is not None - or first_router_need_tokens <= self.batch_max_tokens - ) - # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens): self.cache_len_list.append( @@ -45,7 +39,10 @@ def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens new_batch_first_router_need_tokens += req.get_first_router_need_tokens() # 长短请求模式由 Infer 控制单轮 prefill token 上限。 - ok_prefill = self._can_add_first_router_tokens(new_batch_first_router_need_tokens) + ok_prefill = ( + self.args.short_prefill_token_threshold is not None + or new_batch_first_router_need_tokens <= self.batch_max_tokens + ) if ok_token_num and ok_req_num and ok_prefill: self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index bf3b790737..b67344f182 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -2,6 +2,7 @@ from unittest.mock import MagicMock from lightllm.server.core.objs import FinishStatus +from lightllm.server.router.model_infer.mode_backend import base_backend from lightllm.server.router.model_infer.mode_backend.pd.decode_node_impl import ( decode_impl as pd_decode_impl, ) @@ -12,34 +13,65 @@ def _make_infer_req(cur_output_len: int, shm_output_len: int): return SimpleNamespace( req_id=1, + filter_mark=False, + wait_pause=False, + paused=False, + infer_aborted=False, + finish_status=FinishStatus(), + cur_kv_len=10, cur_output_len=cur_output_len, shm_req=SimpleNamespace(shm_cur_output_len=shm_output_len), sampling_param=SimpleNamespace( shm_param=SimpleNamespace(max_new_tokens=65535), ), - wait_pause=False, + get_cur_total_len=MagicMock(return_value=11), + decode_need_token_num=MagicMock(return_value=1), ) -def test_pd_decode_capacity_yield_only_limits_running_overlap_output(): +def _classify_without_token_capacity(monkeypatch, req): backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) + backend.args = SimpleNamespace( + enable_cpu_cache=False, + enable_prefill_decode_mixed=False, + run_mode="decode", + ) backend.support_overlap = True + backend.is_master_in_dp = True + backend.logger = MagicMock() + backend._timer_merge_radix_tree = MagicMock() + backend._filter_not_ready_reqs = MagicMock(return_value=[req]) + backend._reorder_pd_high_priority_reqs = MagicMock(side_effect=lambda reqs: reqs) + backend._reorder_long_prefill_reqs = MagicMock(side_effect=lambda reqs: reqs) + + infer_context = base_backend.g_infer_context + monkeypatch.setattr(infer_context, "get_can_alloc_token_num", MagicMock(return_value=0)) + monkeypatch.setattr( + infer_context, + "cache_placement_controller", + SimpleNamespace(set_req_cache_way=MagicMock()), + ) + monkeypatch.setattr(infer_context, "filter_reqs", MagicMock()) + monkeypatch.setattr(infer_context, "pause_reqs", MagicMock()) + + backend._get_classed_reqs(req_ids=[req.req_id]) + + +def test_pd_decode_capacity_yield_only_limits_running_overlap_output(monkeypatch): running_req = _make_infer_req(cur_output_len=5, shm_output_len=4) not_running_req = _make_infer_req(cur_output_len=4, shm_output_len=4) - assert backend._handle_decode_alloc_failure(running_req) == 1 + _classify_without_token_capacity(monkeypatch, running_req) assert running_req.sampling_param.shm_param.max_new_tokens == 5 assert not running_req.wait_pause - assert backend._handle_decode_alloc_failure(not_running_req) == 0 + _classify_without_token_capacity(monkeypatch, not_running_req) assert not_running_req.sampling_param.shm_param.max_new_tokens == 65535 assert not not_running_req.wait_pause -def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): - backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) - backend.support_overlap = True +def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(monkeypatch): req = _make_infer_req(cur_output_len=3, shm_output_len=1) req.stop_sequences = [] req.shm_req.input_len = 10 @@ -48,7 +80,7 @@ def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): req.finish_status = FinishStatus() req._stop_sequences_matched = MagicMock(return_value=False) - assert backend._handle_decode_alloc_failure(req) == 1 + _classify_without_token_capacity(monkeypatch, req) assert req.sampling_param.shm_param.max_new_tokens == 3 pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=2) @@ -58,20 +90,11 @@ def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(): assert req.finish_status.is_finished_length() -def test_pd_decode_capacity_limit_never_extends_original_length(): - backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) - backend.support_overlap = True - req = SimpleNamespace( - req_id=1, - cur_output_len=5, - sampling_param=SimpleNamespace( - shm_param=SimpleNamespace(max_new_tokens=3), - ), - shm_req=SimpleNamespace(shm_cur_output_len=1), - finish_status=FinishStatus(), - ) +def test_pd_decode_capacity_limit_never_extends_original_length(monkeypatch): + req = _make_infer_req(cur_output_len=5, shm_output_len=1) + req.sampling_param.shm_param.max_new_tokens = 3 - assert backend._handle_decode_alloc_failure(req) == 1 + _classify_without_token_capacity(monkeypatch, req) assert req.sampling_param.shm_param.max_new_tokens == 3 From 8b82046dad06f7b2192d45bc9b77a84ebac23ad7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 13:06:44 +0000 Subject: [PATCH 11/17] refactor: restore PDQueue name --- lightllm/server/router/req_queue/__init__.py | 4 ++-- .../router/req_queue/chunked_prefill/impl_for_pd.py | 2 +- .../model_infer/mode_backend/test_pd_dynamic_split.py | 8 ++++---- 3 files changed, 7 insertions(+), 7 deletions(-) diff --git a/lightllm/server/router/req_queue/__init__.py b/lightllm/server/router/req_queue/__init__.py index 76157aa6cb..35b4ff6a77 100644 --- a/lightllm/server/router/req_queue/__init__.py +++ b/lightllm/server/router/req_queue/__init__.py @@ -1,12 +1,12 @@ from .chunked_prefill.impl import ChunkedPrefillQueue from .chunked_prefill.beam_impl import ChunkedBeamContinuesBatchQueue -from .chunked_prefill.impl_for_pd import PDPQueue +from .chunked_prefill.impl_for_pd import PDQueue from .dp_base_queue import DpQueue def _get_req_queue_class(args, router, dp_size_in_node: int): if args.run_mode in ["prefill", "decode"]: - return PDPQueue + return PDQueue if args.diverse_mode: return ChunkedBeamContinuesBatchQueue diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index 90bc26cdd1..d91c1ee55b 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py @@ -5,7 +5,7 @@ from lightllm.server.router.req_queue.base_queue import BaseQueue -class PDPQueue(BaseQueue): +class PDQueue(BaseQueue): def __init__(self, args, router, dp_index, dp_size_in_node) -> None: super().__init__(args, router, dp_index, dp_size_in_node) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index b67344f182..7984c40f18 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -7,7 +7,7 @@ decode_impl as pd_decode_impl, ) from lightllm.server.router.req_queue import _get_req_queue_class -from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd import PDPQueue +from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd import PDQueue def _make_infer_req(cur_output_len: int, shm_output_len: int): @@ -110,12 +110,12 @@ def test_pd_nodes_use_pd_queue(): prefill_args = SimpleNamespace(**base_args, run_mode="prefill") decode_args = SimpleNamespace(**base_args, run_mode="decode") - assert _get_req_queue_class(prefill_args, router=None, dp_size_in_node=1) is PDPQueue - assert _get_req_queue_class(decode_args, router=None, dp_size_in_node=1) is PDPQueue + assert _get_req_queue_class(prefill_args, router=None, dp_size_in_node=1) is PDQueue + assert _get_req_queue_class(decode_args, router=None, dp_size_in_node=1) is PDQueue def test_pd_decode_queue_uses_ema_for_prefill_stage_output_length(): - queue = PDPQueue.__new__(PDPQueue) + queue = PDQueue.__new__(PDQueue) queue.args = SimpleNamespace(run_mode="decode") queue.dp_index = 0 queue.max_total_tokens = 4096 From 9cc8f4452de05ab7ff9f2e893cbb8a64bc31cf2e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 13:21:11 +0000 Subject: [PATCH 12/17] fix(pd): force early finish on decode capacity shortage --- .../model_infer/mode_backend/base_backend.py | 28 ++++---- .../mode_backend/test_pd_dynamic_split.py | 72 +++++++++++-------- 2 files changed, 58 insertions(+), 42 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 5297e453f2..185dbe4908 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -694,8 +694,8 @@ def _get_classed_reqs( prefill_reqs = [] decode_reqs = [] - # 单轮最多处理少量因 token 容量不足而无法继续的请求。 - # 普通 Decode 会暂停请求,PD Decode 会缩短当前分段;避免单轮过度处理。 + # 单轮最多处理少量因 token 容量不足而无法继续的请求,避免一次性影响大量请求。 + # 普通 Decode 请求进入暂停队列等待恢复;PD Decode 请求则强制提前结束并进入清理流程。 pause_max_req_num = 2 wait_pause_count = 0 prefill_tokens = 0 @@ -740,17 +740,19 @@ def _get_classed_reqs( else: if wait_pause_count < pause_max_req_num: if self.args.run_mode == "decode": - # overlap 已产生新 token 时,收紧当前分段的输出上限,让请求尽快结束。 - if req_obj.cur_output_len > req_obj.shm_req.shm_cur_output_len: - sampling_params = req_obj.sampling_param.shm_param - sampling_params.max_new_tokens = min( - sampling_params.max_new_tokens, req_obj.cur_output_len - ) - wait_pause_count += 1 - self.logger.info( - f"yield running pd decode req_id={req_obj.req_id} " - f"at output_len={req_obj.cur_output_len} because token memory is insufficient" - ) + # PD Decode 节点的 token 容量不足时,强制当前请求提前结束以释放资源。 + # 单轮只处理 pause_max_req_num 个请求,避免所有资源不足的请求同时退出。 + wait_pause_count += 1 + if support_overlap: + # overlap 模式可能仍有异步计算在访问请求,先标记,下一轮再安全清理。 + req_obj.filter_mark = True + else: + # 非 overlap 模式没有在途的异步计算,可以在本轮直接清理。 + finished_reqs.append(req_obj) + self.logger.info( + f"force early finish for PD decode req_id={req_obj.req_id} " + f"because token capacity is insufficient" + ) else: req_obj.wait_pause = True wait_pause_count += 1 diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 7984c40f18..39955d95c6 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -18,6 +18,7 @@ def _make_infer_req(cur_output_len: int, shm_output_len: int): paused=False, infer_aborted=False, finish_status=FinishStatus(), + cpu_cache_task_status=SimpleNamespace(is_not_started=MagicMock(return_value=True)), cur_kv_len=10, cur_output_len=cur_output_len, shm_req=SimpleNamespace(shm_cur_output_len=shm_output_len), @@ -29,18 +30,20 @@ def _make_infer_req(cur_output_len: int, shm_output_len: int): ) -def _classify_without_token_capacity(monkeypatch, req): +def _classify_without_token_capacity(monkeypatch, req, support_overlap=True): + reqs = req if isinstance(req, list) else [req] backend = pd_decode_impl.PDDecodeNode.__new__(pd_decode_impl.PDDecodeNode) backend.args = SimpleNamespace( enable_cpu_cache=False, enable_prefill_decode_mixed=False, run_mode="decode", ) - backend.support_overlap = True + backend.support_overlap = support_overlap backend.is_master_in_dp = True - backend.logger = MagicMock() + logger = MagicMock() + backend.logger = logger backend._timer_merge_radix_tree = MagicMock() - backend._filter_not_ready_reqs = MagicMock(return_value=[req]) + backend._filter_not_ready_reqs = MagicMock(return_value=reqs) backend._reorder_pd_high_priority_reqs = MagicMock(side_effect=lambda reqs: reqs) backend._reorder_long_prefill_reqs = MagicMock(side_effect=lambda reqs: reqs) @@ -51,43 +54,54 @@ def _classify_without_token_capacity(monkeypatch, req): "cache_placement_controller", SimpleNamespace(set_req_cache_way=MagicMock()), ) - monkeypatch.setattr(infer_context, "filter_reqs", MagicMock()) + filter_reqs = MagicMock() + monkeypatch.setattr(infer_context, "filter_reqs", filter_reqs) monkeypatch.setattr(infer_context, "pause_reqs", MagicMock()) - backend._get_classed_reqs(req_ids=[req.req_id]) + backend._get_classed_reqs(req_ids=[req.req_id for req in reqs]) + return filter_reqs, logger -def test_pd_decode_capacity_yield_only_limits_running_overlap_output(monkeypatch): +def test_pd_decode_capacity_shortage_is_delayed_for_overlap(monkeypatch): + req = _make_infer_req(cur_output_len=5, shm_output_len=4) - running_req = _make_infer_req(cur_output_len=5, shm_output_len=4) - not_running_req = _make_infer_req(cur_output_len=4, shm_output_len=4) + filter_reqs, logger = _classify_without_token_capacity(monkeypatch, req, support_overlap=True) - _classify_without_token_capacity(monkeypatch, running_req) - assert running_req.sampling_param.shm_param.max_new_tokens == 5 - assert not running_req.wait_pause + assert req.filter_mark + assert req.sampling_param.shm_param.max_new_tokens == 65535 + assert not req.wait_pause + filter_reqs.assert_called_once_with(finished_reqs=[]) + assert logger.info.call_args.args[0] == ( + "force early finish for PD decode req_id=1 because token capacity is insufficient" + ) - _classify_without_token_capacity(monkeypatch, not_running_req) - assert not_running_req.sampling_param.shm_param.max_new_tokens == 65535 - assert not not_running_req.wait_pause + filter_reqs, _ = _classify_without_token_capacity(monkeypatch, req, support_overlap=True) + filter_reqs.assert_called_once_with(finished_reqs=[req]) -def test_pd_decode_capacity_limit_finishes_at_committed_mtp_tail(monkeypatch): - req = _make_infer_req(cur_output_len=3, shm_output_len=1) - req.stop_sequences = [] - req.shm_req.input_len = 10 - req.shm_req.shm_prompt_ids = SimpleNamespace(arr=[0] * 13) - req.sampling_param.shm_param.ignore_eos = False - req.finish_status = FinishStatus() - req._stop_sequences_matched = MagicMock(return_value=False) +def test_pd_decode_capacity_shortage_is_filtered_without_overlap(monkeypatch): + req = _make_infer_req(cur_output_len=5, shm_output_len=4) - _classify_without_token_capacity(monkeypatch, req) - assert req.sampling_param.shm_param.max_new_tokens == 3 + filter_reqs, logger = _classify_without_token_capacity(monkeypatch, req, support_overlap=False) + + assert not req.filter_mark + assert req.sampling_param.shm_param.max_new_tokens == 65535 + assert not req.wait_pause + filter_reqs.assert_called_once_with(finished_reqs=[req]) + assert logger.info.call_args.args[0] == ( + "force early finish for PD decode req_id=1 because token capacity is insufficient" + ) + + +def test_pd_decode_capacity_shortage_only_handles_two_requests_per_iteration(monkeypatch): + reqs = [_make_infer_req(cur_output_len=5, shm_output_len=4) for _ in range(3)] + for req_id, req in enumerate(reqs, start=1): + req.req_id = req_id - pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=2) - assert not req.finish_status.is_finished() + filter_reqs, _ = _classify_without_token_capacity(monkeypatch, reqs, support_overlap=True) - pd_decode_impl.InferReq.update_finish_status(req, eos_ids=[], output_len=3) - assert req.finish_status.is_finished_length() + assert [req.filter_mark for req in reqs] == [True, True, False] + filter_reqs.assert_called_once_with(finished_reqs=[]) def test_pd_decode_capacity_limit_never_extends_original_length(monkeypatch): From a78e9ed5e12ef29bfb277be9ccc3ce9e6d2deed9 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 13:34:04 +0000 Subject: [PATCH 13/17] fix(pd): use EMA output length for decode admission --- .../server/router/req_queue/chunked_prefill/impl_for_pd.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index d91c1ee55b..4fb58db4ef 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py @@ -11,7 +11,11 @@ def __init__(self, args, router, dp_index, dp_size_in_node) -> None: # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_num: int) -> Tuple[bool, int, int]: - estimated_peak_token_num += req.input_len + req.sample_params.max_new_tokens + if self.args.run_mode == "decode": + estimated_output_len = self.router.router_statics.ema_req_out_len + else: + estimated_output_len = req.sample_params.max_new_tokens + estimated_peak_token_num += req.input_len + estimated_output_len ok_token_num = estimated_peak_token_num < self.max_total_tokens batch_req_num += 1 ok_req_num = batch_req_num <= self.running_max_req_size From 05c34d81a2520df6f843e82cb213ac23b8c5aed9 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 14:01:27 +0000 Subject: [PATCH 14/17] fix(pd): distinguish decode capacity segments --- lightllm/server/core/objs/req.py | 15 +++- .../httpserver_for_pd_master/manager.py | 41 ++++----- .../server/router/model_infer/infer_batch.py | 13 ++- .../model_infer/mode_backend/base_backend.py | 1 + .../test_pd_master_multi_choice.py | 89 ++++++++++++++++++- unit_tests/server/core/objs/test_req.py | 6 ++ .../mode_backend/test_pd_dynamic_split.py | 46 +++++++++- 7 files changed, 183 insertions(+), 28 deletions(-) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 670fd5d252..9729a8205c 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -27,6 +27,8 @@ class FinishStatus(ctypes.Structure): - ``NO_FINISH``: 未结束 - ``FINISHED_STOP``: 正常停止(EOS / stop 序列等),finish_reason=``stop`` - ``FINISHED_LENGTH``: 达到 max_new_tokens 等长度上限,finish_reason=``length`` + - ``FINISHED_PD_DECODE_CAPACITY``: PD Decode 因 token 容量不足提前结束当前分段。 + 该状态仅用于 PD 内部分段续跑,不应直接暴露给 API 用户。 - ``FINISHED_ABORTED``: 客户端/调度主动 abort,finish_reason=``abort`` - ``FINISHED_ERROR``: 服务端内部错误导致无法继续生成,finish_reason=``error``。 典型场景:PD 分离 decode 节点 KV 传输失败。与 abort 区分:非用户取消, @@ -39,6 +41,9 @@ class FinishStatus(ctypes.Structure): NO_FINISH = 0 FINISHED_STOP = 1 FINISHED_LENGTH = 2 + # PD Decode 因 token 容量不足时,用此状态结束当前 segment。PD Master 会吞掉 + # 模拟结束 token,并携带剩余 max_new_tokens 继续下一个 segment。值 5 用于保持旧状态编号兼容。 + FINISHED_PD_DECODE_CAPACITY = 5 FINISHED_ABORTED = 3 # 内部错误结束(如 PD KV 传输失败);见类文档。 FINISHED_ERROR = 4 @@ -47,14 +52,14 @@ def __init__(self, init_state=NO_FINISH): self.status = init_state def set_status(self, new_status): - assert 0 <= new_status <= 4 + assert 0 <= new_status <= self.FINISHED_PD_DECODE_CAPACITY self.status = new_status def get_status(self): return self.status def is_finished(self): - return self.FINISHED_STOP <= self.status <= self.FINISHED_ERROR + return self.FINISHED_STOP <= self.status <= self.FINISHED_PD_DECODE_CAPACITY def is_stopped(self): return self.status == self.FINISHED_STOP @@ -65,6 +70,9 @@ def is_finished_length(self): def is_finished_error(self): return self.status == self.FINISHED_ERROR + def is_finished_pd_decode_capacity(self): + return self.status == self.FINISHED_PD_DECODE_CAPACITY + def is_error_finished(self): return self.status in (self.FINISHED_ABORTED, self.FINISHED_ERROR) @@ -77,6 +85,9 @@ def get_finish_reason(self): return "abort" elif self.status == self.FINISHED_ERROR: return "error" + elif self.status == self.FINISHED_PD_DECODE_CAPACITY: + # 该内部状态正常会被 PD Master 吞掉;泄漏时按 length 降级处理。 + return "length" return None diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 565f43fa61..1df9e9d4bf 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -233,9 +233,8 @@ async def _generate_one( remaining_max_new_tokens = origin_sampling_params.max_new_tokens segment_index = 0 - # Decode 节点容量不足时会将已运行请求以 length 结束。 - # PD master 把这个内部 length 当作动态分段边界,用剩余 token - # 限额在同一组 P/D 节点上继续。 + # Decode 节点容量不足时会用专用状态结束当前分段。 + # PD Master 吞掉该内部分段 marker,并用剩余 token 限额在同一组 P/D 节点上继续。 while remaining_max_new_tokens > 0: sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) block_group_request_id = self.id_gen.generate_id() @@ -269,18 +268,27 @@ async def _generate_one( request, ) prompt_tokens = sys.maxsize # 因为分段的原因 - segment_output_tokens = 0 segment_finish_status = FinishStatus() async for sub_req_id, request_output, metadata, raw_finish_status in results_generator: # PD 分离模式下 metadata 中的 token 序号可能不准确,按实际产出计数。 assert sub_req_id == block_group_request_id - segment_output_tokens += 1 + + # 收到当前分段的任意输出,说明该请求已经完成 P 节点的 prefill 派发阶段。 + # 立即归还 selector 中记录的在途 prompt 字符数和请求数,并通过置空确保每段只更新一次。 + 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 segment_finish_status = raw_finish_status - finish_status = raw_finish_status - if raw_finish_status.is_finished_length() and segment_output_tokens < remaining_max_new_tokens: - # Decode 节点因容量不足产生的内部分段,不暴露给用户。 - finish_status = FinishStatus() + if raw_finish_status.is_finished_pd_decode_capacity(): + # 容量不足状态是 PD 内部分段边界:吞掉模拟结束 token,继续生成剩余 token。 + break + + # 容量 marker 已在上方过滤,能走到这里的每个 token 都立即扣减全局剩余输出额度。 + remaining_max_new_tokens -= 1 history_gen_token_strs.append(request_output) prompt_tokens = min(prompt_tokens, metadata["prompt_tokens"]) metadata["prompt_tokens"] = prompt_tokens @@ -288,24 +296,17 @@ async def _generate_one( 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) - if not finish_status.is_error_finished(): + if not raw_finish_status.is_error_finished(): # 只有收到成功的推理结果后才将 prompt 写入前缀树,避免尚未进入 # 推理或已失败的请求被后续请求误判为可复用 cache。 self.pd_manager.selector.insert_prompt_cache(prompt, p_node) metadata["prompt_cache_len"] = origin_prompt_cache_len or 0 - 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 + yield origin_request_id, request_output, metadata, raw_finish_status await self.remove_req(group_request_id=block_group_request_id) - remaining_max_new_tokens -= segment_output_tokens segment_index += 1 - # 非 length 状态表示请求已正常结束,无需继续生成下一段。 - if not segment_finish_status.is_finished_length(): + # 只有 PD Decode 容量不足产生的内部分段需要续跑;其他状态都结束整个请求。 + if not segment_finish_status.is_finished_pd_decode_capacity(): break except (ClientDisconnected, BaseException) as e: diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 0cddd84616..3585c223e2 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -907,11 +907,20 @@ def mark_shm_aborted_finished(self): ``Req.mark_simulated_finished``(在已有输出末尾追加 EOS)。不回写本地 ``cur_output_len`` / ``finish_status``(本 InferReq 即将释放)。 """ - # 本地已由 stop / eos / length 等正确结束,保留原 shm finish_status。 + # 请求本身已由 stop / eos / length / error 等状态结束时, + # 请求自身的结束原因优先,不能被后续的容量不足标记覆盖。 if self.finish_status.is_finished(): return + + # 仅在请求本身尚未结束时,才将 finished_by_pd_decode_capacity + # 转换为 PD 内部分段状态,补模拟结束 token 并交给 PD Master 续跑。 + if getattr(self, "finished_by_pd_decode_capacity", False): + finish_status = FinishStatus.FINISHED_PD_DECODE_CAPACITY + else: + finish_status = FinishStatus.FINISHED_ABORTED + self.shm_req.mark_simulated_finished( - FinishStatus.FINISHED_ABORTED, + finish_status, output_len=self.cur_output_len, ) return diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 185dbe4908..560c3b6de4 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -743,6 +743,7 @@ def _get_classed_reqs( # PD Decode 节点的 token 容量不足时,强制当前请求提前结束以释放资源。 # 单轮只处理 pause_max_req_num 个请求,避免所有资源不足的请求同时退出。 wait_pause_count += 1 + setattr(req_obj, "finished_by_pd_decode_capacity", True) if support_overlap: # overlap 模式可能仍有异步计算在访问请求,先标记,下一轮再安全清理。 req_obj.filter_mark = True 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 f2d3b3483c..836734ab56 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -181,6 +181,68 @@ async def generate_one(*_args, **_kwargs): asyncio.run(asyncio.wait_for(run(), timeout=2)) +def test_pd_master_hides_capacity_finish_token_and_continues_next_segment(): + async def run(): + manager = _manager() + manager.id_gen.generate_id.side_effect = [808, 816] + manager.remove_req = AsyncMock() + manager.abort = AsyncMock() + 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, 0.0)) + segment_index = 0 + + async def wait_to_token_package(_p_node, _d_node, _start_time, prompt, params, *_args): + nonlocal segment_index + segment_index += 1 + if segment_index == 1: + assert prompt == "prompt" + assert params.max_new_tokens == 4 + yield ( + 808, + "visible", + {"prompt_tokens": 1, "id": 10, "logprob": -0.1, "logprobs": {"visible": -0.1}}, + FinishStatus(), + ) + yield ( + 808, + "simulated-eos", + {"prompt_tokens": 1, "id": 11, "logprob": 0.0, "logprobs": {"eos": 0.0}}, + FinishStatus(FinishStatus.FINISHED_PD_DECODE_CAPACITY), + ) + else: + assert prompt == "promptvisible" + assert params.max_new_tokens == 3 + yield ( + 816, + "continued", + {"prompt_tokens": 2, "id": 12, "logprob": -0.2, "logprobs": {"continued": -0.2}}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) + + manager._wait_to_token_package = wait_to_token_package + sampling_params = SamplingParams() + sampling_params.max_new_tokens = 4 + + results = [] + async for result in manager._generate_one( + "prompt", + sampling_params, + MagicMock(), + MagicMock(), + 0, + 800, + ): + results.append(result) + + assert segment_index == 2 + assert [result[1] for result in results] == ["visible", "continued"] + assert all(result[3].status != FinishStatus.FINISHED_PD_DECODE_CAPACITY for result in results) + assert results[-1][3].status == FinishStatus.FINISHED_STOP + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_multi_choice_failure_closes_other_generators(): async def run(): manager = _manager() @@ -319,8 +381,19 @@ async def wait_to_token_package( sampling_params.group_request_id, "x", {"prompt_tokens": 1, "count_output_tokens": 1}, - FinishStatus(FinishStatus.FINISHED_LENGTH), + ( + FinishStatus() + if len(dispatched_max_new_tokens) == 1 + else FinishStatus(FinishStatus.FINISHED_LENGTH) + ), ) + if len(dispatched_max_new_tokens) == 1: + yield ( + sampling_params.group_request_id, + "simulated-eos", + {"prompt_tokens": 1, "count_output_tokens": 1}, + FinishStatus(FinishStatus.FINISHED_PD_DECODE_CAPACITY), + ) manager._wait_to_token_package = wait_to_token_package @@ -368,10 +441,13 @@ async def run(): async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling_params, *_args): dispatched_max_new_tokens.append(sampling_params.max_new_tokens) - token_count = 2 if len(dispatched_max_new_tokens) == 1 else 1 + is_first_segment = len(dispatched_max_new_tokens) == 1 + token_count = 2 if is_first_segment else 1 for token_index in range(token_count): finish_status = ( - FinishStatus(FinishStatus.FINISHED_LENGTH) if token_index == token_count - 1 else FinishStatus() + FinishStatus(FinishStatus.FINISHED_LENGTH) + if not is_first_segment and token_index == token_count - 1 + else FinishStatus() ) yield ( sampling_params.group_request_id, @@ -379,6 +455,13 @@ async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling {"prompt_tokens": 1, "count_output_tokens": 100}, finish_status, ) + if is_first_segment: + yield ( + sampling_params.group_request_id, + "simulated-eos", + {"prompt_tokens": 1, "count_output_tokens": 100}, + FinishStatus(FinishStatus.FINISHED_PD_DECODE_CAPACITY), + ) manager._wait_to_token_package = wait_to_token_package diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 0259c0c3e3..3ec1e78ef9 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -98,6 +98,12 @@ def test_finish_status(req): assert req.finish_status.is_error_finished() assert req.finish_status.get_finish_reason() == "error" + req.finish_status.set_status(req.finish_status.FINISHED_PD_DECODE_CAPACITY) + assert req.finish_status.is_finished() + assert req.finish_status.is_finished_pd_decode_capacity() + assert not req.finish_status.is_error_finished() + assert req.finish_status.get_finish_reason() == "length" + if __name__ == "__main__": pytest.main() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 39955d95c6..3921cfe62e 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -68,6 +68,8 @@ def test_pd_decode_capacity_shortage_is_delayed_for_overlap(monkeypatch): filter_reqs, logger = _classify_without_token_capacity(monkeypatch, req, support_overlap=True) assert req.filter_mark + assert req.finished_by_pd_decode_capacity + assert not req.finish_status.is_finished() assert req.sampling_param.shm_param.max_new_tokens == 65535 assert not req.wait_pause filter_reqs.assert_called_once_with(finished_reqs=[]) @@ -85,6 +87,8 @@ def test_pd_decode_capacity_shortage_is_filtered_without_overlap(monkeypatch): filter_reqs, logger = _classify_without_token_capacity(monkeypatch, req, support_overlap=False) assert not req.filter_mark + assert req.finished_by_pd_decode_capacity + assert not req.finish_status.is_finished() assert req.sampling_param.shm_param.max_new_tokens == 65535 assert not req.wait_pause filter_reqs.assert_called_once_with(finished_reqs=[req]) @@ -104,6 +108,37 @@ def test_pd_decode_capacity_shortage_only_handles_two_requests_per_iteration(mon filter_reqs.assert_called_once_with(finished_reqs=[]) +def test_pd_decode_capacity_finish_status_is_written_to_shm(): + shm_req = SimpleNamespace(mark_simulated_finished=MagicMock()) + req = SimpleNamespace( + finish_status=FinishStatus(), + finished_by_pd_decode_capacity=True, + shm_req=shm_req, + cur_output_len=3, + ) + + pd_decode_impl.InferReq.mark_shm_aborted_finished(req) + + shm_req.mark_simulated_finished.assert_called_once_with( + FinishStatus.FINISHED_PD_DECODE_CAPACITY, + output_len=3, + ) + + +def test_request_finish_status_takes_priority_over_pd_decode_capacity_marker(): + shm_req = SimpleNamespace(mark_simulated_finished=MagicMock()) + req = SimpleNamespace( + finish_status=FinishStatus(FinishStatus.FINISHED_STOP), + finished_by_pd_decode_capacity=True, + shm_req=shm_req, + cur_output_len=3, + ) + + pd_decode_impl.InferReq.mark_shm_aborted_finished(req) + + shm_req.mark_simulated_finished.assert_not_called() + + def test_pd_decode_capacity_limit_never_extends_original_length(monkeypatch): req = _make_infer_req(cur_output_len=5, shm_output_len=1) req.sampling_param.shm_param.max_new_tokens = 3 @@ -133,7 +168,14 @@ def test_pd_decode_queue_uses_ema_for_prefill_stage_output_length(): queue.args = SimpleNamespace(run_mode="decode") queue.dp_index = 0 queue.max_total_tokens = 4096 - queue.router = SimpleNamespace(router_statics=SimpleNamespace(ema_req_out_len=128)) + queue.running_max_req_size = 8 + queue.router = SimpleNamespace( + router_statics=SimpleNamespace(ema_req_out_len=128), + shared_token_load=SimpleNamespace( + set_estimated_peak_token_count=MagicMock(), + set_dynamic_max_load=MagicMock(), + ), + ) queue.is_busy = MagicMock(return_value=False) req = SimpleNamespace( @@ -144,6 +186,8 @@ def test_pd_decode_queue_uses_ema_for_prefill_stage_output_length(): batch = SimpleNamespace(reqs=[req]) assert queue._caclu_batch_estimated_peak_token_num(batch) == 138 + assert queue._can_add_new_req(req, estimated_peak_token_num=0, batch_req_num=0) == (True, 138, 1) queue.args.run_mode = "prefill" assert queue._caclu_batch_estimated_peak_token_num(batch) == 1034 + assert queue._can_add_new_req(req, estimated_peak_token_num=0, batch_req_num=0) == (True, 1034, 1) From 922df2917f39b52c0c5c908d4da02be12f899ef2 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 14:11:02 +0000 Subject: [PATCH 15/17] refactor(pd): simplify segment finish handling --- lightllm/server/httpserver_for_pd_master/manager.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 1df9e9d4bf..5085116483 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -268,7 +268,7 @@ async def _generate_one( request, ) prompt_tokens = sys.maxsize # 因为分段的原因 - segment_finish_status = FinishStatus() + raw_finish_status = FinishStatus() async for sub_req_id, request_output, metadata, raw_finish_status in results_generator: # PD 分离模式下 metadata 中的 token 序号可能不准确,按实际产出计数。 assert sub_req_id == block_group_request_id @@ -282,7 +282,6 @@ async def _generate_one( p_node.dispatched_req_num = max(0, p_node.dispatched_req_num - 1) pending_prefill_load_chars = None - segment_finish_status = raw_finish_status if raw_finish_status.is_finished_pd_decode_capacity(): # 容量不足状态是 PD 内部分段边界:吞掉模拟结束 token,继续生成剩余 token。 break @@ -306,7 +305,7 @@ async def _generate_one( await self.remove_req(group_request_id=block_group_request_id) segment_index += 1 # 只有 PD Decode 容量不足产生的内部分段需要续跑;其他状态都结束整个请求。 - if not segment_finish_status.is_finished_pd_decode_capacity(): + if not raw_finish_status.is_finished_pd_decode_capacity(): break except (ClientDisconnected, BaseException) as e: From b09c681d27156fb0617c0248cd073cfebddc9a52 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 14:30:50 +0000 Subject: [PATCH 16/17] fix(pd): cap decode admission estimate by request limit --- .../router/req_queue/chunked_prefill/impl_for_pd.py | 10 ++++++++-- .../model_infer/mode_backend/test_pd_dynamic_split.py | 7 +++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index 4fb58db4ef..637c0454f0 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py @@ -12,7 +12,10 @@ def __init__(self, args, router, dp_index, dp_size_in_node) -> None: # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_num: int) -> Tuple[bool, int, int]: if self.args.run_mode == "decode": - estimated_output_len = self.router.router_statics.ema_req_out_len + estimated_output_len = min( + self.router.router_statics.ema_req_out_len, + req.sample_params.max_new_tokens, + ) else: estimated_output_len = req.sample_params.max_new_tokens estimated_peak_token_num += req.input_len + estimated_output_len @@ -43,7 +46,10 @@ def _caclu_batch_estimated_peak_token_num(self, batch: Batch): ) else: if self.args.run_mode == "decode": - estimated_output_len = self.router.router_statics.ema_req_out_len + estimated_output_len = min( + self.router.router_statics.ema_req_out_len, + req.sample_params.max_new_tokens, + ) else: estimated_output_len = req.sample_params.max_new_tokens estimated_peak_token_num += req.input_len + estimated_output_len diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 3921cfe62e..72ef24342d 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -188,6 +188,13 @@ def test_pd_decode_queue_uses_ema_for_prefill_stage_output_length(): assert queue._caclu_batch_estimated_peak_token_num(batch) == 138 assert queue._can_add_new_req(req, estimated_peak_token_num=0, batch_req_num=0) == (True, 138, 1) + # 接近上下文上限的请求可能只剩很少输出额度,估算值不能超过请求自身的 max_new_tokens, + # 否则本来能够运行的请求会因为 EMA 偏大而永久滞留在 Decode 等待队列。 + req.sample_params.max_new_tokens = 20 + assert queue._caclu_batch_estimated_peak_token_num(batch) == 30 + assert queue._can_add_new_req(req, estimated_peak_token_num=0, batch_req_num=0) == (True, 30, 1) + + req.sample_params.max_new_tokens = 1024 queue.args.run_mode = "prefill" assert queue._caclu_batch_estimated_peak_token_num(batch) == 1034 assert queue._can_add_new_req(req, estimated_peak_token_num=0, batch_req_num=0) == (True, 1034, 1) From e844deb3064eaaeb278d3476c894c678157efbff Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 3 Sep 2026 01:37:52 +0000 Subject: [PATCH 17/17] fix --- lightllm/server/httpserver_for_pd_master/manager.py | 4 +++- test/test_pd_selector/test_pd_master_multi_choice.py | 12 ++++++------ 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 5085116483..a88615f529 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -232,6 +232,9 @@ async def _generate_one( origin_prompt_cache_len = None remaining_max_new_tokens = origin_sampling_params.max_new_tokens segment_index = 0 + # 后续分段的 prompt 会追加已生成内容;始终保留所有分段中最小的 prompt token 数, + # 对外 usage 才能反映用户的原始输入长度,而不是最后一次续跑的 block_prompt 长度。 + prompt_tokens = sys.maxsize # Decode 节点容量不足时会用专用状态结束当前分段。 # PD Master 吞掉该内部分段 marker,并用剩余 token 限额在同一组 P/D 节点上继续。 @@ -267,7 +270,6 @@ async def _generate_one( multimodal_params, request, ) - prompt_tokens = sys.maxsize # 因为分段的原因 raw_finish_status = FinishStatus() async for sub_req_id, request_output, metadata, raw_finish_status in results_generator: # PD 分离模式下 metadata 中的 token 序号可能不准确,按实际产出计数。 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 836734ab56..6955faa734 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -380,12 +380,11 @@ async def wait_to_token_package( yield ( sampling_params.group_request_id, "x", - {"prompt_tokens": 1, "count_output_tokens": 1}, - ( - FinishStatus() - if len(dispatched_max_new_tokens) == 1 - else FinishStatus(FinishStatus.FINISHED_LENGTH) - ), + { + "prompt_tokens": 1 if len(dispatched_max_new_tokens) == 1 else 2, + "count_output_tokens": 1, + }, + (FinishStatus() if len(dispatched_max_new_tokens) == 1 else FinishStatus(FinishStatus.FINISHED_LENGTH)), ) if len(dispatched_max_new_tokens) == 1: yield ( @@ -422,6 +421,7 @@ async def wait_to_token_package( assert p_node.dispatched_prompt_chars == other_request_load assert p_node.dispatched_req_num == other_request_count assert len(results) == 2 + assert [result[2]["prompt_tokens"] for result in results] == [1, 1] assert not results[0][3].is_finished() assert results[1][3].is_finished_length()