diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index fbb63d09f0..26a053f4ab 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -87,6 +87,19 @@ PD 分离模式参数 推理进度健康检查:当仍有在途请求,且整个 PD Master 连续 ``HEALTH_TIMEOUT`` 秒 没有任何请求成功返回 token 时,接口将返回 HTTP 503。 +.. option:: --enable_pd_node_self_request_limit + + 在 Prefill/Decode 节点上启用本地请求限流。PD Master 当前不执行请求准入限流。 + HTTP server 申请本地 ``shm_req`` 对象的超时时间由 + ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` 控制(默认 20 秒);请求进入 Router 后等待 + 进入推理系统的超时时间由 ``LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS`` 控制(默认 20 秒)。 + 超时会导致 ``Server is busy``;其中已进入 Router 但仍未进入推理系统的请求会主动标记为 aborted, + 由 PD Master 转换为 HTTP 429。未开启限流时请求会持续等待资源;PD 高优先级请求 + (分段续跑请求或预计输入 cache 命中率高于 0.8 的请求)由 PD Master 通过 + ``pd_high_priority_request_time_out_seconds`` 下发一个统一的超时时间下限。P/D 节点分别取 + 该值与本地 ``shm_req``、Router 超时的较大值;该字段为 0 时不延长本地超时。PD Master 下发值由 + ``LIGHTLLM_PD_HIGH_PRIORITY_REQUEST_TIMEOUT_SECONDS`` 控制,默认 60 秒。该参数默认关闭。 + .. option:: --config_server_host 配置服务器模式下的主机地址 diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 69edf50a86..20aecb973c 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -89,6 +89,20 @@ PD disaggregation Mode Parameters the endpoints return HTTP 503 if no request on the PD Master successfully returns a token for ``HEALTH_TIMEOUT`` consecutive seconds. +.. option:: --enable_pd_node_self_request_limit + + Enable local request limiting on Prefill/Decode nodes. PD Master does not currently perform request admission + limiting. The local ``shm_req`` allocation timeout is controlled by + ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` (20 seconds by default), while the timeout from Router entry + to inference entry is controlled by ``LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS`` (20 seconds by default). + A timeout reports ``Server is busy``; a request that has entered the Router but not inference is proactively + marked aborted, and PD Master converts this to HTTP 429. Requests continue waiting when local admission is + disabled. For PD high-priority requests (segmented continuation requests or requests whose estimated input + cache hit rate is above 0.8), PD Master supplies a shared timeout floor through + ``pd_high_priority_request_time_out_seconds``. Each P/D node uses the greater of this value and its local + ``shm_req`` or Router timeout; zero does not extend the local timeout. The value supplied by PD Master is controlled by + ``LIGHTLLM_PD_HIGH_PRIORITY_REQUEST_TIMEOUT_SECONDS`` and defaults to 60 seconds. Disabled by default. + .. option:: --config_server_host Host address in configuration server mode diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e8..b27736a62f 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -69,9 +69,12 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: ), ) parser.add_argument( - "--disable_pd_master_decode_capacity_limit", + "--enable_pd_node_self_request_limit", action="store_true", - help="Disable PD master admission control based on the total capacity of registered decode nodes.", + help=( + "Enable local request limiting on Prefill/Decode nodes by enforcing shm_req allocation and Router " + "scheduling wait timeouts. PD Master admission limiting is not currently enabled. Default: disabled." + ), ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 87f54fd9e7..670fd5d252 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -65,6 +65,9 @@ def is_finished_length(self): def is_finished_error(self): return self.status == self.FINISHED_ERROR + def is_error_finished(self): + return self.status in (self.FINISHED_ABORTED, self.FINISHED_ERROR) + def get_finish_reason(self): if self.status == self.FINISHED_STOP: return "stop" @@ -83,6 +86,10 @@ class Req(ctypes.Structure): ("index_in_shm_mem", ctypes.c_int), ("ref_count", ctypes.c_int), # 个人不要操作这个计数 # 个人不要操作这个引用计数 ("recv_time", ctypes.c_double), # 用于记录请求到达服务的时间,主要用于调试 + # Router 收到请求和请求被调度为新 batch 的单调时钟时间戳,用于 HTTP server + # 判断请求是否在 Router 等待过久。时间戳写入共享内存,供不同进程读取。 + ("router_arrival_time", ctypes.c_double), + ("infer_start_time", ctypes.c_double), ("request_id", ctypes.c_int64), # 引用计数 ("group_req_id", ctypes.c_int64), ("input_len", ctypes.c_int), @@ -157,6 +164,8 @@ def init( self.index_in_shm_mem: int = self.index_in_shm_mem self.ref_count: int = self.ref_count self.recv_time: float = time.time() + self.router_arrival_time = 0.0 + self.infer_start_time = 0.0 self.request_id = request_id self.group_req_id = convert_sub_id_to_group_id(request_id) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 8e31c50624..a7b4547d8d 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -294,6 +294,12 @@ class SamplingParams(ctypes.Structure): ("stop_sequences", StopSequenceGroups), ("exponential_decay_length_penalty", ExponentialDecayLengthPenalty), ("group_request_id", ctypes.c_int64), # p d mode used params + # 由 PD Master 为分段续跑或预计 cache 命中率较高的请求设置,表示请求需 + # 以高优先级插入 Router 调度队列。 + ("pd_high_priority_request", ctypes.c_bool), + # PD 高优先级请求在开启本地限流的 P/D 节点上的等待时间下限。节点分别取 + # 该值与本地超时的较大值,用于 shm_req 申请和 Router 等待进入推理系统。 + ("pd_high_priority_request_time_out_seconds", ctypes.c_int), ("suggested_dp_index", ctypes.c_int), # suggest dp index, deepseekv2 dp mode, use to suggest used dp_index # in pd split mode, use to keep the id of pd master ("pd_master_node_id", NodeUUId), @@ -337,6 +343,9 @@ def init(self, tokenizer, **kwargs): self.min_new_tokens = kwargs.get("min_new_tokens", 1) self.input_penalty = kwargs.get("input_penalty", DEFAULT_INPUT_PENALTY) self.group_request_id = kwargs.get("group_request_id", -1) + # 这两个字段是 PD Master 的内部调度信息,不能由外部请求参数开启或修改。 + self.pd_high_priority_request = False + self.pd_high_priority_request_time_out_seconds = 0 self.suggested_dp_index = kwargs.get("suggested_dp_index", -1) self.skip_special_tokens = kwargs.get("skip_special_tokens", SKIP_SPECIAL_TOKENS) @@ -503,6 +512,8 @@ def to_dict(self): "allowed_token_ids": self.allowed_token_ids.to_list(), "invalid_token_ids": self.invalid_token_ids.to_list(), "group_request_id": self.group_request_id, + "pd_high_priority_request": self.pd_high_priority_request, + "pd_high_priority_request_time_out_seconds": self.pd_high_priority_request_time_out_seconds, "skip_special_tokens": self.skip_special_tokens, "add_special_tokens": self.add_special_tokens, "add_spaces_between_special_tokens": self.add_spaces_between_special_tokens, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index a9aef608bd..0d8cc4ac7e 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -22,7 +22,7 @@ class StartArgs: pd_master_ip: str = field(default="0.0.0.0") pd_master_port: int = field(default=1212) pd_master_mode: str = field(default="elastic") - disable_pd_master_decode_capacity_limit: bool = field(default=False) + enable_pd_node_self_request_limit: bool = field(default=False) pd_trans_mode: str = field(default="nccl", metadata={"choices": ["nccl", "nixl"]}) config_server_host: str = field(default=None) config_server_port: int = field(default=None) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 9e6f77e58d..8656c41fc1 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -36,9 +36,13 @@ from .manager_ext import HttpRlManagerHelper from lightllm.utils.statics_utils import MovingAverage from lightllm.utils.config_utils import get_vocab_size -from lightllm.utils.envs_utils import get_unique_server_name +from lightllm.utils.envs_utils import ( + get_pd_node_router_wait_timeout_seconds, + get_pd_node_shm_req_alloc_timeout_seconds, + get_unique_server_name, +) from lightllm.utils.shm_port_args import get_shm_port_args -from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken, ServerBusyError from rpyc.utils.classic import obtain logger = init_logger(__name__) @@ -117,6 +121,13 @@ def __init__( self.pd_mode: NodeRole = NodeRole(self.args.run_mode) assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D] + # HTTP server 只负责在本地 shm_req 或 Router 等待过久时快速返回繁忙,PD Master 负责 QPS 准入限流。 + # 该开关控制 P/D 节点是否启用这两类本地等待超时;多机 TP 从节点不独立拒绝请求。 + self.pd_node_request_limit_enabled: bool = ( + self.args.enable_pd_node_self_request_limit and self.pd_mode.is_P_or_D() and not self.is_multinode_tp_slave + ) + self.pd_node_shm_req_alloc_timeout_seconds = get_pd_node_shm_req_alloc_timeout_seconds() + self.pd_node_router_wait_timeout_seconds = get_pd_node_router_wait_timeout_seconds() self.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() @@ -422,22 +433,17 @@ async def generate( # # 这样会缩小 Prefill 节点自身健康检查的覆盖范围:prompt encode、资源上报及 Decode # 资源等待阶段不再计入本地推理健康状态。资源分配异常应由 PD master 侧的运行请求计数、 - # Decode 节点健康检查和等待资源的超时逻辑负责监控,不能依赖 Prefill 推理计数判断。 + # Decode 节点健康检查和本地 shm_req 等待超时负责监控,不能依赖 Prefill 推理计数判断。 await self._register_running_request() running_request_registered = True - # 申请资源并存储 - alloced_req_indexes = [] - while len(alloced_req_indexes) < sampling_params.n: - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - sleep_time = 0.1 - while alloc_req_index is None: - await asyncio.sleep(sleep_time) - sleep_time *= 1.1 - sleep_time = min(1, sleep_time) - - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - alloced_req_indexes.append(alloc_req_index) + # 申请资源并存储。PD 高优先级请求仍以更短的间隔抢占资源;开启本地限流时, + # 使用 PD Master 下发的较长超时时间,避免资源异常时一直等待。 + alloced_req_indexes = await self._alloc_shm_req_indexes( + sampling_params.n, + pd_high_priority_request=sampling_params.pd_high_priority_request, + pd_high_priority_request_time_out_seconds=sampling_params.pd_high_priority_request_time_out_seconds, + ) req_objs: List[Req] = [] for i, req_index in enumerate(alloced_req_indexes): req_obj = await self.shm_req_manager.async_get_req_obj_by_index(req_index) @@ -543,6 +549,55 @@ def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple return image_tokens, audio_tokens + async def _alloc_shm_req_indexes( + self, + req_num: int, + pd_high_priority_request: bool = False, + pd_high_priority_request_time_out_seconds: int = 0, + ) -> List[int]: + """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。 + + 未开启本地限流时无限等待。开启限流后,普通请求使用节点的 shm_req 申请 + 超时时间;高优先级请求取本地超时与 PD Master 下发值中的较大值。 + """ + alloced_req_indexes = [] + alloc_timeout_seconds = None + if self.pd_node_request_limit_enabled: + alloc_timeout_seconds = self.pd_node_shm_req_alloc_timeout_seconds + if pd_high_priority_request: + alloc_timeout_seconds = max( + alloc_timeout_seconds, + pd_high_priority_request_time_out_seconds, + ) + alloc_deadline = time.monotonic() + alloc_timeout_seconds if alloc_timeout_seconds is not None else None + + try: + while len(alloced_req_indexes) < req_num: + alloc_req_index = await self.shm_req_manager.async_alloc_req_index() + # 保持相同的退避起点,仅通过系数让高优先级请求更快地重新尝试获取 shm_req。 + sleep_time_factor = 0.2 if pd_high_priority_request else 1 + sleep_time = 0.1 + while alloc_req_index is None: + if alloc_deadline is not None and time.monotonic() >= alloc_deadline: + logger.warning( + f"{self.args.run_mode} node shm_req allocation timed out after " + f"{alloc_timeout_seconds} seconds" + ) + raise ServerBusyError( + f"PD {self.args.run_mode} node is busy: unable to allocate a shm_req object " + f"within {alloc_timeout_seconds} seconds" + ) + await asyncio.sleep(sleep_time * sleep_time_factor) + sleep_time = min(1, sleep_time * 1.1) + alloc_req_index = await self.shm_req_manager.async_alloc_req_index() + alloced_req_indexes.append(alloc_req_index) + return alloced_req_indexes + except BaseException: + # 批量申请中途失败时,释放已申请的索引,避免 shm_req 资源泄漏。 + for req_index in alloced_req_indexes: + await self.shm_req_manager.async_release_req_index(req_index) + raise + async def _log_req_header(self, request_headers, group_request_id: int): x_request_id = request_headers.get("X-Request-Id", "") x_session_id = request_headers.get("X-Session-Id", "") @@ -736,6 +791,16 @@ async def _wait_to_token_package( except asyncio.TimeoutError: pass + if ( + self.pd_node_request_limit_enabled + and is_first_token + and req_status.has_timed_out_waiting_for_inference(self.pd_node_router_wait_timeout_seconds) + ): + raise ServerBusyError( + f"PD {self.args.run_mode} node is busy: request did not enter inference " + f"within {self.pd_node_router_wait_timeout_seconds} seconds" + ) + if request is not None and await request.is_disconnected(): await self.abort(group_request_id) raise ClientDisconnected( @@ -1043,6 +1108,26 @@ def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], sta ) self.out_token_info_list = [] + def has_timed_out_waiting_for_inference(self, timeout_seconds: float) -> bool: + """判断请求组是否已在 Router 中等待进入推理系统超时。""" + current_time = time.monotonic() + reqs = self.group_req_objs.shm_req_objs + # 组内任一请求已经进入新 batch,说明整个请求组已经开始执行,不能再按 Router 等待超时清理。 + if any(req.infer_start_time > 0 for req in reqs): + return False + # 高优先级请求取本地 Router 超时与 PD Master 下发值中的较大值,既保证 + # 它比普通请求拥有更充足的等待机会,也避免资源异常时永久滞留。 + if any(req.sample_params.pd_high_priority_request for req in reqs): + timeout_seconds = max( + timeout_seconds, + reqs[0].sample_params.pd_high_priority_request_time_out_seconds, + ) + + for req in reqs: + if req.router_arrival_time > 0 and current_time - req.router_arrival_time >= timeout_seconds: + return True + return False + def can_release(self): for req in self.group_req_objs.shm_req_objs: if not req.can_release(): diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index e66540cb5e..51b7eea33b 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -20,7 +20,7 @@ from ..pd_io_struct import PD_Master_Obj from lightllm.server.core.objs import StartArgs from lightllm.server.core.objs import SamplingParams -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError from lightllm.utils.shm_port_args import get_shm_port_args logger = init_logger(__name__) @@ -247,6 +247,13 @@ async def _pd_process_generate( await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") + except ServerBusyError as e: + group_request_id = sampling_params.group_request_id + logger.warning(f"pd node rejected request {group_request_id}: {e.message}") + try: + await pd_upload_websocket.send(pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, group_request_id, e.message))) + except Exception: + logger.exception(f"report pd node request rejection failed, group_request_id: {group_request_id}") except asyncio.CancelledError: # PD master 主动 abort 或连接断开清理任务时会走取消路径,不需要反向重复上报。 pass diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index b422bf7703..80c114e116 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_split_max_new_tokens +from lightllm.utils.envs_utils import get_pd_high_priority_request_timeout_seconds, get_pd_split_max_new_tokens from lightllm.utils.shm_port_args import get_shm_port_args from .pd_selector import create_selector @@ -48,6 +48,9 @@ def __init__( self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 + # 高优先级请求仍可比普通请求等待更久,但通过请求参数向开启本地限流的 + # P/D 节点传递有限的等待时间,避免资源异常时永久占用请求链路。 + self.pd_high_priority_request_time_out_seconds = get_pd_high_priority_request_timeout_seconds() self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -119,9 +122,13 @@ def tokens(self, prompt, multimodal_params, samping_params: SamplingParams, kwar async def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: return self.pd_manager.select_p_d_node(prompt, sampling_params, multimodal_params) + async def _wait_for_pd_master_request_slot(self) -> None: + """PD Master 请求准入的预留接口,当前不执行限流。""" + return + async def generate( self, prompt: Union[str, List[int]], @@ -129,10 +136,7 @@ async def generate( multimodal_params: MultimodalParams, request: Request, ): - if not self.args.disable_pd_master_decode_capacity_limit: - decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - if self.running_request_count >= decode_capacity: - raise ServerBusyError() + await self._wait_for_pd_master_request_slot() was_idle = self.running_request_count == 0 self.running_request_count += 1 @@ -192,9 +196,14 @@ async def _generate( ) ) + request_finished_successfully = True async for result in self._merge_choice_generators(generators): + finish_status = result[3] + if finish_status.is_error_finished(): + request_finished_successfully = False yield result - self.metric_client.counter_inc("lightllm_request_success") + if request_finished_successfully: + self.metric_client.counter_inc("lightllm_request_success") return async def _generate_one( @@ -219,7 +228,9 @@ async def _generate_one( pending_prefill_load_chars = None try: - p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) + p_node, d_node, estimated_cache_hit_rate = await self.select_p_d_node( + prompt, origin_sampling_params, multimodal_params + ) if not p_node or not d_node: logger.error(f"{origin_request_id}: No p_node or d_node found") raise Exception(f"{origin_request_id}: No p_node or d_node found") @@ -233,6 +244,16 @@ async def _generate_one( 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 + # 预计输入 cache 命中率高于 0.8 时,将首段也标记为 PD 高优先级请求, + # 使其优先进入 Router 调度队列,尽快复用已命中的 KV cache。第二段及后续 + # 分段仍统一使用高优先级,避免因临时资源紧张导致分段续跑失败。 + sampling_params.pd_high_priority_request = iter_index > 0 or estimated_cache_hit_rate > 0.8 + # 为高优先级请求下发较长的有限等待时间;P/D 节点仅在自身开启 + # 本地限流时使用该值,未开启限流时仍保持无限等待。 + if sampling_params.pd_high_priority_request: + sampling_params.pd_high_priority_request_time_out_seconds = ( + self.pd_high_priority_request_time_out_seconds + ) # 分段请求始终复用循环外选定的 P 节点;这里只按每段实际发送的 # prompt 更新该节点的在途 prefill 负载,不会重新选点。 @@ -263,6 +284,10 @@ 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(): + # 只有收到成功的推理结果后才将 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( @@ -677,6 +702,16 @@ async def handle_loop(self): ) else: await req_status.set_error(error_info) + elif obj[0] == ObjType.PD_UPLOAD_SERVER_BUSY: + _, group_req_id, error_info = obj + logger.warning( + f"received PD node server busy, group_req_id: {group_req_id}, reason: {error_info}" + ) + req_status = self.req_id_to_out_inf.get(group_req_id) + if req_status is None: + logger.error(f"PD_UPLOAD_SERVER_BUSY fail find req status for group_req_id: {group_req_id}") + else: + await req_status.set_error(error_info, is_server_busy=True) else: logger.error(f"recevie error obj {obj}") except BaseException as e: @@ -703,6 +738,7 @@ def __init__(self, req_id, p_node, d_node) -> None: self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node self.error_info: Optional[str] = None + self.is_server_busy = False async def wait_to_ready(self): try: @@ -710,9 +746,10 @@ async def wait_to_ready(self): except asyncio.TimeoutError: pass - async def set_error(self, error_info: str): + async def set_error(self, error_info: str, is_server_busy: bool = False): async with self.lock: self.error_info = error_info + self.is_server_busy = is_server_busy # 请求可能正在等待 Prefill prompt ids、Decode KV 资源或输出 token, # 设置全部事件,让请求自己的执行循环立即醒来并抛出异常。 self.event.set() @@ -721,6 +758,8 @@ async def set_error(self, error_info: str): def raise_if_error(self): if self.error_info is not None: + if self.is_server_busy: + raise ServerBusyError(self.error_info) logger.error( f"group_request_id: {self.req_id} detected PD node generate error, " f"raise exception to end the request flow early: {self.error_info}" @@ -885,6 +924,8 @@ def update_node_load_info(self, load_info: Optional[dict]): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: - p_node, d_node = self.selector.select_p_d_node(prompt, sampling_params, multimodal_params) - return p_node, d_node + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: + p_node, d_node, estimated_cache_hit_rate = self.selector.select_p_d_node( + prompt, sampling_params, multimodal_params + ) + return p_node, d_node, estimated_cache_hit_rate diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index adf911bc21..d0e07448a0 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -7,7 +7,7 @@ 负载均衡,避免热点。 实现要点: - - 用前缀树(见 PromptCacheTree)记录「历史 prompt -> 处理它的 worker」; + - 用前缀树(见 PromptCacheTree)记录「成功进入推理的 prompt -> 处理它的 worker」; - 树中的 prefill_node 对应 worker.client_ip_port; - prompt 会按 sample_stride 抽稀后再插入/匹配,降低树的深度与内存; - 根据推理侧返回的平均 prompt cache 命中率,动态调整 cache 亲和与负载均衡的权重; @@ -96,7 +96,7 @@ class CacheAwarePolicy: 维护 prompt 前缀树,并据此为请求选择 prefill worker。 树生命周期: - - 选中 worker 后会把当前 prompt 插入该 worker 对应的 prefill_node; + - 请求成功进入推理后,把当前 prompt 插入实际 worker 对应的 prefill_node; - insert 时若超 max_node_count 会 lazy 触发 LRU 驱逐。 """ @@ -131,7 +131,6 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti 4) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 得到 cache 命中节点; 5) cache 命中节点负载未严重高于最空闲节点 -> 选择 cache 命中节点; 6) 未命中或负载严重失衡 -> 选择最空闲节点; - 7) 将当前 prompt 与最终选中的节点写入前缀树。 """ if not workers: return None @@ -141,15 +140,23 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti # ---- 1. 空闲优先:避免有可用 GPU 闲置 ---- idle_worker = self._select_idle_worker(workers, request_text) if idle_worker is not None: - self.prompt_cache_tree.insert(request_text, idle_worker.client_ip_port) return idle_worker # ---- 2. 所有节点都忙时,在 cache 亲和与负载均衡之间权衡 ---- cache_worker = self._get_cache_worker(workers, request_text) selected_worker = self._select_worker_by_cache_and_load(workers, cache_worker, len(request_text)) + return selected_worker + + def get_estimated_cache_hit_rate(self, selected_worker: PD_Client_Obj, request_text: str) -> float: + """查询最终选中节点的输入 cache 命中率估计。""" + result = self.prompt_cache_tree.prefix_match(request_text) + if result.prefill_node != selected_worker.client_ip_port or result.input_char_count == 0: + return 0.0 + return result.matched_char_count / result.input_char_count + def insert_prompt_cache(self, request_text: str, selected_worker: PD_Client_Obj) -> None: + """在请求成功进入推理后,记录 prompt 与实际执行的 Prefill 节点。""" self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port) - return selected_worker def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: """记录推理侧上报的真实 cache 命中率,并更新动态负载阈值。""" diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index 5474806b7b..ecfdbcc3e4 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py @@ -23,23 +23,27 @@ def update_nodes(self, prefill_nodes, decode_nodes): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: raise NotImplementedError("Subclass must implement this method") def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: """记录推理侧返回的 prompt cache 命中率;非 cache-aware 策略无需处理。""" return + def insert_prompt_cache(self, prompt: str, p_node: PD_Client_Obj) -> None: + """记录成功进入推理的 prompt;非 cache-aware 策略无需处理。""" + return + class RandomSelector(PDSelector): """随机选择器""" def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: p_node = random.choice(self.prefill_nodes) d_node = random.choice(self.decode_nodes) - return p_node, d_node + return p_node, d_node, 0.0 class RoundRobinSelector(PDSelector): @@ -52,14 +56,14 @@ def __init__(self, pd_manager): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: self.prefill_node_index = self.prefill_node_index % len(self.prefill_nodes) self.decode_node_index = self.decode_node_index % len(self.decode_nodes) p_node = self.prefill_nodes[self.prefill_node_index] d_node = self.decode_nodes[self.decode_node_index] self.prefill_node_index += 1 self.decode_node_index += 1 - return p_node, d_node + return p_node, d_node, 0.0 class AdaptiveLoadSelector(PDSelector): @@ -67,11 +71,11 @@ class AdaptiveLoadSelector(PDSelector): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: p_node = self._importance_sampling(self.prefill_nodes) d_node = self._importance_sampling(self.decode_nodes) - return p_node, d_node + return p_node, d_node, 0.0 def _importance_sampling(self, nodes: List[PD_Client_Obj]): return random.choices(nodes, weights=[max(1.0 - e.run_status.total_token_usage_rate, 0.02) for e in nodes])[0] @@ -86,17 +90,23 @@ def __init__(self, pd_manager): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams - ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: + ) -> Tuple[PD_Client_Obj, PD_Client_Obj, float]: assert isinstance(prompt, str), "prompt must be a string for cache-aware selection" p_node = self.policy.select_worker(self.prefill_nodes, request_text=prompt) d_node = self._importance_sampling(self.decode_nodes) + # 选点完成后再查询一次前缀树;只有 cache 实际属于最终选中的 P 节点 + # 时才返回命中率,避免负载均衡改派节点后误判为高命中。 + estimated_cache_hit_rate = self.policy.get_estimated_cache_hit_rate(p_node, prompt) logger.info( f"LoadBalancedCacheAwareSelector: selected p_node={p_node.client_ip_port}, " f"d_node={d_node.client_ip_port}" ) - return p_node, d_node + return p_node, d_node, estimated_cache_hit_rate def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.policy.record_prompt_cache_hit_rate(cache_hit_rate) + + def insert_prompt_cache(self, prompt: str, p_node: PD_Client_Obj) -> None: + self.policy.insert_prompt_cache(prompt, p_node) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 8a1e2bd42b..78f5fedc93 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -42,6 +42,7 @@ class ObjType(enum.Enum): PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 PD_UPLOAD_GENERATE_ERROR = 7 # P/D 节点向 pd master 上报本地请求生成异常。 + PD_UPLOAD_SERVER_BUSY = 8 # P/D 节点向 pd master 上报本地服务繁忙。 @dataclass diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index 11e4e5f688..b1375d754c 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -314,6 +314,11 @@ async def _step(self): async def _add_batch(self, batch: Batch): # 添加新请求 + # 请求被 Router 调度为新 batch 并准备下发到推理系统时记录时间,HTTP server + # 以此判断请求是否在 Router 队列中等待过久;不需要推理进程额外写共享字段。 + infer_start_time = time.monotonic() + for req in batch.reqs: + req.infer_start_time = infer_start_time reqs = [r.to_router_rpc_obj() for r in batch.reqs] while not self.shm_reqs_io_buffer.is_empty(): await asyncio.sleep(0.001) @@ -416,10 +421,12 @@ def get_used_tokens(self, dp_index): def _add_req(self, group_req_indexes: GroupReqIndexes): req_group = [] + router_arrival_time = time.monotonic() for req_index in group_req_indexes.shm_req_indexes: req = self.shm_req_manager.get_req_obj_by_index(req_index) req.multimodal_params = group_req_indexes.multimodal_params req.start_time = group_req_indexes.time_mark + req.router_arrival_time = router_arrival_time # 附加一个私有标记变量,标记请求是否已经被router发送过abort命令给推理进程, # 防止反复发送abort命令给推理进程 req._router_aborted = False 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 07b88471f0..c64cf906bf 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -616,6 +616,14 @@ def _timer_merge_radix_tree(self): ) return + def _reorder_pd_high_priority_reqs(self, ready_reqs: List[InferReq]) -> List[InferReq]: + """将 PD 分段续跑的高优先级请求前置,普通请求保持在其后。""" + # PD 分段续跑请求已经完成前一段推理,需要优先进入本轮调度;将请求拆分后再拼接, + # 保持各自原有顺序,并确保高优先级请求位于普通请求之前。 + high_priority_reqs = [req for req in ready_reqs if req.shm_req.sample_params.pd_high_priority_request] + normal_reqs = [req for req in ready_reqs if not req.shm_req.sample_params.pd_high_priority_request] + return high_priority_reqs + normal_reqs + def _reorder_long_prefill_reqs(self, ready_reqs: List[InferReq]) -> List[InferReq]: """ 提升一个短 prefill 请求的优先级。 @@ -677,6 +685,7 @@ def _get_classed_reqs( ready_reqs = self._filter_not_ready_reqs(req_ids) support_overlap = self.support_overlap + ready_reqs = self._reorder_pd_high_priority_reqs(ready_reqs) ready_reqs = self._reorder_long_prefill_reqs(ready_reqs) wait_pause_reqs = [] 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 f9dc6ee60f..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 @@ -65,8 +65,17 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: for request_id in req_ids: req_obj: InferReq = g_infer_context.requests_mapping[request_id] - if self.is_master_in_dp and req_obj.infer_aborted and req_obj.pd_task_num != 0: + # D 节点的请求未结束时会反复进入该过滤逻辑,最多发送 6 次 + # PDAbortReq,既提高 abort 消息被传输层处理的概率,也避免持续重复发送。 + pd_abort_req_send_count = getattr(req_obj, "pd_abort_req_send_count", 0) + if ( + self.is_master_in_dp + and req_obj.infer_aborted + and req_obj.pd_task_num != 0 + and pd_abort_req_send_count < 6 + ): self.info_queue.put(PDAbortReq(request_id=req_obj.req_id, device_id=req_obj.pd_trans_device_id)) + req_obj.pd_abort_req_send_count = pd_abort_req_send_count + 1 if req_obj.pd_task_num != (req_obj.pd_task_failed_num + req_obj.pd_task_success_num): continue diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 27fe528e9b..b406405e8a 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -271,8 +271,10 @@ def accept_peer_task_loop( local_trans_task = self.waiting_dict.pop(notify_obj.get_key(), None) if local_trans_task is not None: local_trans_task.error_info = notify_obj.error_info - # 软性的调整超时时间,防止一些特殊情况,过快的释放task - # 占用的page 页面,导致多p 复写引起脏内容的问题。 + # TODO: 这里设置 12 秒超时后会立即将任务放入 failed_queue, + # fail_loop 会直接归还任务占用的 page,因此该超时时间实际不会再被检查。 + # 如果底层异步传输尚未结束,page 可能被新任务提前复用,导致脏数据。 + # 后续需要在确认传输已经静默(DONE/ERR)后再回收 page。 local_trans_task.transfer_time_out_secs = 12 self.failed_queue.put(local_trans_task) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 2a501f509b..084e51e1e5 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -2,7 +2,7 @@ import random from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import InferReq -from lightllm.server.pd_io_struct import PDChunckedTransTask +from lightllm.server.pd_io_struct import PDAbortReq, PDChunckedTransTask from lightllm.utils.log_utils import init_logger from lightllm.utils.device_utils import kv_trans_use_p2p from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -32,6 +32,19 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: ans_list: List[InferReq] = [] for request_id in req_ids: req_obj: InferReq = g_infer_context.requests_mapping[request_id] + + # P 节点的推理请求收到 abort 后,主动通知 KV 传输层停止该请求 + # 尚未完成的传输任务,避免只能等待 D 节点上报错误或传输超时。 + pd_abort_req_send_count = getattr(req_obj, "pd_abort_req_send_count", 0) + if ( + self.is_master_in_dp + and req_obj.infer_aborted + and req_obj.pd_task_num != 0 + and pd_abort_req_send_count < 6 + ): + self.info_queue.put(PDAbortReq(request_id=req_obj.req_id, device_id=req_obj.pd_trans_device_id)) + req_obj.pd_abort_req_send_count = pd_abort_req_send_count + 1 + prefill_finished = req_obj.shm_req.input_len <= req_obj.cur_kv_len if prefill_finished: # 等待所有传输任务都已经完成。 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py index 8d55b0417c..b23a5c4141 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py @@ -4,7 +4,7 @@ import time from typing import List, Dict, Union, Callable from lightllm.utils.log_utils import init_logger -from lightllm.server.pd_io_struct import PDChunckedTransTask +from lightllm.server.pd_io_struct import PDAbortReq, PDChunckedTransTask from lightllm.utils.graceful_utils import graceful_registry from lightllm.server.core.objs import StartArgs from ..trans_process_obj import KVTransProcess @@ -58,12 +58,21 @@ def __init__(self, args: StartArgs, info_queue: mp.Queue, start_trans_process_fu def task_dispatcher_loop(self): # 获取任务,并分发给相关卡的处理队列 while True: - task: PDChunckedTransTask = self.info_queue.get() + task: Union[PDChunckedTransTask, PDAbortReq] = self.info_queue.get() + + if isinstance(task, PDChunckedTransTask): + device_id = task.src_device_id + elif isinstance(task, PDAbortReq): + # P 节点的 abort 命令需要与普通传输任务一样,路由到该请求 + # 所在设备的 KV 传输进程中处理。 + device_id = task.device_id + else: + raise TypeError(f"unsupported prefill kv move task type: {type(task)}") - device_id = task.src_device_id try: trans_process: KVTransProcess = self.kv_trans_processes[device_id] trans_process.task_in_queue.put(task) - logger.info(f"kv move manager dispatch task {task.to_str()} to device {device_id}") + if isinstance(task, PDChunckedTransTask): + logger.info(f"kv move manager dispatch task {task.to_str()} to device {device_id}") except BaseException as e: logger.exception(str(e)) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 46181d6fc2..534966ebe9 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -6,10 +6,10 @@ import torch.multiprocessing as mp import queue import pickle -from typing import List, Dict, Optional +from typing import List, Dict, Optional, Union from lightllm.utils.log_utils import init_logger from lightllm.common.kv_cache_mem_manager import MemoryManager -from lightllm.server.pd_io_struct import PDChunckedTransTask +from lightllm.server.pd_io_struct import PDAbortReq, PDChunckedTransTask from lightllm.utils.graceful_utils import graceful_registry from lightllm.server.core.objs import StartArgs from ..kv_transporter import create_kv_transporter @@ -158,7 +158,9 @@ def _abort(self, request_id: int, error_info: str = "aborted req"): aborted_tasks = [] with self.waiting_dict_lock: for key, trans_task in list(self.waiting_dict.items()): - if trans_task.request_id == request_id: + if trans_task.request_id == request_id and trans_task.xfer_handle is None: + # 已经提交给底层 transporter 的异步传输不能直接失败, + # 否则 fail_loop 可能在传输静默前归还 source page,导致脏数据。 aborted_tasks.append(self.waiting_dict.pop(key)) for trans_task in aborted_tasks: @@ -171,8 +173,17 @@ def recv_task_loop(self): torch.cuda.set_device(self.device_id) while True: + obj: Union[PDChunckedTransTask, PDAbortReq] = self.task_in_queue.get() + if isinstance(obj, PDAbortReq): + # Infer 层已经终止该请求,将 abort 命令下传到 P 节点的 + # KV 传输层,尽快结束尚在等待的传输任务。 + self._abort(request_id=obj.request_id) + continue + + assert isinstance(obj, PDChunckedTransTask), f"unsupported prefill transfer task type: {type(obj)}" + page_index = self.page_index_queue.get() - trans_task: PDChunckedTransTask = self.task_in_queue.get() + trans_task = obj trans_task.src_page_index = page_index # 初次校验 time out diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 9af7afd1b4..27e51a4102 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -83,7 +83,22 @@ def filter_aborted_reqs(self): def extend(self, req_group: List[Req]): for req in req_group: req.sample_params.suggested_dp_index = self.dp_index - self.waiting_req_list.extend(req_group) + # PD 高优先级请求应排在普通请求之前,但高优先级请求之间仍按到达顺序排队, + # 避免后到请求反复插到队头而阻塞先到的高优先级请求。 + if req_group and req_group[0].sample_params.pd_high_priority_request: + first_normal_req_index = len(self.waiting_req_list) + for index, waiting_req in enumerate(self.waiting_req_list): + if not waiting_req.sample_params.pd_high_priority_request: + first_normal_req_index = index + break + # req_group 可能包含同一请求组的多个 Req,整体插入可以保持组内顺序。 + self.waiting_req_list = ( + self.waiting_req_list[:first_normal_req_index] + + req_group + + self.waiting_req_list[first_normal_req_index:] + ) + else: + self.waiting_req_list.extend(req_group) return def get_wait_req_num(self): diff --git a/lightllm/server/router/req_queue/dp_base_queue.py b/lightllm/server/router/req_queue/dp_base_queue.py index af8f875d4e..b2d1ff0f7e 100644 --- a/lightllm/server/router/req_queue/dp_base_queue.py +++ b/lightllm/server/router/req_queue/dp_base_queue.py @@ -60,7 +60,16 @@ def extend(self, req_group: List[Req]): suggested_dp_index = req_group[0].sample_params.suggested_dp_index if suggested_dp_index >= self.dp_size_in_node or suggested_dp_index < 0: # 同一个组的,要分配在同一个 dp 上 - self.reqs_waiting_for_dp_index.append(req_group) + if req_group[0].sample_params.pd_high_priority_request: + # 高优先级请求组插在第一个普通请求组之前,同时保持高优先级组之间的 FIFO 顺序。 + first_normal_group_index = len(self.reqs_waiting_for_dp_index) + for index, waiting_group in enumerate(self.reqs_waiting_for_dp_index): + if not waiting_group[0].sample_params.pd_high_priority_request: + first_normal_group_index = index + break + self.reqs_waiting_for_dp_index.insert(first_normal_group_index, req_group) + else: + self.reqs_waiting_for_dp_index.append(req_group) else: self.inner_queues[suggested_dp_index].extend(req_group) return diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 933728b7ff..fb30f47938 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -311,6 +311,24 @@ 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`` 对象的最长等待时间,单位为秒。""" + return int(os.getenv("LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS", 20)) + + +@lru_cache(maxsize=None) +def get_pd_node_router_wait_timeout_seconds() -> int: + """请求进入 Router 后等待进入推理系统的最长时间,单位为秒。""" + return int(os.getenv("LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS", 20)) + + +@lru_cache(maxsize=None) +def get_pd_high_priority_request_timeout_seconds() -> int: + """PD Master 为高优先级请求设置的等待时间下限,单位为秒。""" + return int(os.getenv("LIGHTLLM_PD_HIGH_PRIORITY_REQUEST_TIMEOUT_SECONDS", 60)) + + @lru_cache(maxsize=None) def get_lightllm_url_pool_maxsize() -> int: return int(os.getenv("LIGHTLLM_URL_POOL_MAXSIZE", 512)) 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 7d0b32ded9..98e0014602 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -18,6 +18,7 @@ def _manager() -> HttpServerManagerForPDMaster: manager.metric_client = MagicMock() manager._log_req_header = AsyncMock() manager.tokens = MagicMock(return_value=2) + manager.pd_high_priority_request_time_out_seconds = 60 return manager @@ -39,7 +40,7 @@ async def run(): captured_params = [] p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + manager.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() @@ -143,6 +144,44 @@ async def generate_one( asyncio.run(asyncio.wait_for(run(), timeout=2)) +@pytest.mark.parametrize( + "failed_finish_status", + [FinishStatus.FINISHED_ABORTED, FinishStatus.FINISHED_ERROR], +) +def test_pd_master_does_not_record_aborted_or_error_request_as_success(failed_finish_status): + async def run(): + manager = _manager() + sampling_params = SamplingParams() + sampling_params.n = 1 + sampling_params.best_of = 1 + sampling_params.max_new_tokens = 4 + + multimodal_params = MagicMock() + multimodal_params.verify_and_preload = AsyncMock() + request = MagicMock() + + async def generate_one(*_args, **_kwargs): + yield ( + 800, + "", + {"prompt_tokens": 2}, + FinishStatus(failed_finish_status), + ) + + manager._generate_one = generate_one + + with patch.object(HttpServerManager, "_check_and_repair_length", new=AsyncMock()): + results = [] + async for result in manager._generate("prompt", sampling_params, multimodal_params, request): + results.append(result) + + assert len(results) == 1 + assert results[0][3].get_status() == failed_finish_status + manager.metric_client.counter_inc.assert_called_once_with("lightllm_request_count") + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_multi_choice_failure_closes_other_generators(): async def run(): manager = _manager() @@ -217,7 +256,7 @@ async def run(): 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)) + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node, 0.0)) async def failing_wait_to_token_package(*_args, **_kwargs): raise RuntimeError("generation failed") @@ -239,6 +278,7 @@ async def failing_wait_to_token_package(*_args, **_kwargs): assert p_node.dispatched_prompt_chars == 0 assert p_node.dispatched_req_num == 0 + manager.pd_manager.selector.insert_prompt_cache.assert_not_called() asyncio.run(asyncio.wait_for(run(), timeout=2)) @@ -257,17 +297,19 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node, 0.0)) dispatched_nodes = [] dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] + high_priority_request_flags = [] async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args): dispatched_nodes.append(selected_p_node) dispatched_prompts.append(block_prompt) dispatched_loads.append(selected_p_node.dispatched_prompt_chars) dispatched_req_counts.append(selected_p_node.dispatched_req_num) + high_priority_request_flags.append(sampling_params.pd_high_priority_request) yield ( sampling_params.group_request_id, "x", @@ -293,6 +335,7 @@ async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_pro 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 p_node.dispatched_prompt_chars == other_request_load assert p_node.dispatched_req_num == other_request_count assert len(results) == 2 @@ -300,6 +343,93 @@ async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_pro asyncio.run(asyncio.wait_for(run(), timeout=2)) +@pytest.mark.parametrize( + ("estimated_cache_hit_rate", "expected_high_priority"), + [(0.8, False), (0.81, True)], +) +def test_pd_master_promotes_request_with_high_estimated_cache_hit_rate( + estimated_cache_hit_rate, + expected_high_priority, +): + 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() + 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, estimated_cache_hit_rate)) + high_priority_request_flags = [] + + async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling_params, *_args): + high_priority_request_flags.append(sampling_params.pd_high_priority_request) + yield ( + sampling_params.group_request_id, + "x", + {"prompt_tokens": 1}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) + + manager._wait_to_token_package = wait_to_token_package + + async for _ in manager._generate_one( + "prompt", + SamplingParams(), + MagicMock(), + MagicMock(), + 0, + 800, + ): + pass + + assert high_priority_request_flags == [expected_high_priority] + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + +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() + 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.81)) + captured_timeout_seconds = [] + + async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling_params, *_args): + captured_timeout_seconds.append(sampling_params.pd_high_priority_request_time_out_seconds) + yield ( + sampling_params.group_request_id, + "x", + {"prompt_tokens": 1}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) + + manager._wait_to_token_package = wait_to_token_package + + async for _ in manager._generate_one( + "prompt", + SamplingParams(), + MagicMock(), + MagicMock(), + 0, + 800, + ): + pass + + assert captured_timeout_seconds == [90] + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_releases_prefill_load_when_stream_is_closed(): async def run(): manager = _manager() @@ -314,7 +444,7 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) + 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() diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py new file mode 100644 index 0000000000..9e071cfa98 --- /dev/null +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -0,0 +1,190 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from lightllm.server.core.objs import SamplingParams +from lightllm.server.httpserver.manager import HttpServerManager, ReqStatus +from lightllm.server.router.req_queue.base_queue import BaseQueue +from lightllm.server.router.req_queue.dp_base_queue import DpQueue +from lightllm.utils.error_utils import ServerBusyError + + +class FakeSharedInt: + def __init__(self, value: int = 0): + self.value = value + + def get_value(self): + return self.value + + def set_value(self, value: int): + self.value = value + + +def _manager() -> HttpServerManager: + manager = HttpServerManager.__new__(HttpServerManager) + manager.args = SimpleNamespace(run_mode="decode", running_max_req_size=64) + manager.pd_node_request_limit_enabled = False + manager._run_reqs_count_lock = asyncio.Lock() + manager.run_reqs_count_mark = FakeSharedInt() + manager.latest_success_infer_time_mark = FakeSharedInt() + manager.pd_node_shm_req_alloc_timeout_seconds = 20 + manager.shm_req_manager = MagicMock() + manager.shm_req_manager.async_release_req_index = AsyncMock() + return manager + + +def test_pd_high_priority_timeout_is_internal_and_defaults_to_zero(): + sampling_params = SamplingParams() + sampling_params.init( + None, + pd_high_priority_request=True, + pd_high_priority_request_time_out_seconds=99, + ) + + assert sampling_params.pd_high_priority_request is False + assert sampling_params.pd_high_priority_request_time_out_seconds == 0 + + +def test_shm_req_partial_allocations_are_released_on_failure(): + async def run(): + manager = _manager() + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[3, RuntimeError("failed")]) + + with pytest.raises(RuntimeError, match="failed"): + await manager._alloc_shm_req_indexes(2) + + manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(3) + + asyncio.run(run()) + + +def test_shm_req_allocation_waits_forever_when_local_limit_is_disabled(): + async def run(): + manager = _manager() + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, None, 3]) + + with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()) as sleep: + assert await manager._alloc_shm_req_indexes(1) == [3] + + assert sleep.await_count == 2 + + asyncio.run(run()) + + +def test_high_priority_shm_req_allocation_uses_shorter_backoff_even_with_local_limit(): + async def run(): + manager = _manager() + manager.pd_node_request_limit_enabled = True + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 3]) + + with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()) as sleep: + assert await manager._alloc_shm_req_indexes(1, pd_high_priority_request=True) == [3] + + assert sleep.await_args_list[0].args[0] == pytest.approx(0.1 * 0.2) + + asyncio.run(run()) + + +def test_high_priority_shm_req_allocation_uses_master_timeout_with_local_limit(): + async def run(): + manager = _manager() + manager.pd_node_request_limit_enabled = True + manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=None) + + with ( + patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 161]), + pytest.raises(ServerBusyError, match="within 60 seconds"), + ): + await manager._alloc_shm_req_indexes( + 1, + pd_high_priority_request=True, + pd_high_priority_request_time_out_seconds=60, + ) + + asyncio.run(run()) + + +def test_high_priority_shm_req_allocation_does_not_shorten_local_timeout(): + async def run(): + manager = _manager() + manager.pd_node_request_limit_enabled = True + manager.pd_node_shm_req_alloc_timeout_seconds = 80 + manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=None) + + with ( + patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 181]), + pytest.raises(ServerBusyError, match="within 80 seconds"), + ): + await manager._alloc_shm_req_indexes( + 1, + pd_high_priority_request=True, + pd_high_priority_request_time_out_seconds=60, + ) + + asyncio.run(run()) + + +@pytest.mark.parametrize( + ("infer_start_time", "local_timeout_seconds", "high_priority_timeout_seconds", "expected"), + [(0, 20, 0, True), (0, 20, 60, True), (0, 80, 60, False), (1, 20, 60, False)], +) +def test_high_priority_router_wait_uses_master_timeout( + infer_start_time, + local_timeout_seconds, + high_priority_timeout_seconds, + expected, +): + req_status = ReqStatus.__new__(ReqStatus) + req_status.group_req_objs = SimpleNamespace( + shm_req_objs=[ + SimpleNamespace( + infer_start_time=infer_start_time, + router_arrival_time=100, + sample_params=SimpleNamespace( + pd_high_priority_request=True, + pd_high_priority_request_time_out_seconds=high_priority_timeout_seconds, + ), + ) + ] + ) + + with patch("lightllm.server.httpserver.manager.time.monotonic", return_value=161): + assert req_status.has_timed_out_waiting_for_inference(local_timeout_seconds) is expected + + +def test_pd_high_priority_request_is_inserted_before_first_normal_request(): + queue = BaseQueue.__new__(BaseQueue) + queue.dp_index = 0 + earlier_high_req = SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=True)) + normal_req = SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=False)) + high_req_1 = SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=True)) + high_req_2 = SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=True)) + queue.waiting_req_list = [earlier_high_req, normal_req] + + queue.extend([high_req_1, high_req_2]) + + assert queue.waiting_req_list == [earlier_high_req, high_req_1, high_req_2, normal_req] + assert high_req_1.sample_params.suggested_dp_index == 0 + assert high_req_2.sample_params.suggested_dp_index == 0 + + +def test_pd_high_priority_request_group_keeps_fifo_order_while_waiting_for_dp_index(): + queue = DpQueue.__new__(DpQueue) + queue.dp_size_in_node = 2 + earlier_high_group = [SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=True))] + normal_group = [SimpleNamespace(sample_params=SimpleNamespace(pd_high_priority_request=False))] + new_high_group = [ + SimpleNamespace( + sample_params=SimpleNamespace( + pd_high_priority_request=True, + suggested_dp_index=-1, + ) + ) + ] + queue.reqs_waiting_for_dp_index = [earlier_high_group, normal_group] + + queue.extend(new_high_group) + + assert queue.reqs_waiting_for_dp_index == [earlier_high_group, new_high_group, normal_group] diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 51f8d1c82c..0259c0c3e3 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -81,11 +81,21 @@ def test_final_token_metadata_read_returns_actual_prompt_tokens(req): def test_finish_status(req): req.finish_status.set_status(req.finish_status.FINISHED_STOP) assert req.finish_status.is_finished() + assert not req.finish_status.is_error_finished() assert req.finish_status.get_finish_reason() == "stop" + req.finish_status.set_status(req.finish_status.FINISHED_LENGTH) + assert req.finish_status.is_finished() + assert not req.finish_status.is_error_finished() + + req.finish_status.set_status(req.finish_status.FINISHED_ABORTED) + assert req.finish_status.is_finished() + assert req.finish_status.is_error_finished() + req.finish_status.set_status(req.finish_status.FINISHED_ERROR) assert req.finish_status.is_finished() assert req.finish_status.is_finished_error() + assert req.finish_status.is_error_finished() assert req.finish_status.get_finish_reason() == "error" diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 3c7a4242a0..ca4515ca15 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -13,7 +13,7 @@ ReqStatus, ) from lightllm.server.pd_io_struct import ObjType -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError class _FailingManager: @@ -59,6 +59,14 @@ async def generate(self, **_kwargs): yield +class _BusyManager: + args = SimpleNamespace(run_mode="decode") + + async def generate(self, **_kwargs): + raise ServerBusyError("decode node could not allocate a shm_req object within 20 seconds") + yield + + def test_pd_node_reports_generate_error_to_master(): async def run(): sampling_params = SamplingParams() @@ -86,6 +94,33 @@ async def run(): asyncio.run(run()) +def test_pd_node_reports_local_request_rejection_to_master(): + async def run(): + sampling_params = SamplingParams() + sampling_params.group_request_id = 123 + websocket = AsyncMock() + + await _pd_process_generate( + manager=_BusyManager(), + prompt="prompt", + sampling_params=sampling_params, + multimodal_params=MagicMock(), + forwarding_queue=MagicMock(), + pd_upload_websocket=websocket, + pd_event=asyncio.Event(), + ) + + websocket.send.assert_awaited_once() + obj = pickle.loads(websocket.send.await_args.args[0]) + assert obj == ( + ObjType.PD_UPLOAD_SERVER_BUSY, + 123, + "decode node could not allocate a shm_req object within 20 seconds", + ) + + asyncio.run(run()) + + def test_pd_node_cancellation_finishes_without_reporting_generate_error(): async def run(): sampling_params = SamplingParams() @@ -244,6 +279,41 @@ async def run(): asyncio.run(run()) +def test_pd_master_request_rejection_becomes_server_busy_error(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None) + manager.pd_manager = MagicMock() + manager.timer_log = AsyncMock() + manager.infos_queues = None + + req_status = ReqStatus(123, MagicMock(), MagicMock()) + manager.req_id_to_out_inf = {123: req_status} + + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + await manager.put_to_handle_queue( + ( + ObjType.PD_UPLOAD_SERVER_BUSY, + 123, + "prefill node could not allocate a shm_req object within 20 seconds", + ) + ) + await asyncio.wait_for(req_status.event.wait(), timeout=1) + + assert req_status.is_server_busy is True + with pytest.raises(ServerBusyError, match="prefill node could not allocate a shm_req object"): + req_status.raise_if_error() + finally: + handle_task.cancel() + with suppress(asyncio.CancelledError): + await handle_task + + asyncio.run(run()) + + @pytest.mark.parametrize("event_name", ["prefill_prompt_ids_event", "up_status_event"]) def test_pd_master_generate_error_wakes_resource_wait(event_name): async def run(): 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 a3dc268a93..d39ff1cc73 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -27,11 +27,15 @@ def gen_id(): mgr.tokens = lambda *a, **k: 10 mgr._log_req_header = lambda *a, **k: asyncio.sleep(0) mgr.recorded_cache_hit_rates = [] + mgr.inserted_prompt_caches = [] mgr.pd_manager = SimpleNamespace( - selector=SimpleNamespace(record_prompt_cache_hit_rate=mgr.recorded_cache_hit_rates.append) + selector=SimpleNamespace( + record_prompt_cache_hit_rate=mgr.recorded_cache_hit_rates.append, + insert_prompt_cache=lambda prompt, p_node: mgr.inserted_prompt_caches.append((prompt, p_node)), + ) ) p_node = SimpleNamespace(dispatched_prompt_chars=0, dispatched_req_num=0) - mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(p_node, 1)) + mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(p_node, 1, 0.0)) mgr.remove_req = lambda *a, **k: asyncio.sleep(0) return mgr @@ -73,6 +77,8 @@ def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3]) assert cached and all(c == 30 for c in cached), cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] + assert len(mgr.inserted_prompt_caches) == 1 + assert mgr.inserted_prompt_caches[0][0] == "hello" def test_multi_block_keeps_first_block_hit(monkeypatch): @@ -85,3 +91,42 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, 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_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 + sampling_params.max_new_tokens = 1 + + 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}, + FinishStatus(FinishStatus.FINISHED_ERROR), + ) + + monkeypatch.setattr(mgr, "_wait_to_token_package", failed_wait) + + async def run(): + async for _ in mgr.generate( + prompt="hello", + sampling_params=sampling_params, + multimodal_params=SimpleNamespace( + images=[], + audios=[], + verify_and_preload=lambda req: asyncio.sleep(0), + ), + request=None, + ): + pass + + asyncio.run(run()) + + assert mgr.recorded_cache_hit_rates == [pytest.approx(0.2)] + assert mgr.inserted_prompt_caches == [] diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..bf154826e8 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -6,9 +6,9 @@ import pytest from lightllm.server.core.objs import SamplingParams -from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.server.httpserver.manager import HttpServerManager, ReqStatus from lightllm.server.pd_io_struct import NodeRole, ObjType -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError class _ValueMark: @@ -24,7 +24,16 @@ def set_value(self, value): def _make_manager(mode: NodeRole): manager = HttpServerManager.__new__(HttpServerManager) + manager.args = SimpleNamespace( + enable_pd_node_self_request_limit=False, + run_mode=mode.value, + running_max_req_size=2, + ) manager.pd_mode = mode + manager.is_multinode_tp_slave = False + manager.pd_node_request_limit_enabled = False + manager.pd_node_shm_req_alloc_timeout_seconds = 20 + manager.pd_node_router_wait_timeout_seconds = 20 manager.alloc_req_id = MagicMock(return_value=123) manager.is_multinode_tp_master = False manager.rl_controller = None @@ -38,6 +47,8 @@ def _make_manager(mode: NodeRole): manager._register_running_request = AsyncMock() manager._unregister_running_request = AsyncMock() manager.metric_client = MagicMock() + manager._run_reqs_count_lock = asyncio.Lock() + manager.run_reqs_count_mark = _ValueMark() manager.shm_req_manager = SimpleNamespace(async_alloc_req_index=AsyncMock(side_effect=RuntimeError("alloc failed"))) return manager @@ -187,6 +198,115 @@ async def run(): asyncio.run(run()) +@pytest.mark.parametrize("mode", [NodeRole.P, NodeRole.D]) +def test_pd_node_returns_busy_when_shm_req_allocation_times_out(mode): + async def run(): + manager = _make_manager(mode) + manager.pd_node_request_limit_enabled = True + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_index=AsyncMock(return_value=None), + async_release_req_index=AsyncMock(), + ) + + websocket = AsyncMock() if mode == NodeRole.P else None + pd_event = None + if mode == NodeRole.P: + pd_event = asyncio.Event() + pd_event.decode_node_info = SimpleNamespace(ready_kv_len=1) + pd_event.set() + + with patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[0, 21]): + with pytest.raises(ServerBusyError, match=f"PD {mode.value} node is busy"): + await _drain_generate( + manager, + _sampling_params(), + _multimodal_params(), + websocket, + pd_event, + ) + + manager.shm_req_manager.async_alloc_req_index.assert_awaited_once() + manager.shm_req_manager.async_release_req_index.assert_not_awaited() + manager._register_running_request.assert_awaited_once() + manager._unregister_running_request.assert_awaited_once() + + asyncio.run(run()) + + +@pytest.mark.parametrize("mode", [NodeRole.P, NodeRole.D]) +def test_pd_node_returns_busy_while_first_token_request_waits_in_router(mode): + async def run(): + manager = _make_manager(mode) + manager.pd_node_request_limit_enabled = True + req = SimpleNamespace( + request_id=123, + is_aborted=False, + router_arrival_time=0, + infer_start_time=0, + sample_params=SimpleNamespace(pd_high_priority_request=False), + ) + req_status = ReqStatus(123, None, [req], 0) + req_status.event.set() + sampling_params = _sampling_params() + + with patch("lightllm.server.httpserver.manager.time.monotonic", return_value=21): + req.router_arrival_time = 1.0 + output_generator = manager._wait_to_token_package( + start_time=0, + prompt_ids=[], + group_request_id=123, + sampling_params=sampling_params, + req_status=req_status, + request=None, + ) + with pytest.raises(ServerBusyError, match="request did not enter inference"): + await output_generator.__anext__() + + asyncio.run(run()) + + +def test_httpserver_keeps_started_and_high_priority_request_groups(): + waiting_req = SimpleNamespace( + router_arrival_time=1.0, + infer_start_time=0.0, + sample_params=SimpleNamespace( + pd_high_priority_request=False, + pd_high_priority_request_time_out_seconds=60, + ), + ) + started_req = SimpleNamespace( + router_arrival_time=1.0, + infer_start_time=2.0, + sample_params=SimpleNamespace(pd_high_priority_request=False), + ) + req_status = ReqStatus(123, None, [waiting_req, started_req], 0) + + with patch("lightllm.server.httpserver.manager.time.monotonic", return_value=30): + assert req_status.has_timed_out_waiting_for_inference(20) is False + + started_req.infer_start_time = 0.0 + waiting_req.sample_params.pd_high_priority_request = True + assert req_status.has_timed_out_waiting_for_inference(20) is False + + +def test_pd_node_self_request_limit_releases_partially_allocated_shm_reqs(): + async def run(): + manager = _make_manager(NodeRole.D) + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_index=AsyncMock(side_effect=[7, RuntimeError("alloc failed")]), + async_release_req_index=AsyncMock(), + ) + sampling_params = _sampling_params() + sampling_params.n = 2 + + with pytest.raises(RuntimeError, match="alloc failed"): + await _drain_generate(manager, sampling_params, _multimodal_params()) + + manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(7) + + asyncio.run(run()) + + def test_running_request_helpers_are_atomic_and_refresh_timestamp_only_on_idle_to_running_transition(): async def run(): manager = HttpServerManager.__new__(HttpServerManager) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_prefill_abort.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_prefill_abort.py new file mode 100644 index 0000000000..bb732a0eba --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_prefill_abort.py @@ -0,0 +1,168 @@ +import queue +import threading +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from lightllm.server.pd_io_struct import PDAbortReq +from lightllm.server.router.model_infer.mode_backend.pd.decode_node_impl import decode_impl +from lightllm.server.router.model_infer.mode_backend.pd.decode_node_impl.decode_impl import PDDecodeNode +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl import prefill_impl, prefill_trans_process +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_impl import ( + PDChunkedPrefillForPrefillNode, +) +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_kv_move_manager import ( + PrefillKVMoveManager, +) +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_trans_process import ( + _PrefillTransModule, +) + + +class _StopLoop(Exception): + pass + + +class _SequenceQueue: + def __init__(self, *values): + self.values = list(values) + + def get(self): + if self.values: + return self.values.pop(0) + raise _StopLoop() + + +def test_prefill_infer_forwards_abort_to_kv_move_manager(monkeypatch): + request_id = 123 + req = SimpleNamespace( + req_id=request_id, + infer_aborted=True, + pd_task_num=2, + pd_task_failed_num=0, + pd_task_success_num=0, + pd_trans_device_id=1, + cur_kv_len=4, + shm_req=SimpleNamespace(input_len=8), + ) + monkeypatch.setattr(prefill_impl.g_infer_context, "requests_mapping", {request_id: req}) + + backend = PDChunkedPrefillForPrefillNode.__new__(PDChunkedPrefillForPrefillNode) + backend.is_master_in_dp = True + backend.info_queue = queue.Queue() + + ready_reqs = backend._filter_not_ready_reqs([request_id]) + + assert ready_reqs == [] + abort_req = backend.info_queue.get_nowait() + assert abort_req == PDAbortReq(request_id=request_id, device_id=1) + assert req.pd_abort_req_send_count == 1 + + +def test_prefill_infer_sends_abort_at_most_six_times(monkeypatch): + request_id = 123 + req = SimpleNamespace( + req_id=request_id, + infer_aborted=True, + pd_task_num=2, + pd_task_failed_num=0, + pd_task_success_num=0, + pd_trans_device_id=1, + cur_kv_len=4, + shm_req=SimpleNamespace(input_len=8), + ) + monkeypatch.setattr(prefill_impl.g_infer_context, "requests_mapping", {request_id: req}) + + backend = PDChunkedPrefillForPrefillNode.__new__(PDChunkedPrefillForPrefillNode) + backend.is_master_in_dp = True + backend.info_queue = queue.Queue() + + for _ in range(8): + assert backend._filter_not_ready_reqs([request_id]) == [] + + abort_reqs = [backend.info_queue.get_nowait() for _ in range(6)] + assert all(abort_req == PDAbortReq(request_id=request_id, device_id=1) for abort_req in abort_reqs) + assert backend.info_queue.empty() + assert req.pd_abort_req_send_count == 6 + + +def test_decode_infer_sends_abort_at_most_six_times(monkeypatch): + request_id = 123 + req = SimpleNamespace( + req_id=request_id, + infer_aborted=True, + pd_task_num=2, + pd_task_failed_num=0, + pd_task_success_num=0, + pd_trans_device_id=1, + ) + monkeypatch.setattr(decode_impl.g_infer_context, "requests_mapping", {request_id: req}) + + backend = PDDecodeNode.__new__(PDDecodeNode) + backend.is_master_in_dp = True + backend.info_queue = queue.Queue() + + for _ in range(8): + assert backend._filter_not_ready_reqs([request_id]) == [] + + abort_reqs = [backend.info_queue.get_nowait() for _ in range(6)] + assert all(abort_req == PDAbortReq(request_id=request_id, device_id=1) for abort_req in abort_reqs) + assert backend.info_queue.empty() + assert req.pd_abort_req_send_count == 6 + + +def test_prefill_kv_move_manager_routes_abort_to_target_device(): + abort_req = PDAbortReq(request_id=123, device_id=1) + target_queue = queue.Queue() + manager = PrefillKVMoveManager.__new__(PrefillKVMoveManager) + manager.info_queue = _SequenceQueue(abort_req) + manager.kv_trans_processes = [ + SimpleNamespace(task_in_queue=queue.Queue()), + SimpleNamespace(task_in_queue=target_queue), + ] + + with pytest.raises(_StopLoop): + manager.task_dispatcher_loop() + + assert target_queue.get_nowait() is abort_req + assert manager.kv_trans_processes[0].task_in_queue.empty() + + +def test_prefill_trans_process_handles_abort_without_allocating_page(monkeypatch): + abort_req = PDAbortReq(request_id=123, device_id=0) + module = _PrefillTransModule.__new__(_PrefillTransModule) + module.device_id = 0 + module.task_in_queue = _SequenceQueue(abort_req) + module.page_index_queue = MagicMock() + module._abort = MagicMock() + monkeypatch.setattr(prefill_trans_process.torch.cuda, "set_device", lambda _device_id: None) + + with pytest.raises(_StopLoop): + module.recv_task_loop() + + module._abort.assert_called_once_with(request_id=123) + module.page_index_queue.get.assert_not_called() + + +def test_prefill_abort_only_fails_tasks_not_submitted_to_transporter(): + pending_task = SimpleNamespace(request_id=123, xfer_handle=None, error_info=None) + active_task = SimpleNamespace(request_id=123, xfer_handle=object(), error_info=None) + unrelated_task = SimpleNamespace(request_id=456, xfer_handle=None, error_info=None) + module = _PrefillTransModule.__new__(_PrefillTransModule) + module.waiting_dict_lock = threading.Lock() + module.waiting_dict = { + "pending": pending_task, + "active": active_task, + "unrelated": unrelated_task, + } + module.failed_queue = queue.Queue() + + module._abort(request_id=123) + + assert module.failed_queue.get_nowait() is pending_task + assert pending_task.error_info == "aborted req" + assert module.waiting_dict == { + "active": active_task, + "unrelated": unrelated_task, + } diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 6bdde78576..8cc4d8775e 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -7,6 +7,7 @@ CacheAwareConfig, CacheAwarePolicy, ) +from lightllm.server.httpserver_for_pd_master.pd_selector.pd_selector import LoadBalancedCacheAwareSelector def _worker(address: str, dispatched_prompt_chars: int = 0, dispatched_req_num: int = 0): @@ -73,6 +74,37 @@ def test_cache_aware_updates_threshold_from_inference_cache_hit_rate(): assert policy.config.balance_rel_threshold == pytest.approx(1.55) +def test_cache_aware_selector_returns_selected_nodes_and_estimated_cache_hit_rate(): + selector = LoadBalancedCacheAwareSelector(pd_manager=None) + p_node = _worker("10.0.0.1:8000") + d_node = SimpleNamespace( + client_ip_port="10.0.0.2:8000", + run_status=SimpleNamespace(total_token_usage_rate=0.0), + ) + selector.update_nodes([p_node], [d_node]) + prompt = "x" * 1025 + selector.policy.prompt_cache_tree.insert(prompt, p_node.client_ip_port) + + selected_p_node, selected_d_node, estimated_cache_hit_rate = selector.select_p_d_node(prompt, None, None) + + assert selected_p_node is p_node + assert selected_d_node is d_node + assert estimated_cache_hit_rate == pytest.approx(1.0) + + +def test_cache_aware_inserts_prompt_only_when_explicitly_recorded(): + policy = CacheAwarePolicy() + worker = _worker("10.0.0.1:8000") + prompt = "new prompt" + + assert policy.select_worker([worker], prompt) is worker + assert policy.prompt_cache_tree.prefix_match(prompt).prefill_node is None + + policy.insert_prompt_cache(prompt, worker) + + assert policy.prompt_cache_tree.prefix_match(prompt).prefill_node == worker.client_ip_port + + def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy = CacheAwarePolicy() cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) @@ -81,8 +113,11 @@ def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt) + estimated_cache_hit_rate = policy.get_estimated_cache_hit_rate(selected_worker, prompt) assert selected_worker is cache_worker + # 前缀树每 512 个字符抽样一次,命中率保持原有的保守估算方式。 + assert estimated_cache_hit_rate == pytest.approx(1025 / len(prompt)) def test_cache_aware_uses_least_loaded_worker_when_cache_worker_is_overloaded(): @@ -93,8 +128,10 @@ def test_cache_aware_uses_least_loaded_worker_when_cache_worker_is_overloaded(): policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt) + estimated_cache_hit_rate = policy.get_estimated_cache_hit_rate(selected_worker, prompt) assert selected_worker is least_loaded_worker + assert estimated_cache_hit_rate == 0.0 def test_cache_aware_keeps_cache_worker_when_both_workers_are_idle(): diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94b..8d064c357b 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -5,10 +5,25 @@ import pytest from easydict import EasyDict +from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled(): + parser = make_argument_parser() + + assert parser.parse_args([]).enable_pd_node_self_request_limit is False + assert parser.parse_args(["--enable_pd_node_self_request_limit"]).enable_pd_node_self_request_limit is True + assert StartArgs().enable_pd_node_self_request_limit is False + + +def test_pd_master_request_slot_is_reserved_noop(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + + asyncio.run(manager._wait_for_pd_master_request_slot()) + + def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): from lightllm.utils.config_utils import auto_set_response_parsers diff --git a/unit_tests/utils/test_envs_utils.py b/unit_tests/utils/test_envs_utils.py new file mode 100644 index 0000000000..a806f8c6de --- /dev/null +++ b/unit_tests/utils/test_envs_utils.py @@ -0,0 +1,40 @@ +from lightllm.utils.envs_utils import ( + get_pd_node_router_wait_timeout_seconds, + get_pd_node_shm_req_alloc_timeout_seconds, +) + + +def test_pd_node_shm_req_alloc_timeout_defaults_to_20_seconds(monkeypatch): + monkeypatch.delenv("LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS", raising=False) + get_pd_node_shm_req_alloc_timeout_seconds.cache_clear() + + assert get_pd_node_shm_req_alloc_timeout_seconds() == 20 + + get_pd_node_shm_req_alloc_timeout_seconds.cache_clear() + + +def test_pd_node_shm_req_alloc_timeout_reads_environment_variable(monkeypatch): + monkeypatch.setenv("LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS", "30") + get_pd_node_shm_req_alloc_timeout_seconds.cache_clear() + + assert get_pd_node_shm_req_alloc_timeout_seconds() == 30 + + get_pd_node_shm_req_alloc_timeout_seconds.cache_clear() + + +def test_pd_node_router_wait_timeout_defaults_to_20_seconds(monkeypatch): + monkeypatch.delenv("LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS", raising=False) + get_pd_node_router_wait_timeout_seconds.cache_clear() + + assert get_pd_node_router_wait_timeout_seconds() == 20 + + get_pd_node_router_wait_timeout_seconds.cache_clear() + + +def test_pd_node_router_wait_timeout_reads_environment_variable(monkeypatch): + monkeypatch.setenv("LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS", "45") + get_pd_node_router_wait_timeout_seconds.cache_clear() + + assert get_pd_node_router_wait_timeout_seconds() == 45 + + get_pd_node_router_wait_timeout_seconds.cache_clear()