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/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: 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 80c114e116..a88615f529 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,24 @@ async def _generate_one( history_gen_token_strs = [] origin_prompt_cache_len = None - - for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list): + 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 节点上继续。 + 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,35 +270,44 @@ 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: - # pd 分离模式下,返回的 metadata 可能序号信息可能存在不准确性。 + 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 - if finish_status.is_finished_length() and not is_last_block: - finish_status = FinishStatus() # 转换为NoFinished + + # 收到当前分段的任意输出,说明该请求已经完成 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 + + 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 - 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) - 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) - if finish_status.is_finished(): + segment_index += 1 + # 只有 PD Decode 容量不足产生的内部分段需要续跑;其他状态都结束整个请求。 + if not raw_finish_status.is_finished_pd_decode_capacity(): break except (ClientDisconnected, BaseException) as e: @@ -718,14 +727,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/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 c64cf906bf..560c3b6de4 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -694,10 +694,8 @@ def _get_classed_reqs( prefill_reqs = [] decode_reqs = [] - # 一次性最多暂停请求的数量, 防止盲目暂停大量请求 - # 因为部分请求释放占用的token容量后,就会使推理可以正常进行。 - # 如果因为一次推理容量不足,就以当前token容量的判断暂停了大量 - # 请求,其逻辑是不适合的。 + # 单轮最多处理少量因 token 容量不足而无法继续的请求,避免一次性影响大量请求。 + # 普通 Decode 请求进入暂停队列等待恢复;PD Decode 请求则强制提前结束并进入清理流程。 pause_max_req_num = 2 wait_pause_count = 0 prefill_tokens = 0 @@ -741,8 +739,24 @@ def _get_classed_reqs( can_alloc_token_num -= token_num else: if wait_pause_count < pause_max_req_num: - req_obj.wait_pause = True - wait_pause_count += 1 + if self.args.run_mode == "decode": + # 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 + 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 else: # 在 diverse mode 模式下,prefill 只会使用 master 状态的请求,slave 请求依靠后续 # 的推理代码中将master请求的状态复制到slave请求中去, 所以这里 slave 状态的请求,不 diff --git a/lightllm/server/router/req_queue/__init__.py b/lightllm/server/router/req_queue/__init__.py index c0de01db28..35b4ff6a77 100644 --- a/lightllm/server/router/req_queue/__init__.py +++ b/lightllm/server/router/req_queue/__init__.py @@ -5,15 +5,15 @@ def _get_req_queue_class(args, router, dp_size_in_node: int): + if args.run_mode in ["prefill", "decode"]: + return PDQueue + 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_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py index f6e33144ef..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 @@ -11,7 +11,14 @@ 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 = 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 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 @@ -38,7 +45,14 @@ 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 = 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 if decoding_req_list: decoding_req_list.sort(key=lambda x: -x[1]) 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..6955faa734 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, ) @@ -182,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() @@ -250,7 +311,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 +323,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 +345,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 +360,124 @@ 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}, - 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 ( + 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 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 [result[2]["prompt_tokens"] for result in results] == [1, 1] + assert not results[0][3].is_finished() + assert results[1][3].is_finished_length() + + 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) + 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 not is_first_segment and token_index == token_count - 1 + else FinishStatus() + ) + yield ( + sampling_params.group_request_id, + "x", + {"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 + + 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)) @@ -353,7 +492,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 +505,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 +532,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 +548,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 +574,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 +587,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/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/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..72ef24342d --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -0,0 +1,200 @@ +from types import SimpleNamespace +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, +) +from lightllm.server.router.req_queue import _get_req_queue_class +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): + return SimpleNamespace( + req_id=1, + filter_mark=False, + wait_pause=False, + 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), + sampling_param=SimpleNamespace( + shm_param=SimpleNamespace(max_new_tokens=65535), + ), + get_cur_total_len=MagicMock(return_value=11), + decode_need_token_num=MagicMock(return_value=1), + ) + + +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 = support_overlap + backend.is_master_in_dp = True + logger = MagicMock() + backend.logger = logger + backend._timer_merge_radix_tree = MagicMock() + 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) + + 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()), + ) + 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 for req in reqs]) + return filter_reqs, logger + + +def test_pd_decode_capacity_shortage_is_delayed_for_overlap(monkeypatch): + req = _make_infer_req(cur_output_len=5, shm_output_len=4) + + 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=[]) + assert logger.info.call_args.args[0] == ( + "force early finish for PD decode req_id=1 because token capacity is insufficient" + ) + + 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_shortage_is_filtered_without_overlap(monkeypatch): + req = _make_infer_req(cur_output_len=5, shm_output_len=4) + + 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]) + 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 + + filter_reqs, _ = _classify_without_token_capacity(monkeypatch, reqs, support_overlap=True) + + assert [req.filter_mark for req in reqs] == [True, True, False] + 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 + + _classify_without_token_capacity(monkeypatch, req) + assert req.sampling_param.shm_param.max_new_tokens == 3 + + +def test_pd_nodes_use_pd_queue(): + 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 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 = PDQueue.__new__(PDQueue) + queue.args = SimpleNamespace(run_mode="decode") + queue.dp_index = 0 + queue.max_total_tokens = 4096 + 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( + input_len=10, + sample_params=SimpleNamespace(suggested_dp_index=0, max_new_tokens=1024), + is_infer_decode=MagicMock(return_value=False), + ) + 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) + + # 接近上下文上限的请求可能只剩很少输出额度,估算值不能超过请求自身的 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)