From a2fb61f619b71fb64406e788e7b0a4102efa3ae3 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 05:36:52 +0000 Subject: [PATCH 01/20] fix(pd): remove master decode capacity limit --- docs/CN/source/tutorial/api_server_args.rst | 2 ++ docs/EN/source/tutorial/api_server_args.rst | 2 ++ lightllm/server/api_cli.py | 5 ----- lightllm/server/core/objs/start_args_type.py | 1 - lightllm/server/httpserver_for_pd_master/manager.py | 5 ----- 5 files changed, 4 insertions(+), 11 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index fbb63d09f0..a6e9cb7eb5 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -148,6 +148,8 @@ PD 分离模式参数 .. option:: --running_max_req_size 同时进行前向推理的最大请求数量,默认为 ``1000`` + 在 PD 分离模式的 Decode 节点上,该限制仅在各节点本地生效; + PD Master 不会汇总各 Decode 节点的值作为全局请求准入上限。 .. option:: --max_req_total_len diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 69edf50a86..6a205411b0 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -150,6 +150,8 @@ Memory and Batch Processing Parameters .. option:: --running_max_req_size Maximum number of requests for simultaneous forward inference, default is ``1000`` + On Decode nodes in PD disaggregation mode, this limit applies locally to each node; + PD Master does not aggregate the Decode-node values into a global admission limit. .. option:: --max_req_total_len diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e8..d2d6ddb943 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -68,11 +68,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "dynamically; use

pd for a fixed topology, for example 2p4d. Default: elastic." ), ) - parser.add_argument( - "--disable_pd_master_decode_capacity_limit", - action="store_true", - help="Disable PD master admission control based on the total capacity of registered decode nodes.", - ) parser.add_argument( "--pd_trans_mode", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index a9aef608bd..5786de04d9 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -22,7 +22,6 @@ 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) 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_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index b422bf7703..aebe039426 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -129,11 +129,6 @@ 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() - was_idle = self.running_request_count == 0 self.running_request_count += 1 if was_idle: From 85a54dc949706463c0e9ccc861871b626fed5fae Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 06:21:34 +0000 Subject: [PATCH 02/20] feat(pd): reject requests on local allocation timeout --- docs/CN/source/tutorial/api_server_args.rst | 8 +++ docs/EN/source/tutorial/api_server_args.rst | 9 +++ lightllm/server/api_cli.py | 8 +++ lightllm/server/core/objs/start_args_type.py | 1 + lightllm/server/httpserver/manager.py | 59 +++++++++++---- lightllm/server/httpserver/pd_loop.py | 9 ++- .../httpserver_for_pd_master/manager.py | 16 ++++- lightllm/server/pd_io_struct.py | 1 + lightllm/utils/envs_utils.py | 6 ++ .../httpserver/test_pd_generate_error.py | 72 ++++++++++++++++++- .../test_running_request_lifecycle.py | 51 ++++++++++++- unit_tests/server/test_pd_master_mode.py | 9 +++ unit_tests/utils/test_envs_utils.py | 19 +++++ 13 files changed, 251 insertions(+), 17 deletions(-) create mode 100644 unit_tests/utils/test_envs_utils.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index a6e9cb7eb5..195b4b1753 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -87,6 +87,14 @@ PD 分离模式参数 推理进度健康检查:当仍有在途请求,且整个 PD Master 连续 ``HEALTH_TIMEOUT`` 秒 没有任何请求成功返回 token 时,接口将返回 HTTP 503。 +.. option:: --enable_pd_node_self_request_limit + + 在 PD 分离模式的 Prefill 和 Decode 节点上启用本地请求准入控制。启用后,如果节点在指定超时时间内 + 无法为请求分配本地 ``shm_req`` 对象,会主动拒绝该请求,并由 PD Master 向客户端返回 + HTTP 429。超时时间通过环境变量 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` 设置,单位为秒, + 默认值为 20。例如,设置 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS=30`` 表示最长等待 30 秒。 + 该参数默认关闭,在 ``normal`` 和 ``pd_master`` 模式下不生效。 + .. 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 6a205411b0..04862d23cc 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -89,6 +89,15 @@ 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 admission control on Prefill and Decode nodes in PD disaggregation mode. When enabled, + a node rejects a request through PD Master with HTTP 429 if it cannot allocate the request's local + ``shm_req`` object within the configured timeout. The timeout is controlled by the + ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` environment variable and defaults to 20 seconds. + For example, set ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS=30`` to wait for 30 seconds. + The option is disabled by default and has no effect in ``normal`` or ``pd_master`` mode. + .. 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 d2d6ddb943..0870dd1b21 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -68,6 +68,14 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "dynamically; use

pd for a fixed topology, for example 2p4d. Default: elastic." ), ) + parser.add_argument( + "--enable_pd_node_self_request_limit", + action="store_true", + help=( + "Allow Prefill and Decode nodes in PD mode to reject a request when no local shm_req object " + "can be allocated within 20 seconds. Default: disabled." + ), + ) parser.add_argument( "--pd_trans_mode", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 5786de04d9..0d8cc4ac7e 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -22,6 +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") + 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..229bc6f793 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -36,9 +36,9 @@ 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_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 +117,13 @@ def __init__( self.pd_mode: NodeRole = NodeRole(self.args.run_mode) assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D] + # 该开关只对 PD 分离模式的 Prefill/Decode 服务主节点生效。 + # 多机 TP 的从节点不直接与 PD Master 通信,不能独立拒绝请求。 + 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 + ) + # 超时时间由环境变量 LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS 控制,默认为 20 秒。 + self.pd_node_shm_req_alloc_timeout_seconds = get_pd_node_shm_req_alloc_timeout_seconds() self.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() @@ -427,17 +434,7 @@ async def generate( 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) + alloced_req_indexes = await self._alloc_shm_req_indexes(sampling_params.n) 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 +540,42 @@ def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple return image_tokens, audio_tokens + async def _alloc_shm_req_indexes(self, req_num: int) -> List[int]: + """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。""" + alloced_req_indexes = [] + alloc_deadline = ( + time.monotonic() + self.pd_node_shm_req_alloc_timeout_seconds + if self.pd_node_request_limit_enabled + else None + ) + + try: + while len(alloced_req_indexes) < req_num: + alloc_req_index = await self.shm_req_manager.async_alloc_req_index() + sleep_time = 0.1 + while alloc_req_index is None: + if alloc_deadline is not None: + remaining_time = alloc_deadline - time.monotonic() + if remaining_time <= 0: + logger.warning( + f"PD {self.args.run_mode} node shm_req allocation timed out after " + f"{self.pd_node_shm_req_alloc_timeout_seconds} seconds" + ) + raise ServerBusyError( + f"PD {self.args.run_mode} node is busy: unable to allocate a shm_req object " + f"within {self.pd_node_shm_req_alloc_timeout_seconds} seconds" + ) + await asyncio.sleep(sleep_time) + 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", "") 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 aebe039426..f518109f11 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -672,6 +672,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: @@ -698,6 +708,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: @@ -705,9 +716,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() @@ -716,6 +728,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}" 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/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 933728b7ff..230df3bcee 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -311,6 +311,12 @@ 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_lightllm_url_pool_maxsize() -> int: return int(os.getenv("LIGHTLLM_URL_POOL_MAXSIZE", 512)) 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_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..5d99f7215a 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -8,7 +8,7 @@ from lightllm.server.core.objs import SamplingParams from lightllm.server.httpserver.manager import HttpServerManager 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,15 @@ 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.alloc_req_id = MagicMock(return_value=123) manager.is_multinode_tp_master = False manager.rl_controller = None @@ -187,6 +195,47 @@ async def run(): asyncio.run(run()) +def test_pd_node_self_request_limit_rejects_after_shm_req_allocation_timeout(): + async def run(): + manager = _make_manager(NodeRole.D) + 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(), + ) + + with patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 120]): + with pytest.raises(ServerBusyError, match="unable to allocate a shm_req object within 20 seconds"): + await _drain_generate(manager, _sampling_params(), _multimodal_params()) + + 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()) + + +def test_pd_node_self_request_limit_releases_partially_allocated_shm_reqs(): + async def run(): + manager = _make_manager(NodeRole.D) + manager.pd_node_request_limit_enabled = True + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_index=AsyncMock(side_effect=[7, None]), + async_release_req_index=AsyncMock(), + ) + sampling_params = _sampling_params() + sampling_params.n = 2 + + with patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 120]): + with pytest.raises(ServerBusyError, match="unable to allocate a shm_req object within 20 seconds"): + 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/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94b..981b250f4e 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -5,10 +5,19 @@ 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_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..7da944314a --- /dev/null +++ b/unit_tests/utils/test_envs_utils.py @@ -0,0 +1,19 @@ +from lightllm.utils.envs_utils import 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() From c5b9a10297b52c160a85b390ef2576d911a5d991 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 06:39:05 +0000 Subject: [PATCH 03/20] docs(pd): note unsafe transfer page recycling --- .../pd/decode_node_impl/decode_trans_process.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) 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) From 69b62acccf59c1a5865e5fef58833f96d164cfe2 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 07:13:42 +0000 Subject: [PATCH 04/20] fix(pd): propagate aborts to prefill transfers --- .../pd/decode_node_impl/decode_impl.py | 11 +- .../pd/prefill_node_impl/prefill_impl.py | 15 +- .../prefill_kv_move_manager.py | 17 +- .../prefill_trans_process.py | 19 +- .../mode_backend/test_pd_prefill_abort.py | 168 ++++++++++++++++++ 5 files changed, 220 insertions(+), 10 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_pd_prefill_abort.py 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/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/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, + } From 76e4f99c2e045cde766e1db8e01f43851bfdb454 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 07:45:53 +0000 Subject: [PATCH 05/20] fix(pd): preserve split requests under node rate limiting --- lightllm/server/core/objs/sampling_params.py | 5 ++ lightllm/server/httpserver/manager.py | 21 ++++++-- .../httpserver_for_pd_master/manager.py | 4 ++ .../test_pd_master_multi_choice.py | 3 ++ .../test_pd_node_request_limit.py | 49 +++++++++++++++++++ 5 files changed, 78 insertions(+), 4 deletions(-) create mode 100644 test/test_pd_selector/test_pd_node_request_limit.py diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 8e31c50624..9a13f44842 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -294,6 +294,8 @@ 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 为分段续跑请求设置,避免已经完成首段的请求因本地 shm_req 限流失败。 + ("bypass_pd_node_request_limit", ctypes.c_bool), ("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 +339,8 @@ 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.bypass_pd_node_request_limit = False self.suggested_dp_index = kwargs.get("suggested_dp_index", -1) self.skip_special_tokens = kwargs.get("skip_special_tokens", SKIP_SPECIAL_TOKENS) @@ -503,6 +507,7 @@ 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, + "bypass_pd_node_request_limit": self.bypass_pd_node_request_limit, "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/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 229bc6f793..2e74f78604 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -433,8 +433,17 @@ async def generate( await self._register_running_request() running_request_registered = True - # 申请资源并存储 - alloced_req_indexes = await self._alloc_shm_req_indexes(sampling_params.n) + # 申请资源并存储。PD 分段续跑请求可以绕过本地 shm_req 分配超时,避免 + # 已经成功完成首段的用户请求因为下一段暂时拿不到对象而被 429 中断。 + if self.pd_node_request_limit_enabled and sampling_params.bypass_pd_node_request_limit: + logger.info( + f"PD {self.args.run_mode} node request {group_request_id} bypasses the local shm_req " + f"allocation timeout and will wait for {sampling_params.n} object(s)" + ) + alloced_req_indexes = await self._alloc_shm_req_indexes( + sampling_params.n, + bypass_pd_node_request_limit=sampling_params.bypass_pd_node_request_limit, + ) 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) @@ -540,12 +549,16 @@ def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple return image_tokens, audio_tokens - async def _alloc_shm_req_indexes(self, req_num: int) -> List[int]: + async def _alloc_shm_req_indexes( + self, + req_num: int, + bypass_pd_node_request_limit: bool = False, + ) -> List[int]: """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。""" alloced_req_indexes = [] alloc_deadline = ( time.monotonic() + self.pd_node_shm_req_alloc_timeout_seconds - if self.pd_node_request_limit_enabled + if self.pd_node_request_limit_enabled and not bypass_pd_node_request_limit else None ) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index f518109f11..89505be014 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -228,6 +228,10 @@ 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 + # 首段仍然遵守 P/D 节点的本地限流。第二段及后续分段说明该用户请求已经 + # 成功执行过一段,为避免续跑请求因暂时申请不到 shm_req 而中途失败, + # 允许它们等待可用对象,不应用本地 shm_req 分配超时。 + sampling_params.bypass_pd_node_request_limit = iter_index > 0 # 分段请求始终复用循环外选定的 P 节点;这里只按每段实际发送的 # prompt 更新该节点的在途 prefill 负载,不会重新选点。 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..088bdb74cc 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -262,12 +262,14 @@ async def run(): dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] + bypass_request_limit_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) + bypass_request_limit_flags.append(sampling_params.bypass_pd_node_request_limit) yield ( sampling_params.group_request_id, "x", @@ -293,6 +295,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 bypass_request_limit_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 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..14d8ca00d1 --- /dev/null +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -0,0 +1,49 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.utils.error_utils import ServerBusyError + + +def _manager(*, timeout_seconds: int = 1) -> HttpServerManager: + manager = HttpServerManager.__new__(HttpServerManager) + manager.args = SimpleNamespace(run_mode="decode") + manager.pd_node_request_limit_enabled = True + manager.pd_node_shm_req_alloc_timeout_seconds = timeout_seconds + manager.shm_req_manager = MagicMock() + manager.shm_req_manager.async_release_req_index = AsyncMock() + return manager + + +def test_pd_node_regular_request_times_out_when_shm_req_is_unavailable(): + async def run(): + manager = _manager(timeout_seconds=0) + manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=None) + + with pytest.raises(ServerBusyError): + await manager._alloc_shm_req_indexes( + 1, + bypass_pd_node_request_limit=False, + ) + + asyncio.run(run()) + + +def test_pd_split_continuation_bypasses_shm_req_allocation_timeout(): + async def run(): + manager = _manager(timeout_seconds=0) + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 7]) + + with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()): + indexes = await manager._alloc_shm_req_indexes( + 1, + bypass_pd_node_request_limit=True, + ) + + assert indexes == [7] + manager.shm_req_manager.async_release_req_index.assert_not_awaited() + + asyncio.run(run()) From 87ad3e285583f3924c0aa507ce73fbd41b5c1133 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 09:53:46 +0000 Subject: [PATCH 06/20] feat(pd): add QPS-based node admission control --- docs/CN/source/tutorial/api_server_args.rst | 20 ++- docs/EN/source/tutorial/api_server_args.rst | 21 +++- lightllm/server/api_cli.py | 4 +- lightllm/server/httpserver/manager.py | 49 ++++---- lightllm/server/httpserver/qps_recorder.py | 79 ++++++++++++ lightllm/utils/envs_utils.py | 24 +++- test/test_api/test_qps_recorder.py | 118 ++++++++++++++++++ .../test_pd_node_request_limit.py | 79 ++++++++++-- 8 files changed, 345 insertions(+), 49 deletions(-) create mode 100644 lightllm/server/httpserver/qps_recorder.py create mode 100644 test/test_api/test_qps_recorder.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 195b4b1753..308b0b9d4e 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -89,10 +89,22 @@ PD 分离模式参数 .. option:: --enable_pd_node_self_request_limit - 在 PD 分离模式的 Prefill 和 Decode 节点上启用本地请求准入控制。启用后,如果节点在指定超时时间内 - 无法为请求分配本地 ``shm_req`` 对象,会主动拒绝该请求,并由 PD Master 向客户端返回 - HTTP 429。超时时间通过环境变量 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` 设置,单位为秒, - 默认值为 20。例如,设置 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS=30`` 表示最长等待 30 秒。 + 在 PD 分离模式的 Prefill 和 Decode 节点上启用本地请求准入控制。启用后,节点根据动态 QPS 和 + 当前运行请求总数决定新请求是否可以进入本地 ``shm_req`` 申请流程。超过上限时,节点会主动拒绝新请求,并由 PD Master + 向客户端返回 HTTP 429;已经准入的请求会持续等待可用的 ``shm_req`` 对象,不再进行等待超时判断。 + 服务启动后,QPS 记录器累计完成的请求数未达到 ``running_max_req_size`` 时,最大允许进入请求数直接使用 + ``running_max_req_size``,使冷启动阶段可以快速接收请求并积累足够的完成样本。累计完成请求数达到 + ``running_max_req_size`` 且首个 16 请求 QPS 窗口生成后,最大允许进入请求数改为 + ``int(QPS * 平均整包时间秒数) + 6``。如果 ``running_max_req_size`` 小于 16,节点会继续使用基础容量, + 直到 QPS 完成初始化,避免使用尚未初始化的零值 QPS 过早切换到仅 6 个探测请求。 + 稳定运行阶段额外放行 6 个请求作为探测余量,避免低流量或长时间空闲导致 QPS 降低后, + 系统被限制在过低的并发并且难以重新爬升。 + 这里的平均整包时间表示希望一个请求从进入 + 节点到整包完成所处的时间范围,用于将完成 QPS 换算为合理的在途请求数量。Prefill 阶段通常只负责输入 + 处理和首 token,默认值为 20 秒,避免短阶段积压过多请求;Decode 阶段需要持续生成 token,默认值为 + 60 秒,以覆盖更长的整包处理时间。可以通过环境变量 + ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 显式设置统一值;设置后 Prefill 和 Decode + 节点都会采用该值。 该参数默认关闭,在 ``normal`` 和 ``pd_master`` 模式下不生效。 .. 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 04862d23cc..4a0bbd2fd0 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -92,10 +92,23 @@ PD disaggregation Mode Parameters .. option:: --enable_pd_node_self_request_limit Enable local admission control on Prefill and Decode nodes in PD disaggregation mode. When enabled, - a node rejects a request through PD Master with HTTP 429 if it cannot allocate the request's local - ``shm_req`` object within the configured timeout. The timeout is controlled by the - ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` environment variable and defaults to 20 seconds. - For example, set ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS=30`` to wait for 30 seconds. + each node uses the dynamically measured QPS and its total running request count to decide whether a new request + may enter local ``shm_req`` allocation. Once the limit is exceeded, the node rejects new requests through PD Master with HTTP 429. Admitted requests + keep waiting for an available ``shm_req`` object without an allocation timeout. While the QPS recorder has completed + fewer than ``running_max_req_size`` requests since startup, that base capacity is returned + directly so the cold-start phase can quickly collect enough completion samples. Once the cumulative completed request + count reaches ``running_max_req_size`` and the first 16-request QPS window is available, the maximum admitted request + count becomes ``int(QPS * average whole-request seconds) + 6``. If + ``running_max_req_size`` is below 16, the base + capacity remains in effect until QPS initialization, preventing an uninitialized zero QPS from reducing the limit + prematurely. During steady state, six extra requests are admitted as probing headroom, preventing low traffic or a + long idle period from trapping the service at very low concurrency. + This duration represents the desired + average time range from node admission until the whole request finishes and converts completion QPS into a + reasonable in-flight request count. Prefill defaults to 20 seconds because it mainly handles input processing and + the first token; Decode defaults to 60 seconds because incremental token generation usually keeps the request active + longer. Set ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` to override both node defaults with one + explicit value. The option is disabled by default and has no effect in ``normal`` or ``pd_master`` mode. .. option:: --config_server_host diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 0870dd1b21..e43d5b6004 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -72,8 +72,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--enable_pd_node_self_request_limit", action="store_true", help=( - "Allow Prefill and Decode nodes in PD mode to reject a request when no local shm_req object " - "can be allocated within 20 seconds. Default: disabled." + "Allow Prefill and Decode nodes in PD mode to limit the total running request count " + "according to the dynamically measured QPS. Default: disabled." ), ) parser.add_argument( diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 2e74f78604..40f9eeab86 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -34,9 +34,10 @@ from lightllm.server.metrics.manager import MetricClient from .rl_controller import HttpRlController from .manager_ext import HttpRlManagerHelper +from .qps_recorder import QPSRecorder from lightllm.utils.statics_utils import MovingAverage from lightllm.utils.config_utils import get_vocab_size -from lightllm.utils.envs_utils import get_pd_node_shm_req_alloc_timeout_seconds, get_unique_server_name +from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken, ServerBusyError from rpyc.utils.classic import obtain @@ -122,11 +123,10 @@ def __init__( 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 ) - # 超时时间由环境变量 LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS 控制,默认为 20 秒。 - self.pd_node_shm_req_alloc_timeout_seconds = get_pd_node_shm_req_alloc_timeout_seconds() self.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() + self.qps_recorder = QPSRecorder(self.args) # 有的模型的vocab size 读取tokenizer和config.json中不一致 self.vocab_size = max(get_vocab_size(args.model_dir), self.tokenizer.vocab_size) @@ -429,16 +429,16 @@ async def generate( # # 这样会缩小 Prefill 节点自身健康检查的覆盖范围:prompt encode、资源上报及 Decode # 资源等待阶段不再计入本地推理健康状态。资源分配异常应由 PD master 侧的运行请求计数、 - # Decode 节点健康检查和等待资源的超时逻辑负责监控,不能依赖 Prefill 推理计数判断。 + # Decode 节点健康检查和本地准入控制负责监控,不能依赖 Prefill 推理计数判断。 await self._register_running_request() running_request_registered = True - # 申请资源并存储。PD 分段续跑请求可以绕过本地 shm_req 分配超时,避免 - # 已经成功完成首段的用户请求因为下一段暂时拿不到对象而被 429 中断。 + # 申请资源并存储。PD 分段续跑请求可以绕过本地进入并发限制,避免 + # 已经成功完成首段的用户请求因为下一段暂时无法进入等待流程而被 429 中断。 if self.pd_node_request_limit_enabled and sampling_params.bypass_pd_node_request_limit: logger.info( - f"PD {self.args.run_mode} node request {group_request_id} bypasses the local shm_req " - f"allocation timeout and will wait for {sampling_params.n} object(s)" + f"PD {self.args.run_mode} node request {group_request_id} bypasses the local request " + f"concurrency limit and will wait for {sampling_params.n} shm_req object(s)" ) alloced_req_indexes = await self._alloc_shm_req_indexes( sampling_params.n, @@ -556,28 +556,28 @@ async def _alloc_shm_req_indexes( ) -> List[int]: """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。""" alloced_req_indexes = [] - alloc_deadline = ( - time.monotonic() + self.pd_node_shm_req_alloc_timeout_seconds - if self.pd_node_request_limit_enabled and not bypass_pd_node_request_limit - else None - ) try: + if self.pd_node_request_limit_enabled and not bypass_pd_node_request_limit: + current_request_count = self.run_reqs_count_mark.get_value() + # QPS 记录器累计完成的请求数未达到 running_max_req_size 时返回基础容量, + # 使服务冷启动后可以快速积累足够样本;达到后再根据完成 QPS 和节点平均 + # 整包处理时间估算准入上限,并额外保留 6 个请求的探测余量,避免系统在 + # 低 QPS 状态下恢复过慢。Prefill 默认按 20 秒估算,Decode 默认按 + # 60 秒估算,环境变量可以覆盖对应节点的默认时间。 + max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count() + if current_request_count > max_allowed_request_count: + logger.warning( + f"PD {self.args.run_mode} node rejects a request before shm_req allocation: " + f"running_request_count={current_request_count}, " + f"max_allowed_request_count={max_allowed_request_count}" + ) + raise ServerBusyError(f"PD {self.args.run_mode} node is busy") + while len(alloced_req_indexes) < req_num: alloc_req_index = await self.shm_req_manager.async_alloc_req_index() sleep_time = 0.1 while alloc_req_index is None: - if alloc_deadline is not None: - remaining_time = alloc_deadline - time.monotonic() - if remaining_time <= 0: - logger.warning( - f"PD {self.args.run_mode} node shm_req allocation timed out after " - f"{self.pd_node_shm_req_alloc_timeout_seconds} seconds" - ) - raise ServerBusyError( - f"PD {self.args.run_mode} node is busy: unable to allocate a shm_req object " - f"within {self.pd_node_shm_req_alloc_timeout_seconds} seconds" - ) await asyncio.sleep(sleep_time) sleep_time = min(1, sleep_time * 1.1) alloc_req_index = await self.shm_req_manager.async_alloc_req_index() @@ -823,6 +823,7 @@ async def _wait_to_token_package( unfinished_count -= 1 if unfinished_count == 0: + self.qps_recorder.mark_one_req_finish() total_cost_time_ms = (time.time() - start_time) * 1000 mean_per_token_cost_time_ms = (total_cost_time_ms - first_token_cost_ms) / out_token_counter self.per_token_costs.add(mean_per_token_cost_time_ms) diff --git a/lightllm/server/httpserver/qps_recorder.py b/lightllm/server/httpserver/qps_recorder.py new file mode 100644 index 0000000000..12b7b389c8 --- /dev/null +++ b/lightllm/server/httpserver/qps_recorder.py @@ -0,0 +1,79 @@ +import time +from collections import deque +from threading import Lock +from typing import Deque, Optional + +from lightllm.utils.envs_utils import ( + get_pd_request_limit_max_allowed_request_count_seconds, +) + + +class QPSRecorder: + """根据最近完成的请求计算系统动态 QPS。""" + + def __init__(self, args, ema_alpha: float = 0.1): + if not 0 < ema_alpha <= 1: + raise ValueError("ema_alpha must be in the range (0, 1]") + + self.args = args + self.ema_alpha = float(ema_alpha) + # 保存最近 16 个请求的完成时间。16 个时间点之间包含 15 个完成间隔。 + self._finished_timestamps: Deque[float] = deque(maxlen=16) + # 记录服务启动后已经完成的请求总数,用于判断冷启动阶段是否已收集足够样本。 + self._finished_request_count = 0 + self._qps = 0.0 + self._initialized = False + self._last_qps_update_time: Optional[float] = None + self._lock = Lock() + + def mark_one_req_finish(self) -> None: + """记录一个请求完成事件,并在样本充足时更新全局 QPS。""" + finished_time = time.monotonic() + with self._lock: + self._finished_timestamps.append(finished_time) + self._finished_request_count += 1 + self._update_qps() + + def get_qps(self) -> float: + """返回经过 EMA 平滑后的全局 QPS。""" + with self._lock: + if self._last_qps_update_time is not None: + current_time = time.monotonic() + if current_time - self._last_qps_update_time > 30: + self._update_qps() + return self._qps + + def get_max_allowed_request_count(self) -> int: + """根据冷启动样本数和动态 QPS 返回最大允许进入请求数。 + + 服务启动后累计完成的请求数尚未达到 ``running_max_req_size`` 时,直接返回 + 节点的基础运行容量,使冷启动阶段能够快速接收请求并积累足够的 QPS 样本。 + 当 ``running_max_req_size`` 小于 16 时,还需要等待首个完整 QPS 窗口生成, + 避免在 QPS 尚未初始化时过早切换到仅 6 个探测请求。满足两个条件后,才根据 + 完成 QPS 和平均整包时间估算允许进入的请求数,并额外放行 6 个请求作为 + 探测余量,避免系统在低 QPS 状态下恢复过慢。 + """ + with self._lock: + finished_request_count = self._finished_request_count + qps_initialized = self._initialized + if finished_request_count < self.args.running_max_req_size or not qps_initialized: + return self.args.running_max_req_size + + return int(self.get_qps() * get_pd_request_limit_max_allowed_request_count_seconds(self.args.run_mode)) + 6 + + def _update_qps(self) -> None: + if len(self._finished_timestamps) < self._finished_timestamps.maxlen: + return + + current_time = time.monotonic() + elapsed_time = current_time - self._finished_timestamps[0] + if elapsed_time <= 0: + return + + average_qps = (len(self._finished_timestamps) - 1) / elapsed_time + if not self._initialized: + self._qps = average_qps + self._initialized = True + else: + self._qps = self.ema_alpha * average_qps + (1 - self.ema_alpha) * self._qps + self._last_qps_update_time = current_time diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 230df3bcee..09b656655f 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -312,9 +312,27 @@ def get_pd_split_max_new_tokens() -> int: @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)) +def get_pd_request_limit_max_allowed_request_count_seconds(run_mode: str) -> int: + """获取 PD 节点根据 QPS 估算最大在途请求数时使用的平均整包时间。 + + 该值与完成 QPS 相乘,用于估算节点在目标平均整包时间内可以容纳的在途请求数。 + Prefill 节点主要处理输入和首 token,整包占用时间通常较短,默认使用 20 秒; + Decode 节点负责持续生成 token,整包占用时间通常更长,默认使用 60 秒。 + 显式设置环境变量时,Prefill 和 Decode 节点都使用用户提供的值。 + """ + configured_seconds = os.getenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS") + if configured_seconds is None: + default_seconds = {"prefill": 20, "decode": 60} + if run_mode not in default_seconds: + raise ValueError(f"unsupported PD run mode for request limiting: {run_mode}") + seconds = default_seconds[run_mode] + else: + seconds = int(configured_seconds) + if seconds < 0: + raise ValueError( + "LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS must be greater than or equal to 0" + ) + return seconds @lru_cache(maxsize=None) diff --git a/test/test_api/test_qps_recorder.py b/test/test_api/test_qps_recorder.py new file mode 100644 index 0000000000..9b8065680c --- /dev/null +++ b/test/test_api/test_qps_recorder.py @@ -0,0 +1,118 @@ +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from lightllm.server.httpserver.qps_recorder import QPSRecorder +from lightllm.utils.envs_utils import ( + get_pd_request_limit_max_allowed_request_count_seconds, +) + + +def _args(run_mode="decode", running_max_req_size=16): + return SimpleNamespace(run_mode=run_mode, running_max_req_size=running_max_req_size) + + +def test_qps_recorder_waits_for_sixteen_finished_requests(): + recorder = QPSRecorder(_args()) + + with patch("lightllm.server.httpserver.qps_recorder.time.monotonic", side_effect=range(15)): + for _ in range(15): + recorder.mark_one_req_finish() + + assert recorder.get_qps() == 0.0 + + +def test_qps_recorder_calculates_qps_and_updates_ema(): + recorder = QPSRecorder(_args(), ema_alpha=0.25) + + with patch( + "lightllm.server.httpserver.qps_recorder.time.monotonic", + side_effect=[*range(16), 15, 15, 15.5, 15.5, 15.5], + ): + for _ in range(16): + recorder.mark_one_req_finish() + assert recorder.get_qps() == 1.0 + + recorder.mark_one_req_finish() + window_qps = 15 / 14.5 + expected_qps = 0.25 * window_qps + 0.75 * 1.0 + assert recorder.get_qps() == pytest.approx(expected_qps) + + +def test_get_qps_updates_ema_after_thirty_seconds_without_new_request(): + recorder = QPSRecorder(_args(), ema_alpha=0.25) + + with patch( + "lightllm.server.httpserver.qps_recorder.time.monotonic", + side_effect=[*range(16), 15, 45.1, 45.1, 46], + ): + for _ in range(16): + recorder.mark_one_req_finish() + + stale_window_qps = 15 / 45.1 + expected_qps = 0.25 * stale_window_qps + 0.75 * 1.0 + assert recorder.get_qps() == pytest.approx(expected_qps) + assert recorder.get_qps() == pytest.approx(expected_qps) + + +def test_max_allowed_request_count_uses_env(monkeypatch): + monkeypatch.setenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", "12") + get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() + recorder = QPSRecorder(_args(running_max_req_size=1)) + with patch("lightllm.server.httpserver.qps_recorder.time.monotonic", side_effect=range(17)): + for _ in range(16): + recorder.mark_one_req_finish() + + try: + with patch.object(recorder, "get_qps", return_value=2.5): + assert recorder.get_max_allowed_request_count() == 36 + finally: + get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() + + +def test_max_allowed_request_count_uses_running_capacity_during_warmup(): + recorder = QPSRecorder(_args(run_mode="prefill", running_max_req_size=32)) + + with patch.object(recorder, "get_qps") as get_qps: + assert recorder.get_max_allowed_request_count() == 32 + get_qps.assert_not_called() + + +def test_max_allowed_request_count_waits_until_qps_is_initialized(): + recorder = QPSRecorder(_args(run_mode="prefill", running_max_req_size=1)) + recorder.mark_one_req_finish() + + with patch.object(recorder, "get_qps") as get_qps: + assert recorder.get_max_allowed_request_count() == 1 + get_qps.assert_not_called() + + +def test_max_allowed_request_count_keeps_six_probe_requests_at_zero_qps(): + recorder = QPSRecorder(_args(run_mode="decode", running_max_req_size=1)) + with patch( + "lightllm.server.httpserver.qps_recorder.time.monotonic", + side_effect=range(17), + ): + for _ in range(16): + recorder.mark_one_req_finish() + + with patch.object(recorder, "get_qps", return_value=0.0): + assert recorder.get_max_allowed_request_count() == 6 + + +@pytest.mark.parametrize(("run_mode", "expected_seconds"), [("prefill", 20), ("decode", 60)]) +def test_pd_request_limit_max_allowed_request_count_seconds_uses_node_default(monkeypatch, run_mode, expected_seconds): + monkeypatch.delenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", raising=False) + get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() + + try: + assert get_pd_request_limit_max_allowed_request_count_seconds(run_mode) == expected_seconds + finally: + get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() + + +@pytest.mark.parametrize("ema_alpha", [0, -0.1, 1.1]) +def test_qps_recorder_rejects_invalid_ema_alpha(ema_alpha): + with pytest.raises(ValueError, match=r"ema_alpha must be in the range \(0, 1\]"): + QPSRecorder(_args(), ema_alpha=ema_alpha) diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py index 14d8ca00d1..d8b63786c3 100644 --- a/test/test_pd_selector/test_pd_node_request_limit.py +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -8,33 +8,60 @@ from lightllm.utils.error_utils import ServerBusyError -def _manager(*, timeout_seconds: int = 1) -> HttpServerManager: +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(*, running_request_count: int = 0, max_allowed_request_count: int = 1) -> HttpServerManager: manager = HttpServerManager.__new__(HttpServerManager) - manager.args = SimpleNamespace(run_mode="decode") + manager.args = SimpleNamespace(run_mode="decode", running_max_req_size=64) manager.pd_node_request_limit_enabled = True - manager.pd_node_shm_req_alloc_timeout_seconds = timeout_seconds + manager._run_reqs_count_lock = asyncio.Lock() + manager.run_reqs_count_mark = FakeSharedInt(running_request_count) + manager.latest_success_infer_time_mark = FakeSharedInt() + manager.qps_recorder = MagicMock() + manager.qps_recorder.get_max_allowed_request_count.return_value = max_allowed_request_count manager.shm_req_manager = MagicMock() manager.shm_req_manager.async_release_req_index = AsyncMock() return manager -def test_pd_node_regular_request_times_out_when_shm_req_is_unavailable(): +def test_pd_node_rejects_request_when_running_concurrency_exceeds_limit(): async def run(): - manager = _manager(timeout_seconds=0) - manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=None) + manager = _manager(running_request_count=2, max_allowed_request_count=1) + manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=7) with pytest.raises(ServerBusyError): - await manager._alloc_shm_req_indexes( - 1, - bypass_pd_node_request_limit=False, - ) + await manager._alloc_shm_req_indexes(1) + + manager.shm_req_manager.async_alloc_req_index.assert_not_awaited() + + asyncio.run(run()) + + +def test_pd_node_allows_request_when_running_concurrency_equals_limit(): + async def run(): + manager = _manager(running_request_count=1, max_allowed_request_count=1) + manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=7) + + indexes = await manager._alloc_shm_req_indexes(1) + + assert indexes == [7] + manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with() asyncio.run(run()) -def test_pd_split_continuation_bypasses_shm_req_allocation_timeout(): +def test_pd_split_continuation_bypasses_running_concurrency_limit(): async def run(): - manager = _manager(timeout_seconds=0) + manager = _manager(running_request_count=2, max_allowed_request_count=1) manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 7]) with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()): @@ -44,6 +71,34 @@ async def run(): ) assert indexes == [7] + manager.qps_recorder.get_max_allowed_request_count.assert_not_called() + + asyncio.run(run()) + + +def test_admitted_request_waits_for_shm_req_without_timeout(): + async def run(): + manager = _manager() + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, None, 7]) + + with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()) as sleep: + indexes = await manager._alloc_shm_req_indexes(1) + + assert indexes == [7] + assert sleep.await_count == 2 manager.shm_req_manager.async_release_req_index.assert_not_awaited() asyncio.run(run()) + + +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()) From 7addb308e563a02a53ea66927d3b28e814ffe709 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 1 Sep 2026 13:33:13 +0000 Subject: [PATCH 07/20] fix --- lightllm/server/core/objs/req.py | 3 +++ lightllm/server/httpserver/manager.py | 6 +++++- unit_tests/server/core/objs/test_req.py | 10 ++++++++++ 3 files changed, 18 insertions(+), 1 deletion(-) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 87f54fd9e7..5dc3368546 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" diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 40f9eeab86..2b32d3fa42 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -775,6 +775,7 @@ async def _wait_to_token_package( first_token_cost_ms = sys.float_info.max prompt_tokens = len(prompt_ids) is_first_token = True + all_finished_normally = True while True: try: @@ -821,9 +822,12 @@ async def _wait_to_token_package( # 如果有子请求完成,就更新计数 if finish_status.is_finished(): unfinished_count -= 1 + if finish_status.is_error_finished(): + all_finished_normally = False if unfinished_count == 0: - self.qps_recorder.mark_one_req_finish() + if all_finished_normally: + self.qps_recorder.mark_one_req_finish() total_cost_time_ms = (time.time() - start_time) * 1000 mean_per_token_cost_time_ms = (total_cost_time_ms - first_token_cost_ms) / out_token_counter self.per_token_costs.add(mean_per_token_cost_time_ms) 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" From 5ce005aab240707cddecec051febd38866c44842 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 02:01:10 +0000 Subject: [PATCH 08/20] feat(pd): prioritize segmented requests in router queue --- lightllm/server/core/objs/sampling_params.py | 8 +++---- lightllm/server/httpserver/manager.py | 14 +++++------ .../httpserver_for_pd_master/manager.py | 6 ++--- .../model_infer/mode_backend/base_backend.py | 9 ++++++++ .../server/router/req_queue/base_queue.py | 7 +++++- .../server/router/req_queue/dp_base_queue.py | 5 +++- .../test_pd_master_multi_choice.py | 6 ++--- .../test_pd_node_request_limit.py | 20 ++++++++++++++-- .../test_running_request_lifecycle.py | 23 +++++++++++-------- 9 files changed, 67 insertions(+), 31 deletions(-) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 9a13f44842..8e01a5f174 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -294,8 +294,8 @@ 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 为分段续跑请求设置,避免已经完成首段的请求因本地 shm_req 限流失败。 - ("bypass_pd_node_request_limit", ctypes.c_bool), + # 仅由 PD Master 为分段续跑请求设置,表示请求需以高优先级插入 Router 调度队列。 + ("pd_high_priority_request", ctypes.c_bool), ("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), @@ -340,7 +340,7 @@ def init(self, tokenizer, **kwargs): self.input_penalty = kwargs.get("input_penalty", DEFAULT_INPUT_PENALTY) self.group_request_id = kwargs.get("group_request_id", -1) # 该字段是 PD Master 的内部调度信息,不能由外部请求参数开启。 - self.bypass_pd_node_request_limit = False + self.pd_high_priority_request = False self.suggested_dp_index = kwargs.get("suggested_dp_index", -1) self.skip_special_tokens = kwargs.get("skip_special_tokens", SKIP_SPECIAL_TOKENS) @@ -507,7 +507,7 @@ 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, - "bypass_pd_node_request_limit": self.bypass_pd_node_request_limit, + "pd_high_priority_request": self.pd_high_priority_request, "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/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 2b32d3fa42..2afd843d48 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -433,16 +433,16 @@ async def generate( await self._register_running_request() running_request_registered = True - # 申请资源并存储。PD 分段续跑请求可以绕过本地进入并发限制,避免 - # 已经成功完成首段的用户请求因为下一段暂时无法进入等待流程而被 429 中断。 - if self.pd_node_request_limit_enabled and sampling_params.bypass_pd_node_request_limit: + # 申请资源并存储。PD 分段续跑请求作为高优先级请求,可以绕过本地进入并发限制,避免 + # 已经成功完成首段的用户请求因为下一段暂时无法进入等待流程而被 429 中断;同时它会在 Router 队列中优先调度。 + if self.pd_node_request_limit_enabled and sampling_params.pd_high_priority_request: logger.info( - f"PD {self.args.run_mode} node request {group_request_id} bypasses the local request " + f"PD {self.args.run_mode} high-priority request {group_request_id} bypasses the local request " f"concurrency limit and will wait for {sampling_params.n} shm_req object(s)" ) alloced_req_indexes = await self._alloc_shm_req_indexes( sampling_params.n, - bypass_pd_node_request_limit=sampling_params.bypass_pd_node_request_limit, + pd_high_priority_request=sampling_params.pd_high_priority_request, ) req_objs: List[Req] = [] for i, req_index in enumerate(alloced_req_indexes): @@ -552,13 +552,13 @@ def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple async def _alloc_shm_req_indexes( self, req_num: int, - bypass_pd_node_request_limit: bool = False, + pd_high_priority_request: bool = False, ) -> List[int]: """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。""" alloced_req_indexes = [] try: - if self.pd_node_request_limit_enabled and not bypass_pd_node_request_limit: + if self.pd_node_request_limit_enabled and not pd_high_priority_request: current_request_count = self.run_reqs_count_mark.get_value() # QPS 记录器累计完成的请求数未达到 running_max_req_size 时返回基础容量, # 使服务冷启动后可以快速积累足够样本;达到后再根据完成 QPS 和节点平均 diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 89505be014..0e8a5efbd7 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -229,9 +229,9 @@ async def _generate_one( logger.info(f"pd log gen sub req id {block_group_request_id} for main req id {origin_request_id}") sampling_params.max_new_tokens = block_max_new_tokens # 首段仍然遵守 P/D 节点的本地限流。第二段及后续分段说明该用户请求已经 - # 成功执行过一段,为避免续跑请求因暂时申请不到 shm_req 而中途失败, - # 允许它们等待可用对象,不应用本地 shm_req 分配超时。 - sampling_params.bypass_pd_node_request_limit = iter_index > 0 + # 成功执行过一段,将其标记为 PD 高优先级请求,便于优先进入 Router 调度队列。 + # 高优先级请求仍可等待可用 shm_req 对象,避免因临时资源紧张导致分段续跑失败。 + sampling_params.pd_high_priority_request = iter_index > 0 # 分段请求始终复用循环外选定的 P 节点;这里只按每段实际发送的 # prompt 更新该节点的在途 prefill 负载,不会重新选点。 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/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 9af7afd1b4..af39a869b6 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -83,7 +83,12 @@ 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: + # req_group 可能包含同一请求组的多个 Req,整体前置可以保持组内顺序。 + self.waiting_req_list = req_group + self.waiting_req_list + 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..4986c62c65 100644 --- a/lightllm/server/router/req_queue/dp_base_queue.py +++ b/lightllm/server/router/req_queue/dp_base_queue.py @@ -60,7 +60,10 @@ 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: + self.reqs_waiting_for_dp_index.insert(0, 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/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index 088bdb74cc..250b8ca13a 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -262,14 +262,14 @@ async def run(): dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] - bypass_request_limit_flags = [] + 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) - bypass_request_limit_flags.append(sampling_params.bypass_pd_node_request_limit) + high_priority_request_flags.append(sampling_params.pd_high_priority_request) yield ( sampling_params.group_request_id, "x", @@ -295,7 +295,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 bypass_request_limit_flags == [False, True] + 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 diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py index d8b63786c3..e18726edca 100644 --- a/test/test_pd_selector/test_pd_node_request_limit.py +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -5,6 +5,7 @@ import pytest from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.server.router.req_queue.base_queue import BaseQueue from lightllm.utils.error_utils import ServerBusyError @@ -59,7 +60,7 @@ async def run(): asyncio.run(run()) -def test_pd_split_continuation_bypasses_running_concurrency_limit(): +def test_pd_high_priority_split_continuation_bypasses_running_concurrency_limit(): async def run(): manager = _manager(running_request_count=2, max_allowed_request_count=1) manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 7]) @@ -67,7 +68,7 @@ async def run(): with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()): indexes = await manager._alloc_shm_req_indexes( 1, - bypass_pd_node_request_limit=True, + pd_high_priority_request=True, ) assert indexes == [7] @@ -102,3 +103,18 @@ async def run(): manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(3) asyncio.run(run()) + + +def test_pd_high_priority_request_is_inserted_at_router_queue_head(): + queue = BaseQueue.__new__(BaseQueue) + queue.dp_index = 0 + 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 = [normal_req] + + queue.extend([high_req_1, high_req_2]) + + assert queue.waiting_req_list == [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 diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 5d99f7215a..45befb2090 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -46,6 +46,10 @@ 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.qps_recorder = MagicMock() + manager.qps_recorder.get_max_allowed_request_count.return_value = 2 manager.shm_req_manager = SimpleNamespace(async_alloc_req_index=AsyncMock(side_effect=RuntimeError("alloc failed"))) return manager @@ -195,20 +199,20 @@ async def run(): asyncio.run(run()) -def test_pd_node_self_request_limit_rejects_after_shm_req_allocation_timeout(): +def test_pd_node_self_request_limit_rejects_before_shm_req_allocation_when_concurrency_exceeds_limit(): async def run(): manager = _make_manager(NodeRole.D) manager.pd_node_request_limit_enabled = True + manager.run_reqs_count_mark.set_value(3) manager.shm_req_manager = SimpleNamespace( - async_alloc_req_index=AsyncMock(return_value=None), + async_alloc_req_index=AsyncMock(return_value=7), async_release_req_index=AsyncMock(), ) - with patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 120]): - with pytest.raises(ServerBusyError, match="unable to allocate a shm_req object within 20 seconds"): - await _drain_generate(manager, _sampling_params(), _multimodal_params()) + with pytest.raises(ServerBusyError, match="PD decode node is busy"): + await _drain_generate(manager, _sampling_params(), _multimodal_params()) - manager.shm_req_manager.async_alloc_req_index.assert_awaited_once() + manager.shm_req_manager.async_alloc_req_index.assert_not_awaited() 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() @@ -221,15 +225,14 @@ async def run(): manager = _make_manager(NodeRole.D) manager.pd_node_request_limit_enabled = True manager.shm_req_manager = SimpleNamespace( - async_alloc_req_index=AsyncMock(side_effect=[7, None]), + 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 patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[100, 120]): - with pytest.raises(ServerBusyError, match="unable to allocate a shm_req object within 20 seconds"): - await _drain_generate(manager, sampling_params, _multimodal_params()) + 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) From be17e10f37617a7dd4fb49219e2db529381f4928 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 02:57:32 +0000 Subject: [PATCH 09/20] feat(pd): centralize qps admission on pd master --- docs/CN/source/tutorial/api_server_args.rst | 27 +++---- docs/EN/source/tutorial/api_server_args.rst | 30 +++----- lightllm/server/api_cli.py | 4 +- lightllm/server/httpserver/manager.py | 70 ++++++++---------- .../httpserver_for_pd_master/manager.py | 15 ++++ .../qps_recorder.py | 19 ++--- lightllm/utils/envs_utils.py | 25 +++---- test/test_api/test_qps_recorder.py | 17 +++-- .../test_pd_master_multi_choice.py | 2 + .../test_pd_node_request_limit.py | 71 +++++-------------- .../test_pd_master_cached_tokens.py | 2 + .../test_running_request_lifecycle.py | 15 ++-- unit_tests/server/test_pd_master_mode.py | 22 ++++++ 13 files changed, 142 insertions(+), 177 deletions(-) rename lightllm/server/{httpserver => httpserver_for_pd_master}/qps_recorder.py (73%) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 308b0b9d4e..70a7dd3f76 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -89,23 +89,16 @@ PD 分离模式参数 .. option:: --enable_pd_node_self_request_limit - 在 PD 分离模式的 Prefill 和 Decode 节点上启用本地请求准入控制。启用后,节点根据动态 QPS 和 - 当前运行请求总数决定新请求是否可以进入本地 ``shm_req`` 申请流程。超过上限时,节点会主动拒绝新请求,并由 PD Master - 向客户端返回 HTTP 429;已经准入的请求会持续等待可用的 ``shm_req`` 对象,不再进行等待超时判断。 - 服务启动后,QPS 记录器累计完成的请求数未达到 ``running_max_req_size`` 时,最大允许进入请求数直接使用 - ``running_max_req_size``,使冷启动阶段可以快速接收请求并积累足够的完成样本。累计完成请求数达到 - ``running_max_req_size`` 且首个 16 请求 QPS 窗口生成后,最大允许进入请求数改为 - ``int(QPS * 平均整包时间秒数) + 6``。如果 ``running_max_req_size`` 小于 16,节点会继续使用基础容量, - 直到 QPS 完成初始化,避免使用尚未初始化的零值 QPS 过早切换到仅 6 个探测请求。 - 稳定运行阶段额外放行 6 个请求作为探测余量,避免低流量或长时间空闲导致 QPS 降低后, - 系统被限制在过低的并发并且难以重新爬升。 - 这里的平均整包时间表示希望一个请求从进入 - 节点到整包完成所处的时间范围,用于将完成 QPS 换算为合理的在途请求数量。Prefill 阶段通常只负责输入 - 处理和首 token,默认值为 20 秒,避免短阶段积压过多请求;Decode 阶段需要持续生成 token,默认值为 - 60 秒,以覆盖更长的整包处理时间。可以通过环境变量 - ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 显式设置统一值;设置后 Prefill 和 Decode - 节点都会采用该值。 - 该参数默认关闭,在 ``normal`` 和 ``pd_master`` 模式下不生效。 + 在 PD Master 上启用基于动态 QPS 的请求准入探测。PD Master 根据最近完成请求计算 QPS, + 并限制同时进入的完整 PD 请求数;超过上限时直接向客户端返回 HTTP 429。服务启动后, + 已完成请求数未达到 ``running_max_req_size`` 或 QPS 窗口尚未初始化时,暂时使用 + ``running_max_req_size``,以便快速积累样本;之后使用 ``int(QPS * 平均整包时间秒数) + 6``, + 额外保留 6 个请求作为探测余量,避免低流量后无法恢复。PD Master 统计完整 PD 请求, + ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 未设置时默认使用 60 秒。 + Prefill/Decode 节点不再执行 QPS 准入,但 HTTP server 在申请本地 ``shm_req`` 对象时, + 如果超过 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS``(默认 20 秒)仍未成功, + 会上报 ``Server is busy``,由 PD Master 转换为 HTTP 429;未开启限流以及 PD 高优先级分段续跑请求 + 不受该超时限制,会持续等待资源。该参数默认关闭。 .. 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 4a0bbd2fd0..05167ffec9 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -91,25 +91,17 @@ PD disaggregation Mode Parameters .. option:: --enable_pd_node_self_request_limit - Enable local admission control on Prefill and Decode nodes in PD disaggregation mode. When enabled, - each node uses the dynamically measured QPS and its total running request count to decide whether a new request - may enter local ``shm_req`` allocation. Once the limit is exceeded, the node rejects new requests through PD Master with HTTP 429. Admitted requests - keep waiting for an available ``shm_req`` object without an allocation timeout. While the QPS recorder has completed - fewer than ``running_max_req_size`` requests since startup, that base capacity is returned - directly so the cold-start phase can quickly collect enough completion samples. Once the cumulative completed request - count reaches ``running_max_req_size`` and the first 16-request QPS window is available, the maximum admitted request - count becomes ``int(QPS * average whole-request seconds) + 6``. If - ``running_max_req_size`` is below 16, the base - capacity remains in effect until QPS initialization, preventing an uninitialized zero QPS from reducing the limit - prematurely. During steady state, six extra requests are admitted as probing headroom, preventing low traffic or a - long idle period from trapping the service at very low concurrency. - This duration represents the desired - average time range from node admission until the whole request finishes and converts completion QPS into a - reasonable in-flight request count. Prefill defaults to 20 seconds because it mainly handles input processing and - the first token; Decode defaults to 60 seconds because incremental token generation usually keeps the request active - longer. Set ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` to override both node defaults with one - explicit value. - The option is disabled by default and has no effect in ``normal`` or ``pd_master`` mode. + Enable dynamic-QPS admission probing on the PD Master. The PD Master computes QPS from recently completed + requests and limits the number of complete PD requests admitted concurrently; requests above the limit receive + HTTP 429 directly. During cold start, while completed requests are fewer than ``running_max_req_size`` or the + QPS window is not initialized, the base ``running_max_req_size`` is used to collect samples quickly. Afterwards, + the limit is ``int(QPS * average whole-request seconds) + 6``; six extra requests provide probing headroom so + low traffic does not permanently trap the service at low concurrency. The PD Master measures complete PD + requests, so ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` defaults to 60 seconds when unset. + Prefill/Decode nodes no longer perform QPS admission. Instead, their HTTP server reports ``Server is busy`` if + local ``shm_req`` allocation does not succeed within ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` + (20 seconds by default); PD Master converts this to HTTP 429. The timeout is bypassed when local admission is + disabled and for PD high-priority segmented continuation requests, which continue waiting for resources. Disabled by default. .. option:: --config_server_host diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e43d5b6004..4e90bac0fa 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -72,8 +72,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--enable_pd_node_self_request_limit", action="store_true", help=( - "Allow Prefill and Decode nodes in PD mode to limit the total running request count " - "according to the dynamically measured QPS. Default: disabled." + "Enable PD Master admission probing based on dynamically measured QPS; Prefill/Decode nodes " + "only enforce shm_req allocation timeout. Default: disabled." ), ) parser.add_argument( diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 2afd843d48..06cab06b51 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -34,10 +34,9 @@ from lightllm.server.metrics.manager import MetricClient from .rl_controller import HttpRlController from .manager_ext import HttpRlManagerHelper -from .qps_recorder import QPSRecorder 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_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, ServerBusyError from rpyc.utils.classic import obtain @@ -118,15 +117,15 @@ def __init__( self.pd_mode: NodeRole = NodeRole(self.args.run_mode) assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D] - # 该开关只对 PD 分离模式的 Prefill/Decode 服务主节点生效。 - # 多机 TP 的从节点不直接与 PD Master 通信,不能独立拒绝请求。 + # HTTP server 只负责在本地 shm_req 长时间不可用时快速返回繁忙,PD Master 负责 QPS 准入限流。 + # 该开关控制 P/D 节点是否启用本地 shm_req 等待超时;多机 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.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() - self.qps_recorder = QPSRecorder(self.args) # 有的模型的vocab size 读取tokenizer和config.json中不一致 self.vocab_size = max(get_vocab_size(args.model_dir), self.tokenizer.vocab_size) @@ -429,17 +428,12 @@ async def generate( # # 这样会缩小 Prefill 节点自身健康检查的覆盖范围:prompt encode、资源上报及 Decode # 资源等待阶段不再计入本地推理健康状态。资源分配异常应由 PD master 侧的运行请求计数、 - # Decode 节点健康检查和本地准入控制负责监控,不能依赖 Prefill 推理计数判断。 + # Decode 节点健康检查和本地 shm_req 等待超时负责监控,不能依赖 Prefill 推理计数判断。 await self._register_running_request() running_request_registered = True - # 申请资源并存储。PD 分段续跑请求作为高优先级请求,可以绕过本地进入并发限制,避免 - # 已经成功完成首段的用户请求因为下一段暂时无法进入等待流程而被 429 中断;同时它会在 Router 队列中优先调度。 - if self.pd_node_request_limit_enabled and sampling_params.pd_high_priority_request: - logger.info( - f"PD {self.args.run_mode} high-priority request {group_request_id} bypasses the local request " - f"concurrency limit and will wait for {sampling_params.n} shm_req object(s)" - ) + # 申请资源并存储。PD 分段续跑请求仍在 Router 队列中优先调度,同时绕过本地 + # shm_req 等待超时,避免已经开始执行的请求因临时资源紧张而被中断。 alloced_req_indexes = await self._alloc_shm_req_indexes( sampling_params.n, pd_high_priority_request=sampling_params.pd_high_priority_request, @@ -549,36 +543,35 @@ 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, - ) -> List[int]: - """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。""" + async def _alloc_shm_req_indexes(self, req_num: int, pd_high_priority_request: bool = False) -> List[int]: + """为一个请求申请全部 shm_req 索引,申请失败时回滚已分配的索引。 + + 未开启本地限流或请求为 PD 高优先级请求时无限等待;普通请求在限流开启时, + 最多等待 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` 秒。 + """ alloced_req_indexes = [] + request_limit_applies = self.pd_node_request_limit_enabled and not pd_high_priority_request + alloc_deadline = ( + time.monotonic() + self.pd_node_shm_req_alloc_timeout_seconds if request_limit_applies else None + ) try: - if self.pd_node_request_limit_enabled and not pd_high_priority_request: - current_request_count = self.run_reqs_count_mark.get_value() - # QPS 记录器累计完成的请求数未达到 running_max_req_size 时返回基础容量, - # 使服务冷启动后可以快速积累足够样本;达到后再根据完成 QPS 和节点平均 - # 整包处理时间估算准入上限,并额外保留 6 个请求的探测余量,避免系统在 - # 低 QPS 状态下恢复过慢。Prefill 默认按 20 秒估算,Decode 默认按 - # 60 秒估算,环境变量可以覆盖对应节点的默认时间。 - max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count() - if current_request_count > max_allowed_request_count: - logger.warning( - f"PD {self.args.run_mode} node rejects a request before shm_req allocation: " - f"running_request_count={current_request_count}, " - f"max_allowed_request_count={max_allowed_request_count}" - ) - raise ServerBusyError(f"PD {self.args.run_mode} node is busy") - 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: - await asyncio.sleep(sleep_time) + 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"{self.pd_node_shm_req_alloc_timeout_seconds} seconds" + ) + raise ServerBusyError( + f"PD {self.args.run_mode} node is busy: unable to allocate a shm_req object " + f"within {self.pd_node_shm_req_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) @@ -775,7 +768,6 @@ async def _wait_to_token_package( first_token_cost_ms = sys.float_info.max prompt_tokens = len(prompt_ids) is_first_token = True - all_finished_normally = True while True: try: @@ -822,12 +814,8 @@ async def _wait_to_token_package( # 如果有子请求完成,就更新计数 if finish_status.is_finished(): unfinished_count -= 1 - if finish_status.is_error_finished(): - all_finished_normally = False if unfinished_count == 0: - if all_finished_normally: - self.qps_recorder.mark_one_req_finish() total_cost_time_ms = (time.time() - start_time) * 1000 mean_per_token_cost_time_ms = (total_cost_time_ms - first_token_cost_ms) / out_token_counter self.per_token_costs.add(mean_per_token_cost_time_ms) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 0e8a5efbd7..a06dae1628 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -26,6 +26,7 @@ from lightllm.utils.envs_utils import get_pd_split_max_new_tokens from lightllm.utils.shm_port_args import get_shm_port_args from .pd_selector import create_selector +from .qps_recorder import QPSRecorder logger = init_logger(__name__) @@ -48,6 +49,10 @@ def __init__( self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 + # PD Master 统一统计完整请求的 QPS,并据此控制进入请求数;P/D 节点只负责 + # shm_req 资源申请超时,避免各节点分别探测造成限流判断不一致。 + self.pd_master_request_limit_enabled = args.enable_pd_node_self_request_limit and args.run_mode == "pd_master" + self.qps_recorder = QPSRecorder(args) self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -129,6 +134,15 @@ async def generate( multimodal_params: MultimodalParams, request: Request, ): + if self.pd_master_request_limit_enabled: + max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count() + if self.running_request_count > max_allowed_request_count: + logger.warning( + f"PD Master rejects request before dispatch: running_request_count={self.running_request_count}, " + f"max_allowed_request_count={max_allowed_request_count}" + ) + raise ServerBusyError("PD Master is busy") + was_idle = self.running_request_count == 0 self.running_request_count += 1 if was_idle: @@ -190,6 +204,7 @@ async def _generate( async for result in self._merge_choice_generators(generators): yield result self.metric_client.counter_inc("lightllm_request_success") + self.qps_recorder.mark_one_req_finish() return async def _generate_one( diff --git a/lightllm/server/httpserver/qps_recorder.py b/lightllm/server/httpserver_for_pd_master/qps_recorder.py similarity index 73% rename from lightllm/server/httpserver/qps_recorder.py rename to lightllm/server/httpserver_for_pd_master/qps_recorder.py index 12b7b389c8..b3a478e83f 100644 --- a/lightllm/server/httpserver/qps_recorder.py +++ b/lightllm/server/httpserver_for_pd_master/qps_recorder.py @@ -3,13 +3,11 @@ from threading import Lock from typing import Deque, Optional -from lightllm.utils.envs_utils import ( - get_pd_request_limit_max_allowed_request_count_seconds, -) +from lightllm.utils.envs_utils import get_pd_request_limit_max_allowed_request_count_seconds class QPSRecorder: - """根据最近完成的请求计算系统动态 QPS。""" + """根据最近完成的请求计算 PD Master 的动态 QPS。""" def __init__(self, args, ema_alpha: float = 0.1): if not 0 < ema_alpha <= 1: @@ -44,22 +42,15 @@ def get_qps(self) -> float: return self._qps def get_max_allowed_request_count(self) -> int: - """根据冷启动样本数和动态 QPS 返回最大允许进入请求数。 - - 服务启动后累计完成的请求数尚未达到 ``running_max_req_size`` 时,直接返回 - 节点的基础运行容量,使冷启动阶段能够快速接收请求并积累足够的 QPS 样本。 - 当 ``running_max_req_size`` 小于 16 时,还需要等待首个完整 QPS 窗口生成, - 避免在 QPS 尚未初始化时过早切换到仅 6 个探测请求。满足两个条件后,才根据 - 完成 QPS 和平均整包时间估算允许进入的请求数,并额外放行 6 个请求作为 - 探测余量,避免系统在低 QPS 状态下恢复过慢。 - """ + """根据冷启动样本数和动态 QPS 返回 PD Master 最大允许进入请求数。""" with self._lock: finished_request_count = self._finished_request_count qps_initialized = self._initialized if finished_request_count < self.args.running_max_req_size or not qps_initialized: return self.args.running_max_req_size - return int(self.get_qps() * get_pd_request_limit_max_allowed_request_count_seconds(self.args.run_mode)) + 6 + # PD Master 统计完整 PD 请求,按统一的平均整包时长配置估算在途请求数。 + return int(self.get_qps() * get_pd_request_limit_max_allowed_request_count_seconds()) + 6 def _update_qps(self) -> None: if len(self._finished_timestamps) < self._finished_timestamps.maxlen: diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 09b656655f..9ab1815dc8 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -312,22 +312,19 @@ def get_pd_split_max_new_tokens() -> int: @lru_cache(maxsize=None) -def get_pd_request_limit_max_allowed_request_count_seconds(run_mode: str) -> int: - """获取 PD 节点根据 QPS 估算最大在途请求数时使用的平均整包时间。 +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)) - 该值与完成 QPS 相乘,用于估算节点在目标平均整包时间内可以容纳的在途请求数。 - Prefill 节点主要处理输入和首 token,整包占用时间通常较短,默认使用 20 秒; - Decode 节点负责持续生成 token,整包占用时间通常更长,默认使用 60 秒。 - 显式设置环境变量时,Prefill 和 Decode 节点都使用用户提供的值。 + +@lru_cache(maxsize=None) +def get_pd_request_limit_max_allowed_request_count_seconds() -> int: + """获取根据 QPS 估算 PD Master 最大在途请求数时使用的平均整包时间。 + + 该值与完成 QPS 相乘,用于估算在目标平均整包时间内可以容纳的完整 PD 请求数。 + PD Master 统一统计完整请求,不再区分 Prefill 和 Decode,默认值为 60 秒。 """ - configured_seconds = os.getenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS") - if configured_seconds is None: - default_seconds = {"prefill": 20, "decode": 60} - if run_mode not in default_seconds: - raise ValueError(f"unsupported PD run mode for request limiting: {run_mode}") - seconds = default_seconds[run_mode] - else: - seconds = int(configured_seconds) + seconds = int(os.getenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", 60)) if seconds < 0: raise ValueError( "LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS must be greater than or equal to 0" diff --git a/test/test_api/test_qps_recorder.py b/test/test_api/test_qps_recorder.py index 9b8065680c..b816a5e213 100644 --- a/test/test_api/test_qps_recorder.py +++ b/test/test_api/test_qps_recorder.py @@ -3,7 +3,7 @@ import pytest -from lightllm.server.httpserver.qps_recorder import QPSRecorder +from lightllm.server.httpserver_for_pd_master.qps_recorder import QPSRecorder from lightllm.utils.envs_utils import ( get_pd_request_limit_max_allowed_request_count_seconds, ) @@ -16,7 +16,7 @@ def _args(run_mode="decode", running_max_req_size=16): def test_qps_recorder_waits_for_sixteen_finished_requests(): recorder = QPSRecorder(_args()) - with patch("lightllm.server.httpserver.qps_recorder.time.monotonic", side_effect=range(15)): + with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(15)): for _ in range(15): recorder.mark_one_req_finish() @@ -27,7 +27,7 @@ def test_qps_recorder_calculates_qps_and_updates_ema(): recorder = QPSRecorder(_args(), ema_alpha=0.25) with patch( - "lightllm.server.httpserver.qps_recorder.time.monotonic", + "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=[*range(16), 15, 15, 15.5, 15.5, 15.5], ): for _ in range(16): @@ -44,7 +44,7 @@ def test_get_qps_updates_ema_after_thirty_seconds_without_new_request(): recorder = QPSRecorder(_args(), ema_alpha=0.25) with patch( - "lightllm.server.httpserver.qps_recorder.time.monotonic", + "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=[*range(16), 15, 45.1, 45.1, 46], ): for _ in range(16): @@ -60,7 +60,7 @@ def test_max_allowed_request_count_uses_env(monkeypatch): monkeypatch.setenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", "12") get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() recorder = QPSRecorder(_args(running_max_req_size=1)) - with patch("lightllm.server.httpserver.qps_recorder.time.monotonic", side_effect=range(17)): + with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(17)): for _ in range(16): recorder.mark_one_req_finish() @@ -91,7 +91,7 @@ def test_max_allowed_request_count_waits_until_qps_is_initialized(): def test_max_allowed_request_count_keeps_six_probe_requests_at_zero_qps(): recorder = QPSRecorder(_args(run_mode="decode", running_max_req_size=1)) with patch( - "lightllm.server.httpserver.qps_recorder.time.monotonic", + "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(17), ): for _ in range(16): @@ -101,13 +101,12 @@ def test_max_allowed_request_count_keeps_six_probe_requests_at_zero_qps(): assert recorder.get_max_allowed_request_count() == 6 -@pytest.mark.parametrize(("run_mode", "expected_seconds"), [("prefill", 20), ("decode", 60)]) -def test_pd_request_limit_max_allowed_request_count_seconds_uses_node_default(monkeypatch, run_mode, expected_seconds): +def test_pd_request_limit_max_allowed_request_count_seconds_uses_unified_default(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", raising=False) get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() try: - assert get_pd_request_limit_max_allowed_request_count_seconds(run_mode) == expected_seconds + assert get_pd_request_limit_max_allowed_request_count_seconds() == 60 finally: get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() 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 250b8ca13a..be4b8386f0 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -16,6 +16,8 @@ def _manager() -> HttpServerManagerForPDMaster: manager.id_gen = MagicMock() manager.id_gen.generate_id.return_value = 800 manager.metric_client = MagicMock() + manager.pd_master_request_limit_enabled = False + manager.qps_recorder = MagicMock() manager._log_req_header = AsyncMock() manager.tokens = MagicMock(return_value=2) return manager diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py index e18726edca..091c35cb5b 100644 --- a/test/test_pd_selector/test_pd_node_request_limit.py +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -6,7 +6,6 @@ from lightllm.server.httpserver.manager import HttpServerManager from lightllm.server.router.req_queue.base_queue import BaseQueue -from lightllm.utils.error_utils import ServerBusyError class FakeSharedInt: @@ -20,87 +19,55 @@ def set_value(self, value: int): self.value = value -def _manager(*, running_request_count: int = 0, max_allowed_request_count: int = 1) -> HttpServerManager: +def _manager() -> HttpServerManager: manager = HttpServerManager.__new__(HttpServerManager) manager.args = SimpleNamespace(run_mode="decode", running_max_req_size=64) - manager.pd_node_request_limit_enabled = True + manager.pd_node_request_limit_enabled = False manager._run_reqs_count_lock = asyncio.Lock() - manager.run_reqs_count_mark = FakeSharedInt(running_request_count) + manager.run_reqs_count_mark = FakeSharedInt() manager.latest_success_infer_time_mark = FakeSharedInt() - manager.qps_recorder = MagicMock() - manager.qps_recorder.get_max_allowed_request_count.return_value = max_allowed_request_count + 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_node_rejects_request_when_running_concurrency_exceeds_limit(): - async def run(): - manager = _manager(running_request_count=2, max_allowed_request_count=1) - manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=7) - - with pytest.raises(ServerBusyError): - await manager._alloc_shm_req_indexes(1) - - manager.shm_req_manager.async_alloc_req_index.assert_not_awaited() - - asyncio.run(run()) - - -def test_pd_node_allows_request_when_running_concurrency_equals_limit(): - async def run(): - manager = _manager(running_request_count=1, max_allowed_request_count=1) - manager.shm_req_manager.async_alloc_req_index = AsyncMock(return_value=7) - - indexes = await manager._alloc_shm_req_indexes(1) - - assert indexes == [7] - manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with() - - asyncio.run(run()) - - -def test_pd_high_priority_split_continuation_bypasses_running_concurrency_limit(): +def test_shm_req_partial_allocations_are_released_on_failure(): async def run(): - manager = _manager(running_request_count=2, max_allowed_request_count=1) - manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 7]) + manager = _manager() + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[3, RuntimeError("failed")]) - with patch("lightllm.server.httpserver.manager.asyncio.sleep", new=AsyncMock()): - indexes = await manager._alloc_shm_req_indexes( - 1, - pd_high_priority_request=True, - ) + with pytest.raises(RuntimeError, match="failed"): + await manager._alloc_shm_req_indexes(2) - assert indexes == [7] - manager.qps_recorder.get_max_allowed_request_count.assert_not_called() + manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(3) asyncio.run(run()) -def test_admitted_request_waits_for_shm_req_without_timeout(): +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, 7]) + 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: - indexes = await manager._alloc_shm_req_indexes(1) + assert await manager._alloc_shm_req_indexes(1) == [3] - assert indexes == [7] assert sleep.await_count == 2 - manager.shm_req_manager.async_release_req_index.assert_not_awaited() asyncio.run(run()) -def test_shm_req_partial_allocations_are_released_on_failure(): +def test_high_priority_shm_req_allocation_uses_shorter_backoff_even_with_local_limit(): async def run(): manager = _manager() - manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[3, RuntimeError("failed")]) + manager.pd_node_request_limit_enabled = True + manager.shm_req_manager.async_alloc_req_index = AsyncMock(side_effect=[None, 3]) - with pytest.raises(RuntimeError, match="failed"): - await manager._alloc_shm_req_indexes(2) + 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] - manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(3) + assert sleep.await_args_list[0].args[0] == pytest.approx(0.1 * 0.2) asyncio.run(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..07f0dffa40 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -16,6 +16,8 @@ def _make_manager(monkeypatch): monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) mgr.running_request_count = 0 + mgr.pd_master_request_limit_enabled = False + mgr.qps_recorder = SimpleNamespace(mark_one_req_finish=lambda: None) counter = [0] def gen_id(): diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 45befb2090..12bc9118dc 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -48,8 +48,6 @@ def _make_manager(mode: NodeRole): manager.metric_client = MagicMock() manager._run_reqs_count_lock = asyncio.Lock() manager.run_reqs_count_mark = _ValueMark() - manager.qps_recorder = MagicMock() - manager.qps_recorder.get_max_allowed_request_count.return_value = 2 manager.shm_req_manager = SimpleNamespace(async_alloc_req_index=AsyncMock(side_effect=RuntimeError("alloc failed"))) return manager @@ -199,20 +197,20 @@ async def run(): asyncio.run(run()) -def test_pd_node_self_request_limit_rejects_before_shm_req_allocation_when_concurrency_exceeds_limit(): +def test_httpserver_returns_busy_when_shm_req_allocation_times_out(): async def run(): manager = _make_manager(NodeRole.D) manager.pd_node_request_limit_enabled = True - manager.run_reqs_count_mark.set_value(3) manager.shm_req_manager = SimpleNamespace( - async_alloc_req_index=AsyncMock(return_value=7), + async_alloc_req_index=AsyncMock(return_value=None), async_release_req_index=AsyncMock(), ) - with pytest.raises(ServerBusyError, match="PD decode node is busy"): - await _drain_generate(manager, _sampling_params(), _multimodal_params()) + with patch("lightllm.server.httpserver.manager.time.monotonic", side_effect=[0, 21]): + with pytest.raises(ServerBusyError, match="PD decode node is busy"): + await _drain_generate(manager, _sampling_params(), _multimodal_params()) - manager.shm_req_manager.async_alloc_req_index.assert_not_awaited() + 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() @@ -223,7 +221,6 @@ async def run(): def test_pd_node_self_request_limit_releases_partially_allocated_shm_reqs(): async def run(): manager = _make_manager(NodeRole.D) - manager.pd_node_request_limit_enabled = True manager.shm_req_manager = SimpleNamespace( async_alloc_req_index=AsyncMock(side_effect=[7, RuntimeError("alloc failed")]), async_release_req_index=AsyncMock(), diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 981b250f4e..888bed4013 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,6 +1,7 @@ import asyncio import json from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from easydict import EasyDict @@ -8,6 +9,7 @@ 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 +from lightllm.utils.error_utils import ServerBusyError def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled(): @@ -18,6 +20,24 @@ def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled( assert StartArgs().enable_pd_node_self_request_limit is False +def test_pd_master_qps_limit_rejects_before_dispatch(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_master_request_limit_enabled = True + manager.running_request_count = 3 + manager.qps_recorder = MagicMock() + manager.qps_recorder.get_max_allowed_request_count.return_value = 2 + + async def consume_generate(): + async for _ in manager.generate("prompt", None, None, None): + pass + + with pytest.raises(ServerBusyError, match="PD Master is busy"): + asyncio.run(consume_generate()) + + assert manager.running_request_count == 3 + manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with() + + def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): from lightllm.utils.config_utils import auto_set_response_parsers @@ -278,6 +298,7 @@ async def verify_and_preload(self, request): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.running_request_count = 0 + manager.pd_master_request_limit_enabled = False async def consume_generate(): async for _ in manager.generate("prompt", None, FailingMultimodalParams(), None): @@ -292,6 +313,7 @@ async def consume_generate(): def test_pd_master_request_count_covers_async_generator_lifecycle(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.running_request_count = 0 + manager.pd_master_request_limit_enabled = False inner_generator_closed = False async def fake_generate(prompt, sampling_params, multimodal_params, request): From 9d4fd4a3a7cd743845a5b1c7a73951149dffcd86 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 05:54:51 +0000 Subject: [PATCH 10/20] feat(pd): reject requests stuck in router queue --- docs/CN/source/tutorial/api_server_args.rst | 10 ++-- docs/EN/source/tutorial/api_server_args.rst | 8 ++- lightllm/server/api_cli.py | 2 +- lightllm/server/core/objs/req.py | 6 ++ lightllm/server/httpserver/manager.py | 37 ++++++++++++- lightllm/server/router/manager.py | 7 +++ lightllm/utils/envs_utils.py | 8 ++- .../test_running_request_lifecycle.py | 55 ++++++++++++++++++- unit_tests/utils/test_envs_utils.py | 23 +++++++- 9 files changed, 142 insertions(+), 14 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 70a7dd3f76..c094fd4b3f 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -95,10 +95,12 @@ PD 分离模式参数 ``running_max_req_size``,以便快速积累样本;之后使用 ``int(QPS * 平均整包时间秒数) + 6``, 额外保留 6 个请求作为探测余量,避免低流量后无法恢复。PD Master 统计完整 PD 请求, ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 未设置时默认使用 60 秒。 - Prefill/Decode 节点不再执行 QPS 准入,但 HTTP server 在申请本地 ``shm_req`` 对象时, - 如果超过 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS``(默认 20 秒)仍未成功, - 会上报 ``Server is busy``,由 PD Master 转换为 HTTP 429;未开启限流以及 PD 高优先级分段续跑请求 - 不受该超时限制,会持续等待资源。该参数默认关闭。 + Prefill/Decode 节点不再执行 QPS 准入。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 高优先级分段续跑请求不受该超时限制,会持续等待资源。该参数默认关闭。 .. 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 05167ffec9..0c9dceae15 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -98,9 +98,11 @@ PD disaggregation Mode Parameters the limit is ``int(QPS * average whole-request seconds) + 6``; six extra requests provide probing headroom so low traffic does not permanently trap the service at low concurrency. The PD Master measures complete PD requests, so ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` defaults to 60 seconds when unset. - Prefill/Decode nodes no longer perform QPS admission. Instead, their HTTP server reports ``Server is busy`` if - local ``shm_req`` allocation does not succeed within ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` - (20 seconds by default); PD Master converts this to HTTP 429. The timeout is bypassed when local admission is + Prefill/Decode nodes no longer perform QPS admission. 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. The timeout is bypassed when local admission is disabled and for PD high-priority segmented continuation requests, which continue waiting for resources. Disabled by default. .. option:: --config_server_host diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 4e90bac0fa..73d0405295 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -73,7 +73,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help=( "Enable PD Master admission probing based on dynamically measured QPS; Prefill/Decode nodes " - "only enforce shm_req allocation timeout. Default: disabled." + "enforce both shm_req allocation timeout and Router scheduling wait timeout. Default: disabled." ), ) parser.add_argument( diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 5dc3368546..670fd5d252 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -86,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), @@ -160,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/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 06cab06b51..27ff671e2a 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -36,7 +36,11 @@ 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_pd_node_shm_req_alloc_timeout_seconds, 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, ServerBusyError from rpyc.utils.classic import obtain @@ -117,12 +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 长时间不可用时快速返回繁忙,PD Master 负责 QPS 准入限流。 - # 该开关控制 P/D 节点是否启用本地 shm_req 等待超时;多机 TP 从节点不独立拒绝请求。 + # 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() @@ -775,6 +780,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( @@ -1082,6 +1097,22 @@ 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 短暂拥塞被清理。 + if any(req.sample_params.pd_high_priority_request for req in reqs): + return False + + 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/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/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 9ab1815dc8..83848c69c3 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -313,10 +313,16 @@ def get_pd_split_max_new_tokens() -> int: @lru_cache(maxsize=None) def get_pd_node_shm_req_alloc_timeout_seconds() -> int: - """PD 节点申请 shm_req 对象的最长等待时间,单位为秒。""" + """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_request_limit_max_allowed_request_count_seconds() -> int: """获取根据 QPS 估算 PD Master 最大在途请求数时使用的平均整包时间。 diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 12bc9118dc..08ab652af6 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -6,7 +6,7 @@ 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, ServerBusyError @@ -33,6 +33,7 @@ def _make_manager(mode: NodeRole): 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 @@ -218,6 +219,58 @@ async def run(): asyncio.run(run()) +def test_httpserver_returns_busy_while_first_token_request_waits_in_router(): + async def run(): + manager = _make_manager(NodeRole.D) + 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 anext(output_generator) + + 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), + ) + 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) diff --git a/unit_tests/utils/test_envs_utils.py b/unit_tests/utils/test_envs_utils.py index 7da944314a..a806f8c6de 100644 --- a/unit_tests/utils/test_envs_utils.py +++ b/unit_tests/utils/test_envs_utils.py @@ -1,4 +1,7 @@ -from lightllm.utils.envs_utils import get_pd_node_shm_req_alloc_timeout_seconds +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): @@ -17,3 +20,21 @@ def test_pd_node_shm_req_alloc_timeout_reads_environment_variable(monkeypatch): 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() From 6737d49f7441a8fc531e0dd26d3a46dece76c06f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 06:00:27 +0000 Subject: [PATCH 11/20] fix(pd): exclude failed requests from qps stats --- .../httpserver_for_pd_master/manager.py | 9 ++++- .../test_pd_master_multi_choice.py | 40 +++++++++++++++++++ 2 files changed, 47 insertions(+), 2 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index a06dae1628..f074d689a4 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -201,10 +201,15 @@ 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") - self.qps_recorder.mark_one_req_finish() + if request_finished_successfully: + self.metric_client.counter_inc("lightllm_request_success") + self.qps_recorder.mark_one_req_finish() return async def _generate_one( 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 be4b8386f0..8b75c2e7cf 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -96,6 +96,7 @@ async def wait_to_token_package( call("lightllm_request_count"), call("lightllm_request_success"), ] + manager.qps_recorder.mark_one_req_finish.assert_called_once_with() assert p_node.dispatched_prompt_chars == 0 assert p_node.dispatched_req_num == 0 @@ -145,6 +146,45 @@ 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") + manager.qps_recorder.mark_one_req_finish.assert_not_called() + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_multi_choice_failure_closes_other_generators(): async def run(): manager = _manager() From 9eb1f0a1479870b3a4a0d6256fb08ef1aa92ce5c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 06:12:29 +0000 Subject: [PATCH 12/20] refactor(pd): warm up qps limit from decode capacity --- docs/CN/source/tutorial/api_server_args.rst | 5 +-- docs/EN/source/tutorial/api_server_args.rst | 5 +-- .../httpserver_for_pd_master/manager.py | 4 ++- .../httpserver_for_pd_master/qps_recorder.py | 21 ++++++------ test/test_api/test_qps_recorder.py | 34 +++++++++---------- unit_tests/server/test_pd_master_mode.py | 8 ++++- 6 files changed, 43 insertions(+), 34 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index c094fd4b3f..47b7f0075d 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -91,8 +91,9 @@ PD 分离模式参数 在 PD Master 上启用基于动态 QPS 的请求准入探测。PD Master 根据最近完成请求计算 QPS, 并限制同时进入的完整 PD 请求数;超过上限时直接向客户端返回 HTTP 429。服务启动后, - 已完成请求数未达到 ``running_max_req_size`` 或 QPS 窗口尚未初始化时,暂时使用 - ``running_max_req_size``,以便快速积累样本;之后使用 ``int(QPS * 平均整包时间秒数) + 6``, + 最近完成时间窗口尚未积累满 64 个请求,或 QPS 尚未初始化时,暂时使用所有 Decode 节点 + ``running_max_req_size`` 之和作为并发上限,以便快速积累样本;之后使用 + ``int(QPS * 平均整包时间秒数) + 6``, 额外保留 6 个请求作为探测余量,避免低流量后无法恢复。PD Master 统计完整 PD 请求, ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 未设置时默认使用 60 秒。 Prefill/Decode 节点不再执行 QPS 准入。HTTP server 申请本地 ``shm_req`` 对象的超时时间由 diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 0c9dceae15..2ad2f5e9d1 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -93,8 +93,9 @@ PD disaggregation Mode Parameters Enable dynamic-QPS admission probing on the PD Master. The PD Master computes QPS from recently completed requests and limits the number of complete PD requests admitted concurrently; requests above the limit receive - HTTP 429 directly. During cold start, while completed requests are fewer than ``running_max_req_size`` or the - QPS window is not initialized, the base ``running_max_req_size`` is used to collect samples quickly. Afterwards, + HTTP 429 directly. During cold start, while the recent-completion window contains fewer than 64 requests or QPS + is not initialized, the sum of ``running_max_req_size`` across all Decode nodes is used as the concurrency limit + to collect samples quickly. Afterwards, the limit is ``int(QPS * average whole-request seconds) + 6``; six extra requests provide probing headroom so low traffic does not permanently trap the service at low concurrency. The PD Master measures complete PD requests, so ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` defaults to 60 seconds when unset. diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index f074d689a4..78fb43f994 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -135,7 +135,9 @@ async def generate( request: Request, ): if self.pd_master_request_limit_enabled: - max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count() + # QPS 尚未形成稳定估算时,以当前所有 Decode 节点声明的并发容量之和作为准入上限。 + decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) + max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count(decode_capacity) if self.running_request_count > max_allowed_request_count: logger.warning( f"PD Master rejects request before dispatch: running_request_count={self.running_request_count}, " diff --git a/lightllm/server/httpserver_for_pd_master/qps_recorder.py b/lightllm/server/httpserver_for_pd_master/qps_recorder.py index b3a478e83f..e91e746961 100644 --- a/lightllm/server/httpserver_for_pd_master/qps_recorder.py +++ b/lightllm/server/httpserver_for_pd_master/qps_recorder.py @@ -15,10 +15,8 @@ def __init__(self, args, ema_alpha: float = 0.1): self.args = args self.ema_alpha = float(ema_alpha) - # 保存最近 16 个请求的完成时间。16 个时间点之间包含 15 个完成间隔。 - self._finished_timestamps: Deque[float] = deque(maxlen=16) - # 记录服务启动后已经完成的请求总数,用于判断冷启动阶段是否已收集足够样本。 - self._finished_request_count = 0 + # 保存最近 64 个请求的完成时间。64 个时间点之间包含 63 个完成间隔。 + self._finished_timestamps: Deque[float] = deque(maxlen=64) self._qps = 0.0 self._initialized = False self._last_qps_update_time: Optional[float] = None @@ -29,7 +27,6 @@ def mark_one_req_finish(self) -> None: finished_time = time.monotonic() with self._lock: self._finished_timestamps.append(finished_time) - self._finished_request_count += 1 self._update_qps() def get_qps(self) -> float: @@ -41,13 +38,15 @@ def get_qps(self) -> float: self._update_qps() return self._qps - def get_max_allowed_request_count(self) -> int: - """根据冷启动样本数和动态 QPS 返回 PD Master 最大允许进入请求数。""" + def get_max_allowed_request_count(self, default_max_allowed_request_count: int) -> int: + """返回 PD Master 最大允许进入请求数。 + + QPS 尚未初始化或完成样本不足时,使用调用方传入的 Decode 节点总并发容量; + 样本充足后使用动态 QPS 估算值。 + """ with self._lock: - finished_request_count = self._finished_request_count - qps_initialized = self._initialized - if finished_request_count < self.args.running_max_req_size or not qps_initialized: - return self.args.running_max_req_size + if not self._initialized: + return default_max_allowed_request_count # PD Master 统计完整 PD 请求,按统一的平均整包时长配置估算在途请求数。 return int(self.get_qps() * get_pd_request_limit_max_allowed_request_count_seconds()) + 6 diff --git a/test/test_api/test_qps_recorder.py b/test/test_api/test_qps_recorder.py index b816a5e213..de47d71e38 100644 --- a/test/test_api/test_qps_recorder.py +++ b/test/test_api/test_qps_recorder.py @@ -13,11 +13,11 @@ def _args(run_mode="decode", running_max_req_size=16): return SimpleNamespace(run_mode=run_mode, running_max_req_size=running_max_req_size) -def test_qps_recorder_waits_for_sixteen_finished_requests(): +def test_qps_recorder_waits_for_sixty_four_finished_requests(): recorder = QPSRecorder(_args()) - with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(15)): - for _ in range(15): + with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(63)): + for _ in range(63): recorder.mark_one_req_finish() assert recorder.get_qps() == 0.0 @@ -28,14 +28,14 @@ def test_qps_recorder_calculates_qps_and_updates_ema(): with patch( "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=[*range(16), 15, 15, 15.5, 15.5, 15.5], + side_effect=[*range(64), 63, 63, 63.5, 63.5, 63.5], ): - for _ in range(16): + for _ in range(64): recorder.mark_one_req_finish() assert recorder.get_qps() == 1.0 recorder.mark_one_req_finish() - window_qps = 15 / 14.5 + window_qps = 63 / 62.5 expected_qps = 0.25 * window_qps + 0.75 * 1.0 assert recorder.get_qps() == pytest.approx(expected_qps) @@ -45,12 +45,12 @@ def test_get_qps_updates_ema_after_thirty_seconds_without_new_request(): with patch( "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=[*range(16), 15, 45.1, 45.1, 46], + side_effect=[*range(64), 63, 93.1, 93.1, 94], ): - for _ in range(16): + for _ in range(64): recorder.mark_one_req_finish() - stale_window_qps = 15 / 45.1 + stale_window_qps = 63 / 93.1 expected_qps = 0.25 * stale_window_qps + 0.75 * 1.0 assert recorder.get_qps() == pytest.approx(expected_qps) assert recorder.get_qps() == pytest.approx(expected_qps) @@ -60,13 +60,13 @@ def test_max_allowed_request_count_uses_env(monkeypatch): monkeypatch.setenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", "12") get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() recorder = QPSRecorder(_args(running_max_req_size=1)) - with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(17)): - for _ in range(16): + with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(65)): + for _ in range(64): recorder.mark_one_req_finish() try: with patch.object(recorder, "get_qps", return_value=2.5): - assert recorder.get_max_allowed_request_count() == 36 + assert recorder.get_max_allowed_request_count(1) == 36 finally: get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() @@ -75,7 +75,7 @@ def test_max_allowed_request_count_uses_running_capacity_during_warmup(): recorder = QPSRecorder(_args(run_mode="prefill", running_max_req_size=32)) with patch.object(recorder, "get_qps") as get_qps: - assert recorder.get_max_allowed_request_count() == 32 + assert recorder.get_max_allowed_request_count(48) == 48 get_qps.assert_not_called() @@ -84,7 +84,7 @@ def test_max_allowed_request_count_waits_until_qps_is_initialized(): recorder.mark_one_req_finish() with patch.object(recorder, "get_qps") as get_qps: - assert recorder.get_max_allowed_request_count() == 1 + assert recorder.get_max_allowed_request_count(8) == 8 get_qps.assert_not_called() @@ -92,13 +92,13 @@ def test_max_allowed_request_count_keeps_six_probe_requests_at_zero_qps(): recorder = QPSRecorder(_args(run_mode="decode", running_max_req_size=1)) with patch( "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=range(17), + side_effect=range(65), ): - for _ in range(16): + for _ in range(64): recorder.mark_one_req_finish() with patch.object(recorder, "get_qps", return_value=0.0): - assert recorder.get_max_allowed_request_count() == 6 + assert recorder.get_max_allowed_request_count(1) == 6 def test_pd_request_limit_max_allowed_request_count_seconds_uses_unified_default(monkeypatch): diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 888bed4013..0526cc66f9 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -26,6 +26,12 @@ def test_pd_master_qps_limit_rejects_before_dispatch(): manager.running_request_count = 3 manager.qps_recorder = MagicMock() manager.qps_recorder.get_max_allowed_request_count.return_value = 2 + manager.pd_manager = SimpleNamespace( + decode_nodes=[ + SimpleNamespace(start_args={"running_max_req_size": 2}), + SimpleNamespace(start_args={"running_max_req_size": 3}), + ] + ) async def consume_generate(): async for _ in manager.generate("prompt", None, None, None): @@ -35,7 +41,7 @@ async def consume_generate(): asyncio.run(consume_generate()) assert manager.running_request_count == 3 - manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with() + manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with(5) def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): From 551887a55246a138989adac3f7ef580723227dbe Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 06:24:46 +0000 Subject: [PATCH 13/20] feat(pd): retry master admission before rejection --- docs/CN/source/tutorial/api_server_args.rst | 4 +- docs/EN/source/tutorial/api_server_args.rst | 5 ++- lightllm/server/api_cli.py | 5 ++- .../httpserver_for_pd_master/manager.py | 45 ++++++++++++++----- lightllm/utils/envs_utils.py | 9 ++++ unit_tests/server/test_pd_master_mode.py | 24 +++++++++- unit_tests/utils/test_envs_utils.py | 19 ++++++++ 7 files changed, 93 insertions(+), 18 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 47b7f0075d..673b986564 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -90,7 +90,9 @@ PD 分离模式参数 .. option:: --enable_pd_node_self_request_limit 在 PD Master 上启用基于动态 QPS 的请求准入探测。PD Master 根据最近完成请求计算 QPS, - 并限制同时进入的完整 PD 请求数;超过上限时直接向客户端返回 HTTP 429。服务启动后, + 并限制同时进入的完整 PD 请求数。超过上限时每 2 秒重新尝试一次,最长等待时间由 + ``LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS`` 控制(默认 15 秒);等待结束后仍然 + 无法进入时向客户端返回 HTTP 429。服务启动后, 最近完成时间窗口尚未积累满 64 个请求,或 QPS 尚未初始化时,暂时使用所有 Decode 节点 ``running_max_req_size`` 之和作为并发上限,以便快速积累样本;之后使用 ``int(QPS * 平均整包时间秒数) + 6``, diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 2ad2f5e9d1..45702fc85e 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -92,8 +92,9 @@ PD disaggregation Mode Parameters .. option:: --enable_pd_node_self_request_limit Enable dynamic-QPS admission probing on the PD Master. The PD Master computes QPS from recently completed - requests and limits the number of complete PD requests admitted concurrently; requests above the limit receive - HTTP 429 directly. During cold start, while the recent-completion window contains fewer than 64 requests or QPS + requests and limits the number of complete PD requests admitted concurrently. Requests above the limit retry + every two seconds for up to ``LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS`` (15 seconds by default), + then receive HTTP 429 if admission is still unavailable. During cold start, while the recent-completion window contains fewer than 64 requests or QPS is not initialized, the sum of ``running_max_req_size`` across all Decode nodes is used as the concurrency limit to collect samples quickly. Afterwards, the limit is ``int(QPS * average whole-request seconds) + 6``; six extra requests provide probing headroom so diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 73d0405295..e893c3728e 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -72,8 +72,9 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--enable_pd_node_self_request_limit", action="store_true", help=( - "Enable PD Master admission probing based on dynamically measured QPS; Prefill/Decode nodes " - "enforce both shm_req allocation timeout and Router scheduling wait timeout. Default: disabled." + "Enable PD Master admission probing based on dynamically measured QPS, with bounded admission retry; " + "Prefill/Decode nodes enforce both shm_req allocation timeout and Router scheduling wait timeout. " + "Default: disabled." ), ) parser.add_argument( diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 78fb43f994..12c6017488 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -23,7 +23,10 @@ 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_master_request_limit_wait_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 from .qps_recorder import QPSRecorder @@ -50,8 +53,9 @@ def __init__( self.latest_success_infer_time = time.time() self.running_request_count = 0 # PD Master 统一统计完整请求的 QPS,并据此控制进入请求数;P/D 节点只负责 - # shm_req 资源申请超时,避免各节点分别探测造成限流判断不一致。 + # 本地 shm_req 资源申请及 Router 调度等待超时,避免各节点分别探测造成限流判断不一致。 self.pd_master_request_limit_enabled = args.enable_pd_node_self_request_limit and args.run_mode == "pd_master" + self.pd_master_request_limit_wait_timeout_seconds = get_pd_master_request_limit_wait_timeout_seconds() self.qps_recorder = QPSRecorder(args) self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -127,24 +131,41 @@ async def select_p_d_node( ) -> Tuple[PD_Client_Obj, PD_Client_Obj]: return self.pd_manager.select_p_d_node(prompt, sampling_params, multimodal_params) - async def generate( - self, - prompt: Union[str, List[int]], - sampling_params: SamplingParams, - multimodal_params: MultimodalParams, - request: Request, - ): - if self.pd_master_request_limit_enabled: + async def _wait_for_pd_master_request_slot(self) -> None: + """等待 PD Master 动态并发限制放行,超时后拒绝请求。""" + if not self.pd_master_request_limit_enabled: + return + + wait_timeout_seconds = self.pd_master_request_limit_wait_timeout_seconds + deadline = time.monotonic() + wait_timeout_seconds + retry_interval_seconds = 2 + while True: # QPS 尚未形成稳定估算时,以当前所有 Decode 节点声明的并发容量之和作为准入上限。 decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count(decode_capacity) - if self.running_request_count > max_allowed_request_count: + if self.running_request_count < max_allowed_request_count: + return + + remaining_time = deadline - time.monotonic() + if remaining_time <= 0: logger.warning( - f"PD Master rejects request before dispatch: running_request_count={self.running_request_count}, " + f"PD Master rejects request after waiting {wait_timeout_seconds}s: " + f"running_request_count={self.running_request_count}, " f"max_allowed_request_count={max_allowed_request_count}" ) raise ServerBusyError("PD Master is busy") + await asyncio.sleep(min(retry_interval_seconds, remaining_time)) + + async def generate( + self, + prompt: Union[str, List[int]], + sampling_params: SamplingParams, + multimodal_params: MultimodalParams, + request: Request, + ): + await self._wait_for_pd_master_request_slot() + was_idle = self.running_request_count == 0 self.running_request_count += 1 if was_idle: diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 83848c69c3..7c5d10eab8 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -338,6 +338,15 @@ def get_pd_request_limit_max_allowed_request_count_seconds() -> int: return seconds +@lru_cache(maxsize=None) +def get_pd_master_request_limit_wait_timeout_seconds() -> int: + """PD Master 请求超过动态并发上限后的最长等待时间,单位为秒。""" + seconds = int(os.getenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", 15)) + if seconds < 0: + raise ValueError("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS must be greater than or equal to 0") + return seconds + + @lru_cache(maxsize=None) def get_lightllm_url_pool_maxsize() -> int: return int(os.getenv("LIGHTLLM_URL_POOL_MAXSIZE", 512)) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 0526cc66f9..89bd85236d 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,7 +1,7 @@ import asyncio import json from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest from easydict import EasyDict @@ -23,6 +23,7 @@ def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled( def test_pd_master_qps_limit_rejects_before_dispatch(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.pd_master_request_limit_enabled = True + manager.pd_master_request_limit_wait_timeout_seconds = 0 manager.running_request_count = 3 manager.qps_recorder = MagicMock() manager.qps_recorder.get_max_allowed_request_count.return_value = 2 @@ -44,6 +45,27 @@ async def consume_generate(): manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with(5) +def test_pd_master_qps_limit_retries_until_request_can_enter(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_master_request_limit_enabled = True + manager.pd_master_request_limit_wait_timeout_seconds = 15 + manager.running_request_count = 3 + manager.qps_recorder = MagicMock() + manager.qps_recorder.get_max_allowed_request_count.side_effect = [3, 4] + manager.pd_manager = SimpleNamespace(decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 4})]) + + with pytest.MonkeyPatch.context() as monkeypatch: + sleep = AsyncMock() + monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.asyncio.sleep", sleep) + await manager._wait_for_pd_master_request_slot() + + sleep.assert_awaited_once_with(2) + assert manager.qps_recorder.get_max_allowed_request_count.call_count == 2 + + asyncio.run(run()) + + 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 index a806f8c6de..73fe52fda8 100644 --- a/unit_tests/utils/test_envs_utils.py +++ b/unit_tests/utils/test_envs_utils.py @@ -1,4 +1,5 @@ from lightllm.utils.envs_utils import ( + get_pd_master_request_limit_wait_timeout_seconds, get_pd_node_router_wait_timeout_seconds, get_pd_node_shm_req_alloc_timeout_seconds, ) @@ -38,3 +39,21 @@ def test_pd_node_router_wait_timeout_reads_environment_variable(monkeypatch): assert get_pd_node_router_wait_timeout_seconds() == 45 get_pd_node_router_wait_timeout_seconds.cache_clear() + + +def test_pd_master_request_limit_wait_timeout_defaults_to_15_seconds(monkeypatch): + monkeypatch.delenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", raising=False) + get_pd_master_request_limit_wait_timeout_seconds.cache_clear() + + assert get_pd_master_request_limit_wait_timeout_seconds() == 15 + + get_pd_master_request_limit_wait_timeout_seconds.cache_clear() + + +def test_pd_master_request_limit_wait_timeout_reads_environment_variable(monkeypatch): + monkeypatch.setenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", "21") + get_pd_master_request_limit_wait_timeout_seconds.cache_clear() + + assert get_pd_master_request_limit_wait_timeout_seconds() == 21 + + get_pd_master_request_limit_wait_timeout_seconds.cache_clear() From 713cb156b589f52375866aae0a3f7956172819c0 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 08:31:06 +0000 Subject: [PATCH 14/20] test: cover PD request limiting scenarios --- .../test_running_request_lifecycle.py | 27 +++++++++--- unit_tests/server/test_pd_master_mode.py | 42 ++++++++++++++++++- 2 files changed, 61 insertions(+), 8 deletions(-) diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 08ab652af6..5f9c2e56be 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -198,18 +198,32 @@ async def run(): asyncio.run(run()) -def test_httpserver_returns_busy_when_shm_req_allocation_times_out(): +@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(NodeRole.D) + 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="PD decode node is busy"): - await _drain_generate(manager, _sampling_params(), _multimodal_params()) + 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() @@ -219,9 +233,10 @@ async def run(): asyncio.run(run()) -def test_httpserver_returns_busy_while_first_token_request_waits_in_router(): +@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(NodeRole.D) + manager = _make_manager(mode) manager.pd_node_request_limit_enabled = True req = SimpleNamespace( request_id=123, diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 89bd85236d..070b606eef 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -24,7 +24,8 @@ def test_pd_master_qps_limit_rejects_before_dispatch(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.pd_master_request_limit_enabled = True manager.pd_master_request_limit_wait_timeout_seconds = 0 - manager.running_request_count = 3 + # 当前在途数等于上限时也不能继续放行,避免实际并发突破准入上限。 + manager.running_request_count = 2 manager.qps_recorder = MagicMock() manager.qps_recorder.get_max_allowed_request_count.return_value = 2 manager.pd_manager = SimpleNamespace( @@ -41,7 +42,7 @@ async def consume_generate(): with pytest.raises(ServerBusyError, match="PD Master is busy"): asyncio.run(consume_generate()) - assert manager.running_request_count == 3 + assert manager.running_request_count == 2 manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with(5) @@ -66,6 +67,43 @@ async def run(): asyncio.run(run()) +def test_pd_master_qps_limit_retries_until_timeout_then_rejects(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_master_request_limit_enabled = True + manager.pd_master_request_limit_wait_timeout_seconds = 3 + manager.running_request_count = 4 + manager.qps_recorder = MagicMock() + manager.qps_recorder.get_max_allowed_request_count.return_value = 4 + manager.pd_manager = SimpleNamespace(decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 4})]) + + monotonic_values = iter([0, 0, 2, 3]) + with pytest.MonkeyPatch.context() as monkeypatch: + sleep = AsyncMock() + monkeypatch.setattr( + "lightllm.server.httpserver_for_pd_master.manager.time.monotonic", + lambda: next(monotonic_values), + ) + monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.asyncio.sleep", sleep) + with pytest.raises(ServerBusyError, match="PD Master is busy"): + await manager._wait_for_pd_master_request_slot() + + assert [call.args[0] for call in sleep.await_args_list] == [2, 1] + assert manager.qps_recorder.get_max_allowed_request_count.call_count == 3 + + asyncio.run(run()) + + +def test_pd_master_qps_limit_disabled_does_not_query_capacity(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_master_request_limit_enabled = False + manager.qps_recorder = MagicMock() + + asyncio.run(manager._wait_for_pd_master_request_slot()) + + manager.qps_recorder.get_max_allowed_request_count.assert_not_called() + + def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): from lightllm.utils.config_utils import auto_set_response_parsers From 46dac36f860bc59ef3aec7800699f98c2029c123 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 08:35:22 +0000 Subject: [PATCH 15/20] docs: remove outdated PD capacity note --- docs/CN/source/tutorial/api_server_args.rst | 2 -- docs/EN/source/tutorial/api_server_args.rst | 2 -- 2 files changed, 4 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 673b986564..09d7cb97a3 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -166,8 +166,6 @@ PD 分离模式参数 .. option:: --running_max_req_size 同时进行前向推理的最大请求数量,默认为 ``1000`` - 在 PD 分离模式的 Decode 节点上,该限制仅在各节点本地生效; - PD Master 不会汇总各 Decode 节点的值作为全局请求准入上限。 .. option:: --max_req_total_len diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 45702fc85e..4d3623dce7 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -168,8 +168,6 @@ Memory and Batch Processing Parameters .. option:: --running_max_req_size Maximum number of requests for simultaneous forward inference, default is ``1000`` - On Decode nodes in PD disaggregation mode, this limit applies locally to each node; - PD Master does not aggregate the Decode-node values into a global admission limit. .. option:: --max_req_total_len From 53dc4205e5907f212d715593f3a30295c248871a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 09:03:26 +0000 Subject: [PATCH 16/20] refactor: disable PD Master request limiting --- docs/CN/source/tutorial/api_server_args.rst | 12 +- docs/EN/source/tutorial/api_server_args.rst | 12 +- lightllm/server/api_cli.py | 5 +- .../httpserver_for_pd_master/manager.py | 38 +----- .../httpserver_for_pd_master/qps_recorder.py | 69 ----------- lightllm/utils/envs_utils.py | 24 ---- test/test_api/test_qps_recorder.py | 117 ------------------ .../test_pd_master_multi_choice.py | 4 - .../test_pd_master_cached_tokens.py | 2 - unit_tests/server/test_pd_master_mode.py | 84 +------------ unit_tests/utils/test_envs_utils.py | 19 --- 11 files changed, 10 insertions(+), 376 deletions(-) delete mode 100644 lightllm/server/httpserver_for_pd_master/qps_recorder.py delete mode 100644 test/test_api/test_qps_recorder.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 09d7cb97a3..90b7f357f0 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -89,16 +89,8 @@ PD 分离模式参数 .. option:: --enable_pd_node_self_request_limit - 在 PD Master 上启用基于动态 QPS 的请求准入探测。PD Master 根据最近完成请求计算 QPS, - 并限制同时进入的完整 PD 请求数。超过上限时每 2 秒重新尝试一次,最长等待时间由 - ``LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS`` 控制(默认 15 秒);等待结束后仍然 - 无法进入时向客户端返回 HTTP 429。服务启动后, - 最近完成时间窗口尚未积累满 64 个请求,或 QPS 尚未初始化时,暂时使用所有 Decode 节点 - ``running_max_req_size`` 之和作为并发上限,以便快速积累样本;之后使用 - ``int(QPS * 平均整包时间秒数) + 6``, - 额外保留 6 个请求作为探测余量,避免低流量后无法恢复。PD Master 统计完整 PD 请求, - ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` 未设置时默认使用 60 秒。 - Prefill/Decode 节点不再执行 QPS 准入。HTTP server 申请本地 ``shm_req`` 对象的超时时间由 + 在 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, diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 4d3623dce7..4fc3afb6ee 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -91,16 +91,8 @@ PD disaggregation Mode Parameters .. option:: --enable_pd_node_self_request_limit - Enable dynamic-QPS admission probing on the PD Master. The PD Master computes QPS from recently completed - requests and limits the number of complete PD requests admitted concurrently. Requests above the limit retry - every two seconds for up to ``LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS`` (15 seconds by default), - then receive HTTP 429 if admission is still unavailable. During cold start, while the recent-completion window contains fewer than 64 requests or QPS - is not initialized, the sum of ``running_max_req_size`` across all Decode nodes is used as the concurrency limit - to collect samples quickly. Afterwards, - the limit is ``int(QPS * average whole-request seconds) + 6``; six extra requests provide probing headroom so - low traffic does not permanently trap the service at low concurrency. The PD Master measures complete PD - requests, so ``LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS`` defaults to 60 seconds when unset. - Prefill/Decode nodes no longer perform QPS admission. The local ``shm_req`` allocation timeout is controlled by + 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 diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e893c3728e..b27736a62f 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -72,9 +72,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--enable_pd_node_self_request_limit", action="store_true", help=( - "Enable PD Master admission probing based on dynamically measured QPS, with bounded admission retry; " - "Prefill/Decode nodes enforce both shm_req allocation timeout and Router scheduling wait timeout. " - "Default: disabled." + "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( diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 12c6017488..36e710ea64 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -23,13 +23,9 @@ 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_master_request_limit_wait_timeout_seconds, - get_pd_split_max_new_tokens, -) +from lightllm.utils.envs_utils import get_pd_split_max_new_tokens from lightllm.utils.shm_port_args import get_shm_port_args from .pd_selector import create_selector -from .qps_recorder import QPSRecorder logger = init_logger(__name__) @@ -52,11 +48,6 @@ def __init__( self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 - # PD Master 统一统计完整请求的 QPS,并据此控制进入请求数;P/D 节点只负责 - # 本地 shm_req 资源申请及 Router 调度等待超时,避免各节点分别探测造成限流判断不一致。 - self.pd_master_request_limit_enabled = args.enable_pd_node_self_request_limit and args.run_mode == "pd_master" - self.pd_master_request_limit_wait_timeout_seconds = get_pd_master_request_limit_wait_timeout_seconds() - self.qps_recorder = QPSRecorder(args) self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -132,30 +123,8 @@ async def select_p_d_node( 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 动态并发限制放行,超时后拒绝请求。""" - if not self.pd_master_request_limit_enabled: - return - - wait_timeout_seconds = self.pd_master_request_limit_wait_timeout_seconds - deadline = time.monotonic() + wait_timeout_seconds - retry_interval_seconds = 2 - while True: - # QPS 尚未形成稳定估算时,以当前所有 Decode 节点声明的并发容量之和作为准入上限。 - decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - max_allowed_request_count = self.qps_recorder.get_max_allowed_request_count(decode_capacity) - if self.running_request_count < max_allowed_request_count: - return - - remaining_time = deadline - time.monotonic() - if remaining_time <= 0: - logger.warning( - f"PD Master rejects request after waiting {wait_timeout_seconds}s: " - f"running_request_count={self.running_request_count}, " - f"max_allowed_request_count={max_allowed_request_count}" - ) - raise ServerBusyError("PD Master is busy") - - await asyncio.sleep(min(retry_interval_seconds, remaining_time)) + """PD Master 请求准入的预留接口,当前不执行限流。""" + return async def generate( self, @@ -232,7 +201,6 @@ async def _generate( yield result if request_finished_successfully: self.metric_client.counter_inc("lightllm_request_success") - self.qps_recorder.mark_one_req_finish() return async def _generate_one( diff --git a/lightllm/server/httpserver_for_pd_master/qps_recorder.py b/lightllm/server/httpserver_for_pd_master/qps_recorder.py deleted file mode 100644 index e91e746961..0000000000 --- a/lightllm/server/httpserver_for_pd_master/qps_recorder.py +++ /dev/null @@ -1,69 +0,0 @@ -import time -from collections import deque -from threading import Lock -from typing import Deque, Optional - -from lightllm.utils.envs_utils import get_pd_request_limit_max_allowed_request_count_seconds - - -class QPSRecorder: - """根据最近完成的请求计算 PD Master 的动态 QPS。""" - - def __init__(self, args, ema_alpha: float = 0.1): - if not 0 < ema_alpha <= 1: - raise ValueError("ema_alpha must be in the range (0, 1]") - - self.args = args - self.ema_alpha = float(ema_alpha) - # 保存最近 64 个请求的完成时间。64 个时间点之间包含 63 个完成间隔。 - self._finished_timestamps: Deque[float] = deque(maxlen=64) - self._qps = 0.0 - self._initialized = False - self._last_qps_update_time: Optional[float] = None - self._lock = Lock() - - def mark_one_req_finish(self) -> None: - """记录一个请求完成事件,并在样本充足时更新全局 QPS。""" - finished_time = time.monotonic() - with self._lock: - self._finished_timestamps.append(finished_time) - self._update_qps() - - def get_qps(self) -> float: - """返回经过 EMA 平滑后的全局 QPS。""" - with self._lock: - if self._last_qps_update_time is not None: - current_time = time.monotonic() - if current_time - self._last_qps_update_time > 30: - self._update_qps() - return self._qps - - def get_max_allowed_request_count(self, default_max_allowed_request_count: int) -> int: - """返回 PD Master 最大允许进入请求数。 - - QPS 尚未初始化或完成样本不足时,使用调用方传入的 Decode 节点总并发容量; - 样本充足后使用动态 QPS 估算值。 - """ - with self._lock: - if not self._initialized: - return default_max_allowed_request_count - - # PD Master 统计完整 PD 请求,按统一的平均整包时长配置估算在途请求数。 - return int(self.get_qps() * get_pd_request_limit_max_allowed_request_count_seconds()) + 6 - - def _update_qps(self) -> None: - if len(self._finished_timestamps) < self._finished_timestamps.maxlen: - return - - current_time = time.monotonic() - elapsed_time = current_time - self._finished_timestamps[0] - if elapsed_time <= 0: - return - - average_qps = (len(self._finished_timestamps) - 1) / elapsed_time - if not self._initialized: - self._qps = average_qps - self._initialized = True - else: - self._qps = self.ema_alpha * average_qps + (1 - self.ema_alpha) * self._qps - self._last_qps_update_time = current_time diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 7c5d10eab8..bb9799edb5 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -323,30 +323,6 @@ def get_pd_node_router_wait_timeout_seconds() -> int: return int(os.getenv("LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS", 20)) -@lru_cache(maxsize=None) -def get_pd_request_limit_max_allowed_request_count_seconds() -> int: - """获取根据 QPS 估算 PD Master 最大在途请求数时使用的平均整包时间。 - - 该值与完成 QPS 相乘,用于估算在目标平均整包时间内可以容纳的完整 PD 请求数。 - PD Master 统一统计完整请求,不再区分 Prefill 和 Decode,默认值为 60 秒。 - """ - seconds = int(os.getenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", 60)) - if seconds < 0: - raise ValueError( - "LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS must be greater than or equal to 0" - ) - return seconds - - -@lru_cache(maxsize=None) -def get_pd_master_request_limit_wait_timeout_seconds() -> int: - """PD Master 请求超过动态并发上限后的最长等待时间,单位为秒。""" - seconds = int(os.getenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", 15)) - if seconds < 0: - raise ValueError("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS must be greater than or equal to 0") - return seconds - - @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_api/test_qps_recorder.py b/test/test_api/test_qps_recorder.py deleted file mode 100644 index de47d71e38..0000000000 --- a/test/test_api/test_qps_recorder.py +++ /dev/null @@ -1,117 +0,0 @@ -from types import SimpleNamespace -from unittest.mock import patch - -import pytest - -from lightllm.server.httpserver_for_pd_master.qps_recorder import QPSRecorder -from lightllm.utils.envs_utils import ( - get_pd_request_limit_max_allowed_request_count_seconds, -) - - -def _args(run_mode="decode", running_max_req_size=16): - return SimpleNamespace(run_mode=run_mode, running_max_req_size=running_max_req_size) - - -def test_qps_recorder_waits_for_sixty_four_finished_requests(): - recorder = QPSRecorder(_args()) - - with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(63)): - for _ in range(63): - recorder.mark_one_req_finish() - - assert recorder.get_qps() == 0.0 - - -def test_qps_recorder_calculates_qps_and_updates_ema(): - recorder = QPSRecorder(_args(), ema_alpha=0.25) - - with patch( - "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=[*range(64), 63, 63, 63.5, 63.5, 63.5], - ): - for _ in range(64): - recorder.mark_one_req_finish() - assert recorder.get_qps() == 1.0 - - recorder.mark_one_req_finish() - window_qps = 63 / 62.5 - expected_qps = 0.25 * window_qps + 0.75 * 1.0 - assert recorder.get_qps() == pytest.approx(expected_qps) - - -def test_get_qps_updates_ema_after_thirty_seconds_without_new_request(): - recorder = QPSRecorder(_args(), ema_alpha=0.25) - - with patch( - "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=[*range(64), 63, 93.1, 93.1, 94], - ): - for _ in range(64): - recorder.mark_one_req_finish() - - stale_window_qps = 63 / 93.1 - expected_qps = 0.25 * stale_window_qps + 0.75 * 1.0 - assert recorder.get_qps() == pytest.approx(expected_qps) - assert recorder.get_qps() == pytest.approx(expected_qps) - - -def test_max_allowed_request_count_uses_env(monkeypatch): - monkeypatch.setenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", "12") - get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() - recorder = QPSRecorder(_args(running_max_req_size=1)) - with patch("lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", side_effect=range(65)): - for _ in range(64): - recorder.mark_one_req_finish() - - try: - with patch.object(recorder, "get_qps", return_value=2.5): - assert recorder.get_max_allowed_request_count(1) == 36 - finally: - get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() - - -def test_max_allowed_request_count_uses_running_capacity_during_warmup(): - recorder = QPSRecorder(_args(run_mode="prefill", running_max_req_size=32)) - - with patch.object(recorder, "get_qps") as get_qps: - assert recorder.get_max_allowed_request_count(48) == 48 - get_qps.assert_not_called() - - -def test_max_allowed_request_count_waits_until_qps_is_initialized(): - recorder = QPSRecorder(_args(run_mode="prefill", running_max_req_size=1)) - recorder.mark_one_req_finish() - - with patch.object(recorder, "get_qps") as get_qps: - assert recorder.get_max_allowed_request_count(8) == 8 - get_qps.assert_not_called() - - -def test_max_allowed_request_count_keeps_six_probe_requests_at_zero_qps(): - recorder = QPSRecorder(_args(run_mode="decode", running_max_req_size=1)) - with patch( - "lightllm.server.httpserver_for_pd_master.qps_recorder.time.monotonic", - side_effect=range(65), - ): - for _ in range(64): - recorder.mark_one_req_finish() - - with patch.object(recorder, "get_qps", return_value=0.0): - assert recorder.get_max_allowed_request_count(1) == 6 - - -def test_pd_request_limit_max_allowed_request_count_seconds_uses_unified_default(monkeypatch): - monkeypatch.delenv("LIGHTLLM_PD_REQUEST_LIMIT_MAX_ALLOWED_REQUEST_COUNT_SECONDS", raising=False) - get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() - - try: - assert get_pd_request_limit_max_allowed_request_count_seconds() == 60 - finally: - get_pd_request_limit_max_allowed_request_count_seconds.cache_clear() - - -@pytest.mark.parametrize("ema_alpha", [0, -0.1, 1.1]) -def test_qps_recorder_rejects_invalid_ema_alpha(ema_alpha): - with pytest.raises(ValueError, match=r"ema_alpha must be in the range \(0, 1\]"): - QPSRecorder(_args(), ema_alpha=ema_alpha) 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 8b75c2e7cf..d56c11cdba 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -16,8 +16,6 @@ def _manager() -> HttpServerManagerForPDMaster: manager.id_gen = MagicMock() manager.id_gen.generate_id.return_value = 800 manager.metric_client = MagicMock() - manager.pd_master_request_limit_enabled = False - manager.qps_recorder = MagicMock() manager._log_req_header = AsyncMock() manager.tokens = MagicMock(return_value=2) return manager @@ -96,7 +94,6 @@ async def wait_to_token_package( call("lightllm_request_count"), call("lightllm_request_success"), ] - manager.qps_recorder.mark_one_req_finish.assert_called_once_with() assert p_node.dispatched_prompt_chars == 0 assert p_node.dispatched_req_num == 0 @@ -180,7 +177,6 @@ async def generate_one(*_args, **_kwargs): 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") - manager.qps_recorder.mark_one_req_finish.assert_not_called() asyncio.run(asyncio.wait_for(run(), timeout=2)) diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index 07f0dffa40..a3dc268a93 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -16,8 +16,6 @@ def _make_manager(monkeypatch): monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) mgr.running_request_count = 0 - mgr.pd_master_request_limit_enabled = False - mgr.qps_recorder = SimpleNamespace(mark_one_req_finish=lambda: None) counter = [0] def gen_id(): diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 070b606eef..8d064c357b 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,7 +1,6 @@ import asyncio import json from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock import pytest from easydict import EasyDict @@ -9,7 +8,6 @@ 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 -from lightllm.utils.error_utils import ServerBusyError def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled(): @@ -20,89 +18,11 @@ def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled( assert StartArgs().enable_pd_node_self_request_limit is False -def test_pd_master_qps_limit_rejects_before_dispatch(): +def test_pd_master_request_slot_is_reserved_noop(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.pd_master_request_limit_enabled = True - manager.pd_master_request_limit_wait_timeout_seconds = 0 - # 当前在途数等于上限时也不能继续放行,避免实际并发突破准入上限。 - manager.running_request_count = 2 - manager.qps_recorder = MagicMock() - manager.qps_recorder.get_max_allowed_request_count.return_value = 2 - manager.pd_manager = SimpleNamespace( - decode_nodes=[ - SimpleNamespace(start_args={"running_max_req_size": 2}), - SimpleNamespace(start_args={"running_max_req_size": 3}), - ] - ) - - async def consume_generate(): - async for _ in manager.generate("prompt", None, None, None): - pass - - with pytest.raises(ServerBusyError, match="PD Master is busy"): - asyncio.run(consume_generate()) - - assert manager.running_request_count == 2 - manager.qps_recorder.get_max_allowed_request_count.assert_called_once_with(5) - - -def test_pd_master_qps_limit_retries_until_request_can_enter(): - async def run(): - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.pd_master_request_limit_enabled = True - manager.pd_master_request_limit_wait_timeout_seconds = 15 - manager.running_request_count = 3 - manager.qps_recorder = MagicMock() - manager.qps_recorder.get_max_allowed_request_count.side_effect = [3, 4] - manager.pd_manager = SimpleNamespace(decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 4})]) - - with pytest.MonkeyPatch.context() as monkeypatch: - sleep = AsyncMock() - monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.asyncio.sleep", sleep) - await manager._wait_for_pd_master_request_slot() - - sleep.assert_awaited_once_with(2) - assert manager.qps_recorder.get_max_allowed_request_count.call_count == 2 - - asyncio.run(run()) - - -def test_pd_master_qps_limit_retries_until_timeout_then_rejects(): - async def run(): - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.pd_master_request_limit_enabled = True - manager.pd_master_request_limit_wait_timeout_seconds = 3 - manager.running_request_count = 4 - manager.qps_recorder = MagicMock() - manager.qps_recorder.get_max_allowed_request_count.return_value = 4 - manager.pd_manager = SimpleNamespace(decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 4})]) - - monotonic_values = iter([0, 0, 2, 3]) - with pytest.MonkeyPatch.context() as monkeypatch: - sleep = AsyncMock() - monkeypatch.setattr( - "lightllm.server.httpserver_for_pd_master.manager.time.monotonic", - lambda: next(monotonic_values), - ) - monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.asyncio.sleep", sleep) - with pytest.raises(ServerBusyError, match="PD Master is busy"): - await manager._wait_for_pd_master_request_slot() - - assert [call.args[0] for call in sleep.await_args_list] == [2, 1] - assert manager.qps_recorder.get_max_allowed_request_count.call_count == 3 - - asyncio.run(run()) - - -def test_pd_master_qps_limit_disabled_does_not_query_capacity(): - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.pd_master_request_limit_enabled = False - manager.qps_recorder = MagicMock() asyncio.run(manager._wait_for_pd_master_request_slot()) - manager.qps_recorder.get_max_allowed_request_count.assert_not_called() - def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): from lightllm.utils.config_utils import auto_set_response_parsers @@ -364,7 +284,6 @@ async def verify_and_preload(self, request): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.running_request_count = 0 - manager.pd_master_request_limit_enabled = False async def consume_generate(): async for _ in manager.generate("prompt", None, FailingMultimodalParams(), None): @@ -379,7 +298,6 @@ async def consume_generate(): def test_pd_master_request_count_covers_async_generator_lifecycle(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.running_request_count = 0 - manager.pd_master_request_limit_enabled = False inner_generator_closed = False async def fake_generate(prompt, sampling_params, multimodal_params, request): diff --git a/unit_tests/utils/test_envs_utils.py b/unit_tests/utils/test_envs_utils.py index 73fe52fda8..a806f8c6de 100644 --- a/unit_tests/utils/test_envs_utils.py +++ b/unit_tests/utils/test_envs_utils.py @@ -1,5 +1,4 @@ from lightllm.utils.envs_utils import ( - get_pd_master_request_limit_wait_timeout_seconds, get_pd_node_router_wait_timeout_seconds, get_pd_node_shm_req_alloc_timeout_seconds, ) @@ -39,21 +38,3 @@ def test_pd_node_router_wait_timeout_reads_environment_variable(monkeypatch): assert get_pd_node_router_wait_timeout_seconds() == 45 get_pd_node_router_wait_timeout_seconds.cache_clear() - - -def test_pd_master_request_limit_wait_timeout_defaults_to_15_seconds(monkeypatch): - monkeypatch.delenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", raising=False) - get_pd_master_request_limit_wait_timeout_seconds.cache_clear() - - assert get_pd_master_request_limit_wait_timeout_seconds() == 15 - - get_pd_master_request_limit_wait_timeout_seconds.cache_clear() - - -def test_pd_master_request_limit_wait_timeout_reads_environment_variable(monkeypatch): - monkeypatch.setenv("LIGHTLLM_PD_MASTER_REQUEST_LIMIT_WAIT_TIMEOUT_SECONDS", "21") - get_pd_master_request_limit_wait_timeout_seconds.cache_clear() - - assert get_pd_master_request_limit_wait_timeout_seconds() == 21 - - get_pd_master_request_limit_wait_timeout_seconds.cache_clear() From 0b64634a3fba4cb026a427a473a5525d72eab2e4 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 10:06:57 +0000 Subject: [PATCH 17/20] feat: prioritize PD requests with high cache affinity --- docs/CN/source/tutorial/api_server_args.rst | 3 +- docs/EN/source/tutorial/api_server_args.rst | 3 +- lightllm/server/core/objs/sampling_params.py | 3 +- .../httpserver_for_pd_master/manager.py | 26 +++++---- .../pd_selector/cache_aware.py | 17 ++++-- .../pd_selector/pd_selector.py | 28 ++++++---- .../test_pd_master_multi_choice.py | 54 +++++++++++++++++-- .../test_pd_master_cached_tokens.py | 49 ++++++++++++++++- unit_tests/server/test_pd_cache_aware.py | 37 +++++++++++++ 9 files changed, 188 insertions(+), 32 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 90b7f357f0..1197d474af 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -95,7 +95,8 @@ PD 分离模式参数 进入推理系统的超时时间由 ``LIGHTLLM_PD_NODE_ROUTER_WAIT_TIMEOUT_SECONDS`` 控制(默认 20 秒)。 超时会导致 ``Server is busy``;其中已进入 Router 但仍未进入推理系统的请求会主动标记为 aborted, 由 PD Master 转换为 HTTP 429; - 未开启限流以及 PD 高优先级分段续跑请求不受该超时限制,会持续等待资源。该参数默认关闭。 + 未开启限流以及 PD 高优先级请求(分段续跑请求或预计输入 cache 命中率高于 0.8 + 的请求)不受该超时限制,会持续等待资源。该参数默认关闭。 .. 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 4fc3afb6ee..e80c18702b 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -97,7 +97,8 @@ PD disaggregation Mode Parameters 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. The timeout is bypassed when local admission is - disabled and for PD high-priority segmented continuation requests, which continue waiting for resources. Disabled by default. + disabled and for PD high-priority requests (segmented continuation requests or requests whose estimated input + cache hit rate is above 0.8), which continue waiting for resources. Disabled by default. .. option:: --config_server_host diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 8e01a5f174..d14e1894ff 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -294,7 +294,8 @@ 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 为分段续跑请求设置,表示请求需以高优先级插入 Router 调度队列。 + # 由 PD Master 为分段续跑或预计 cache 命中率较高的请求设置,表示请求需 + # 以高优先级插入 Router 调度队列。 ("pd_high_priority_request", ctypes.c_bool), ("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 diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 36e710ea64..847cd49d2e 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -119,7 +119,7 @@ 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: @@ -225,7 +225,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") @@ -239,10 +241,10 @@ 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 - # 首段仍然遵守 P/D 节点的本地限流。第二段及后续分段说明该用户请求已经 - # 成功执行过一段,将其标记为 PD 高优先级请求,便于优先进入 Router 调度队列。 - # 高优先级请求仍可等待可用 shm_req 对象,避免因临时资源紧张导致分段续跑失败。 - sampling_params.pd_high_priority_request = iter_index > 0 + # 预计输入 cache 命中率高于 0.8 时,将首段也标记为 PD 高优先级请求, + # 使其优先进入 Router 调度队列,尽快复用已命中的 KV cache。第二段及后续 + # 分段仍统一使用高优先级,避免因临时资源紧张导致分段续跑失败。 + sampling_params.pd_high_priority_request = iter_index > 0 or estimated_cache_hit_rate > 0.8 # 分段请求始终复用循环外选定的 P 节点;这里只按每段实际发送的 # prompt 更新该节点的在途 prefill 负载,不会重新选点。 @@ -273,6 +275,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( @@ -909,6 +915,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/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index d56c11cdba..3462d090d3 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -39,7 +39,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() @@ -255,7 +255,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") @@ -277,6 +277,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)) @@ -295,7 +296,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)) dispatched_nodes = [] dispatched_prompts = [] dispatched_loads = [] @@ -341,6 +342,51 @@ 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_releases_prefill_load_when_stream_is_closed(): async def run(): manager = _manager() @@ -355,7 +401,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/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/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(): From 31b17e02b00ad980573f73442a54141da2fa4d17 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 10:38:10 +0000 Subject: [PATCH 18/20] feat(pd): bound high-priority request waits --- docs/CN/source/tutorial/api_server_args.rst | 8 +- docs/EN/source/tutorial/api_server_args.rst | 9 +- lightllm/server/core/objs/sampling_params.py | 7 +- lightllm/server/httpserver/manager.py | 41 ++++++--- .../httpserver_for_pd_master/manager.py | 11 ++- lightllm/utils/envs_utils.py | 6 ++ .../test_pd_master_multi_choice.py | 43 ++++++++++ .../test_pd_node_request_limit.py | 83 ++++++++++++++++++- 8 files changed, 186 insertions(+), 22 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 1197d474af..26a053f4ab 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -94,9 +94,11 @@ PD 分离模式参数 ``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 转换为 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 e80c18702b..20aecb973c 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -96,9 +96,12 @@ PD disaggregation Mode Parameters ``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. The timeout is bypassed when local admission is - disabled and for PD high-priority requests (segmented continuation requests or requests whose estimated input - cache hit rate is above 0.8), which continue waiting for resources. Disabled by default. + 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 diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index d14e1894ff..a7b4547d8d 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -297,6 +297,9 @@ class SamplingParams(ctypes.Structure): # 由 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), @@ -340,8 +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 的内部调度信息,不能由外部请求参数开启。 + # 这两个字段是 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) @@ -509,6 +513,7 @@ def to_dict(self): "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/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 27ff671e2a..8656c41fc1 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -437,11 +437,12 @@ async def generate( await self._register_running_request() running_request_registered = True - # 申请资源并存储。PD 分段续跑请求仍在 Router 队列中优先调度,同时绕过本地 - # shm_req 等待超时,避免已经开始执行的请求因临时资源紧张而被中断。 + # 申请资源并存储。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): @@ -548,17 +549,27 @@ 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) -> List[int]: + 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 索引,申请失败时回滚已分配的索引。 - 未开启本地限流或请求为 PD 高优先级请求时无限等待;普通请求在限流开启时, - 最多等待 ``LIGHTLLM_PD_NODE_SHM_REQ_ALLOC_TIMEOUT_SECONDS`` 秒。 + 未开启本地限流时无限等待。开启限流后,普通请求使用节点的 shm_req 申请 + 超时时间;高优先级请求取本地超时与 PD Master 下发值中的较大值。 """ alloced_req_indexes = [] - request_limit_applies = self.pd_node_request_limit_enabled and not pd_high_priority_request - alloc_deadline = ( - time.monotonic() + self.pd_node_shm_req_alloc_timeout_seconds if request_limit_applies else None - ) + 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: @@ -570,11 +581,11 @@ async def _alloc_shm_req_indexes(self, req_num: int, pd_high_priority_request: b 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"{self.pd_node_shm_req_alloc_timeout_seconds} seconds" + 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 {self.pd_node_shm_req_alloc_timeout_seconds} seconds" + f"within {alloc_timeout_seconds} seconds" ) await asyncio.sleep(sleep_time * sleep_time_factor) sleep_time = min(1, sleep_time * 1.1) @@ -1104,9 +1115,13 @@ def has_timed_out_waiting_for_inference(self, timeout_seconds: float) -> bool: # 组内任一请求已经进入新 batch,说明整个请求组已经开始执行,不能再按 Router 等待超时清理。 if any(req.infer_start_time > 0 for req in reqs): return False - # 高优先级分段续跑请求必须保证可以继续执行,不因 Router 短暂拥塞被清理。 + # 高优先级请求取本地 Router 超时与 PD Master 下发值中的较大值,既保证 + # 它比普通请求拥有更充足的等待机会,也避免资源异常时永久滞留。 if any(req.sample_params.pd_high_priority_request for req in reqs): - return False + 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: diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 847cd49d2e..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) @@ -245,6 +248,12 @@ async def _generate_one( # 使其优先进入 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 负载,不会重新选点。 diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index bb9799edb5..fb30f47938 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -323,6 +323,12 @@ def get_pd_node_router_wait_timeout_seconds() -> int: 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 3462d090d3..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 @@ -387,6 +388,48 @@ async def wait_to_token_package(_p_node, _d_node, _start_time, _prompt, sampling 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() diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py index 091c35cb5b..e79be4675f 100644 --- a/test/test_pd_selector/test_pd_node_request_limit.py +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -4,8 +4,10 @@ import pytest -from lightllm.server.httpserver.manager import HttpServerManager +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.utils.error_utils import ServerBusyError class FakeSharedInt: @@ -32,6 +34,18 @@ def _manager() -> HttpServerManager: 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() @@ -72,6 +86,73 @@ async def run(): 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_at_router_queue_head(): queue = BaseQueue.__new__(BaseQueue) queue.dp_index = 0 From 55f8b4fe88afc97d32854b5e74fe8b5f022ff304 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 10:47:59 +0000 Subject: [PATCH 19/20] fix(pd): preserve FIFO order for priority requests --- .../server/router/req_queue/base_queue.py | 16 +++++++++-- .../server/router/req_queue/dp_base_queue.py | 8 +++++- .../test_pd_node_request_limit.py | 28 +++++++++++++++++-- 3 files changed, 45 insertions(+), 7 deletions(-) diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index af39a869b6..27e51a4102 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -83,10 +83,20 @@ 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 - # PD 分段续跑请求已经执行过前一段,应优先进入调度队列,避免被新请求长时间阻塞。 + # PD 高优先级请求应排在普通请求之前,但高优先级请求之间仍按到达顺序排队, + # 避免后到请求反复插到队头而阻塞先到的高优先级请求。 if req_group and req_group[0].sample_params.pd_high_priority_request: - # req_group 可能包含同一请求组的多个 Req,整体前置可以保持组内顺序。 - self.waiting_req_list = req_group + self.waiting_req_list + 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 diff --git a/lightllm/server/router/req_queue/dp_base_queue.py b/lightllm/server/router/req_queue/dp_base_queue.py index 4986c62c65..b2d1ff0f7e 100644 --- a/lightllm/server/router/req_queue/dp_base_queue.py +++ b/lightllm/server/router/req_queue/dp_base_queue.py @@ -61,7 +61,13 @@ def extend(self, req_group: List[Req]): if suggested_dp_index >= self.dp_size_in_node or suggested_dp_index < 0: # 同一个组的,要分配在同一个 dp 上 if req_group[0].sample_params.pd_high_priority_request: - self.reqs_waiting_for_dp_index.insert(0, req_group) + # 高优先级请求组插在第一个普通请求组之前,同时保持高优先级组之间的 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: diff --git a/test/test_pd_selector/test_pd_node_request_limit.py b/test/test_pd_selector/test_pd_node_request_limit.py index e79be4675f..9e071cfa98 100644 --- a/test/test_pd_selector/test_pd_node_request_limit.py +++ b/test/test_pd_selector/test_pd_node_request_limit.py @@ -7,6 +7,7 @@ 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 @@ -153,16 +154,37 @@ def test_high_priority_router_wait_uses_master_timeout( assert req_status.has_timed_out_waiting_for_inference(local_timeout_seconds) is expected -def test_pd_high_priority_request_is_inserted_at_router_queue_head(): +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 = [normal_req] + queue.waiting_req_list = [earlier_high_req, normal_req] queue.extend([high_req_1, high_req_2]) - assert queue.waiting_req_list == [high_req_1, high_req_2, normal_req] + 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] From 0201cb31fb3b0cac655da40634af669e054fadff Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 2 Sep 2026 11:07:20 +0000 Subject: [PATCH 20/20] fix --- .../server/httpserver/test_running_request_lifecycle.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 5f9c2e56be..bf154826e8 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -260,7 +260,7 @@ async def run(): request=None, ) with pytest.raises(ServerBusyError, match="request did not enter inference"): - await anext(output_generator) + await output_generator.__anext__() asyncio.run(run()) @@ -269,7 +269,10 @@ 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), + sample_params=SimpleNamespace( + pd_high_priority_request=False, + pd_high_priority_request_time_out_seconds=60, + ), ) started_req = SimpleNamespace( router_arrival_time=1.0,