diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e..6500ad53b 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -71,7 +71,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: 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.", + help="Disable the PD master admission queue based on registered decode capacity.", ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 1d8b2112f..26c27e6fe 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -41,6 +41,8 @@ async def register_and_keep_alive(websocket: WebSocket): data = await asyncio.wait_for(websocket.receive_bytes(), timeout=heartbeat_timeout_seconds) obj = pickle.loads(data) if isinstance(obj, tuple) and obj and obj[0] == ObjType.HEARTBEAT: + load_info = obj[1] if len(obj) > 1 else None + g_objs.httpserver_manager.update_node_load_info(load_info) continue await g_objs.httpserver_manager.put_to_handle_queue(obj) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index e66540cb5..95a280bdf 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -9,9 +9,15 @@ import os import signal import sys +import time from typing import Dict, Optional, Union, List from websockets import ClientConnection -from lightllm.server.pd_io_struct import NodeRole, ObjType +from lightllm.server.pd_io_struct import ( + NodeRole, + ObjType, + PD_MASTER_CAPACITY_EPOCH_KEY, + PD_MASTER_CAPACITY_SHARE_KEY, +) from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger @@ -26,6 +32,51 @@ logger = init_logger(__name__) +def _update_pd_master_membership(manager: HttpServerManager, pd_master_ids) -> None: + """更新 Master 成员和容量版本,并立即唤醒心跳。""" + pd_master_ids = tuple(sorted(pd_master_ids)) + if getattr(manager, "pd_master_ids", ()) == pd_master_ids: + return + manager.pd_master_ids = pd_master_ids + manager.pd_master_capacity_epoch = max( + getattr(manager, "pd_master_capacity_epoch", 0) + 1, + time.time_ns(), + ) + membership_changed = getattr(manager, "pd_master_membership_changed", None) + if membership_changed is None: + membership_changed = manager.pd_master_membership_changed = asyncio.Event() + membership_changed.set() + + +def _allocate_capacity_share(total_capacity: int, pd_master_ids, pd_master_node_id: int) -> int: + """把节点容量确定性地切成互不重叠的 PD Master 租约池。""" + pd_master_ids = tuple(sorted(pd_master_ids)) + if not pd_master_ids or pd_master_node_id not in pd_master_ids: + return 0 + base, remainder = divmod(max(0, total_capacity), len(pd_master_ids)) + return base + int(pd_master_ids.index(pd_master_node_id) < remainder) + + +def _build_pd_registration_info(manager: HttpServerManager, pd_master_obj: PD_Master_Obj) -> dict: + """构造保持旧顶层 schema 兼容的 P/D 节点注册信息。""" + # Older Masters expand the registration JSON directly into PD_Client_Obj + # and reject unknown top-level fields during a rolling upgrade. + args_dict = vars(manager.args).copy() + args_dict["host"] = manager.host_ip + args_dict[PD_MASTER_CAPACITY_SHARE_KEY] = _allocate_capacity_share( + manager.args.running_max_req_size, + manager.pd_master_ids, + pd_master_obj.node_id, + ) + args_dict[PD_MASTER_CAPACITY_EPOCH_KEY] = manager.pd_master_capacity_epoch + return { + "node_id": manager.args.pd_node_id, + "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", + "mode": manager.pd_mode.value, + "start_args": args_dict, + } + + async def timer_log(manager: HttpServerManager): while True: await asyncio.sleep(30) @@ -56,7 +107,8 @@ async def pd_handle_loop(manager: HttpServerManager): logger.info(f"get pd_master_objs {id_to_pd_master_obj}") if id_to_pd_master_obj is not None: - for node_id, pd_master_obj in id_to_handle_task.items(): + _update_pd_master_membership(manager, id_to_pd_master_obj) + for node_id, pd_master_obj in list(id_to_handle_task.items()): if node_id not in id_to_pd_master_obj: id_to_handle_task[node_id].cancel() id_to_handle_task.pop(node_id, None) @@ -98,22 +150,19 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O sock = websocket.transport.get_extra_info("socket") sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - args_dict = vars(manager.args) - args_dict["host"] = manager.host_ip # 发送注册信息 - regist_json = { - "node_id": manager.args.pd_node_id, - "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", - "mode": manager.pd_mode.value, - "start_args": args_dict, - } + regist_json = _build_pd_registration_info(manager, pd_master_obj) await websocket.send(json.dumps(regist_json)) logger.info(f"Sent registration JSON: {regist_json}") # 转发任务 - forwarding_tokens_task = asyncio.create_task(_up_tokens_to_pd_master(forwarding_queue, websocket)) - heartbeat_task = asyncio.create_task(_send_heartbeat_to_pd_master(websocket)) + forwarding_tokens_task = asyncio.create_task( + _up_tokens_to_pd_master(forwarding_queue, websocket, pd_master_obj.node_id) + ) + heartbeat_task = asyncio.create_task( + _send_heartbeat_to_pd_master(manager, websocket, pd_master_obj.node_id) + ) group_req_id_to_event: Dict[int, asyncio.Event] = weakref.WeakValueDictionary() # 接收 pd master 发来的请求,并推理后,将生成的token转发回pd master。 @@ -264,24 +313,41 @@ async def _pd_process_generate( # 转发token的task -async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): +async def _up_tokens_to_pd_master( + forwarding_queue: AsyncQueue, + websocket: ClientConnection, + pd_master_node_id: int, +): + """批量向 PD Master 转发生成结果和最新负载。""" while True: handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: - load_info: dict = _get_load_info() + load_info: dict = _get_load_info(pd_master_node_id) await websocket.send(pickle.dumps((ObjType.TOKEN_PACKS, handle_list, load_info))) -async def _send_heartbeat_to_pd_master(websocket: ClientConnection): +async def _send_heartbeat_to_pd_master( + manager: HttpServerManager, + websocket: ClientConnection, + pd_master_node_id: int, +): + """定时或在成员变化时向 PD Master 上报心跳。""" heartbeat_interval_seconds = 15 + membership_changed = manager.pd_master_membership_changed while True: - await websocket.send(pickle.dumps((ObjType.HEARTBEAT,))) - await asyncio.sleep(heartbeat_interval_seconds) + membership_changed.clear() + await websocket.send(pickle.dumps((ObjType.HEARTBEAT, _get_load_info(pd_master_node_id)))) + try: + # Master 集合变化时立即重报份额,缩短新旧容量租约并存的窗口。 + await asyncio.wait_for(membership_changed.wait(), timeout=heartbeat_interval_seconds) + except asyncio.TimeoutError: + pass # 获取节点负载信息 -def _get_load_info() -> dict: +def _get_load_info(pd_master_node_id: int) -> dict: + """汇总当前 Master 对应的容量和节点负载。""" from lightllm.server.api_http import g_objs @@ -295,8 +361,11 @@ def _get_load_info() -> dict: float(g_objs.shared_token_load.get_dynamic_max_load(dp_index)) for dp_index in range(dp_size_in_node) ] mean_node_load = sum(current_load) / len(current_load) + pd_master_ids = getattr(g_objs.httpserver_manager, "pd_master_ids", (pd_master_node_id,)) load_info = { "total_token_usage_rate": mean_node_load, "client_ip_port": f"{g_objs.httpserver_manager.host_ip}:{get_shm_port_args().port}", + "capacity_share": _allocate_capacity_share(args.running_max_req_size, pd_master_ids, pd_master_node_id), + "capacity_epoch": getattr(g_objs.httpserver_manager, "pd_master_capacity_epoch", 0), } return load_info diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py new file mode 100644 index 000000000..bebd02686 --- /dev/null +++ b/lightllm/server/httpserver_for_pd_master/admission.py @@ -0,0 +1,716 @@ +from __future__ import annotations + +import asyncio +import time +from collections import OrderedDict, deque +from dataclasses import dataclass, replace +from enum import IntEnum +from typing import Callable, Deque, Dict, Optional + +from lightllm.utils.error_utils import ServerBusyError + + +class AdmissionPriority(IntEnum): + """请求的业务优先级;数值越大,获得服务的权重越高。""" + + COLD = 0 + PROBABLE_CACHE_HIT = 1 + CONTINUATION = 2 + + +@dataclass(frozen=True, slots=True) +class AdmissionPolicy: + """PD Master 内部准入策略。 + + 这些值表达产品层面的等待与公平策略,因此集中在一个对象中,不散落到 + 调度控制流里。容量本身由已注册的 Decode 节点动态提供。 + """ + + continuation_weight: int = 8 + probable_cache_hit_weight: int = 3 + cold_weight: int = 1 + continuation_max_wait_seconds: float = 30.0 + probable_cache_hit_max_wait_seconds: float = 15.0 + cold_max_wait_seconds: float = 5.0 + waiting_decode_waves: int = 1 + probable_cache_hit_threshold: float = 0.5 + active_session_ttl_seconds: float = 30 * 60.0 + max_tracked_sessions: int = 100_000 + + def __post_init__(self) -> None: + """校验准入策略中的权重、超时和容量参数。""" + if ( + min( + self.continuation_weight, + self.probable_cache_hit_weight, + self.cold_weight, + ) + < 1 + ): + raise ValueError("admission weights must be positive") + if ( + min( + self.continuation_max_wait_seconds, + self.probable_cache_hit_max_wait_seconds, + self.cold_max_wait_seconds, + ) + <= 0 + ): + raise ValueError("admission wait timeouts must be positive") + if self.waiting_decode_waves < 1: + raise ValueError("waiting_decode_waves must be positive") + if not 0.0 <= self.probable_cache_hit_threshold <= 1.0: + raise ValueError("probable_cache_hit_threshold must be between zero and one") + if self.active_session_ttl_seconds <= 0: + raise ValueError("active_session_ttl_seconds must be positive") + if self.max_tracked_sessions < 1: + raise ValueError("max_tracked_sessions must be positive") + + def weight(self, priority: AdmissionPriority) -> int: + """返回指定优先级在轮转调度中的权重。""" + if priority == AdmissionPriority.CONTINUATION: + return self.continuation_weight + if priority == AdmissionPriority.PROBABLE_CACHE_HIT: + return self.probable_cache_hit_weight + return self.cold_weight + + def max_wait_seconds(self, priority: AdmissionPriority) -> float: + """返回指定优先级允许的最长排队时间。""" + if priority == AdmissionPriority.CONTINUATION: + return self.continuation_max_wait_seconds + if priority == AdmissionPriority.PROBABLE_CACHE_HIT: + return self.probable_cache_hit_max_wait_seconds + return self.cold_max_wait_seconds + + +@dataclass(frozen=True, slots=True) +class AdmissionRequest: + session_key: Optional[str] + priority: AdmissionPriority + decode_slots: int = 1 + + def __post_init__(self) -> None: + """校验请求需要原子获取的 Decode 槽位数。""" + if self.decode_slots < 1: + raise ValueError("decode_slots must be positive") + + +class SessionTracker: + """只把服务端已经成功观察过的 Session 视为连续会话。""" + + def __init__( + self, + ttl_seconds: float, + max_sessions: int, + clock: Callable[[], float] = time.monotonic, + ) -> None: + """初始化带 TTL 和数量上限的 Session 记录器。""" + self.ttl_seconds = ttl_seconds + self.max_sessions = max_sessions + self._clock = clock + self._last_success: OrderedDict[str, float] = OrderedDict() + + def is_continuation(self, session_key: Optional[str]) -> bool: + """判断 Session 是否在有效期内成功返回过结果。""" + if not session_key: + return False + now = self._clock() + last_success = self._last_success.get(session_key) + if last_success is None: + return False + if now - last_success > self.ttl_seconds: + self._last_success.pop(session_key, None) + return False + self._last_success.move_to_end(session_key) + return True + + def mark_success(self, session_key: Optional[str]) -> None: + """记录 Session 最近一次成功返回结果的时间。""" + if not session_key: + return + self._last_success[session_key] = self._clock() + self._last_success.move_to_end(session_key) + while len(self._last_success) > self.max_sessions: + self._last_success.popitem(last=False) + + +@dataclass(slots=True) +class _WaitingRequest: + request: AdmissionRequest + enqueue_time: float + deadline: float + deadline_changed: asyncio.Event + future: asyncio.Future + + +class AdmissionLease: + """一次已经获得的 Decode 容量租约。""" + + def __init__( + self, + controller: "PDAdmissionController", + request: AdmissionRequest, + waited_seconds: float, + ) -> None: + """保存本次租约占用的 Decode 槽位。""" + self._controller = controller + self.request = request + self.waited_seconds = waited_seconds + self._released = False + + async def __aenter__(self) -> "AdmissionLease": + """进入异步上下文并返回当前租约。""" + return self + + async def __aexit__(self, _exc_type, _exc, _traceback) -> None: + """退出异步上下文时自动释放租约。""" + self.release() + + def release(self) -> None: + """幂等释放本次占用的准入槽位。""" + if self._released: + return + self._released = True + self._controller._release(self) + + +class PDAdmissionController: + """在请求派发到 P/D 节点之前提供有界、可取消的公平等待队列。""" + + _PRIORITY_ORDER = ( + AdmissionPriority.CONTINUATION, + AdmissionPriority.PROBABLE_CACHE_HIT, + AdmissionPriority.COLD, + ) + + def __init__( + self, + decode_capacity_provider: Callable[[], int], + policy: Optional[AdmissionPolicy] = None, + clock: Callable[[], float] = time.monotonic, + state_change_callback: Optional[Callable[["PDAdmissionController"], None]] = None, + ) -> None: + """初始化容量提供器、优先级队列和调度状态。""" + self.policy = policy or AdmissionPolicy() + self._decode_capacity_provider = decode_capacity_provider + self._clock = clock + self._state_change_callback = state_change_callback + self._active_slots = 0 + self._active_sessions = set() + self._queues: Dict[AdmissionPriority, Deque[_WaitingRequest]] = { + priority: deque() for priority in self._PRIORITY_ORDER + } + self._session_queues: Dict[str, Deque[_WaitingRequest]] = {} + self._queued_slots = 0 + self._deficits: Dict[AdmissionPriority, int] = {priority: 0 for priority in self._PRIORITY_ORDER} + self._priority_index = 0 + self._priority_visit_started = False + self._blocked_waiter: Optional[_WaitingRequest] = None + self._backfilled_slots = 0 + self._backfill_limit = 0 + self._reservation_active = False + + @property + def active_slots(self) -> int: + """返回当前已经发放的 Decode 槽位数。""" + return self._active_slots + + @property + def queued_slots(self) -> int: + """返回等待队列中的 Decode 槽位总数。""" + return self._queued_slots + + @property + def queued_request_count(self) -> int: + """返回三个优先级队列中的请求总数。""" + return sum(len(queue) for queue in self._queues.values()) + + def _capacity(self) -> int: + """读取并规范化当前可用的 Decode 容量。""" + return max(0, int(self._decode_capacity_provider())) + + async def acquire(self, request: AdmissionRequest) -> AdmissionLease: + """立即发放租约或等待队列调度后再返回租约。""" + capacity = self._capacity() + if capacity <= 0 or request.decode_slots > capacity: + raise ServerBusyError("PD decode capacity is unavailable") + + if self.queued_request_count == 0 and self._can_activate(request, capacity): + lease = self._activate(request, waited_seconds=0.0) + self._notify_state_change() + return lease + + idle_fill_lease = self._try_activate_idle_fill(request, capacity) + if idle_fill_lease is not None: + self._notify_state_change() + return idle_fill_lease + + loop = asyncio.get_running_loop() + enqueue_time = self._clock() + waiter = _WaitingRequest( + request=request, + enqueue_time=enqueue_time, + deadline=enqueue_time + self.policy.max_wait_seconds(request.priority), + deadline_changed=asyncio.Event(), + future=loop.create_future(), + ) + + if not self._make_queue_room(waiter): + raise ServerBusyError("PD master admission queue is full") + + self._enqueue(waiter) + self._drain() + + try: + return await self._wait_for_lease(waiter) + except asyncio.TimeoutError as exc: + lease = self._cancel_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise ServerBusyError("PD master admission queue wait timed out") from exc + except BaseException: + lease = self._cancel_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise + + def on_capacity_change(self) -> None: + """Decode 容量变化后重置临时公平状态并重新驱动队列。""" + self._clear_backfill_state() + self._reset_deficits() + self._drain() + + def promote_session(self, session_key: Optional[str]) -> None: + """把同一 Session 尚未派发的请求提升为连续会话优先级。""" + if not session_key: + return + session_queue = self._session_queues.get(session_key) + if not session_queue: + return + + if self._blocked_waiter in session_queue: + self._clear_backfill_state(reset_deficits=True) + for waiter in tuple(session_queue): + old_priority = waiter.request.priority + if old_priority == AdmissionPriority.CONTINUATION: + continue + self._queues[old_priority].remove(waiter) + waiter.request = replace(waiter.request, priority=AdmissionPriority.CONTINUATION) + waiter.deadline = max( + waiter.deadline, + waiter.enqueue_time + self.policy.continuation_max_wait_seconds, + ) + waiter.deadline_changed.set() + self._queues[AdmissionPriority.CONTINUATION].append(waiter) + if not self._queues[old_priority]: + self._deficits[old_priority] = 0 + self._drain() + + def _can_activate(self, request: AdmissionRequest, capacity: Optional[int] = None) -> bool: + """检查 Decode 总容量和 Session 串行约束。""" + if capacity is None: + capacity = self._capacity() + if self._active_slots + request.decode_slots > capacity: + return False + return request.session_key is None or request.session_key not in self._active_sessions + + def _try_activate_idle_fill( + self, + request: AdmissionRequest, + capacity: int, + ) -> Optional[AdmissionLease]: + """在满队列拒绝前,用当前唯一可运行的新请求填充空槽。 + + 只覆盖两种不会越过可运行旧请求的场景:现有队列全部受 + Session 串行约束,或已保护 gang 仍在一波 bounded backfill 预算内。 + """ + if not self._can_activate(request, capacity): + return None + if request.session_key is not None and request.session_key in self._session_queues: + return None + + blocked = self._blocked_waiter + if blocked is None: + has_grantable_waiter = any( + not waiter.future.done() and self._session_is_grantable(waiter) + for queue in self._queues.values() + for waiter in queue + ) + if has_grantable_waiter: + return None + return self._activate(request, waited_seconds=0.0) + + available_slots = capacity - self._active_slots + if ( + self._reservation_active + or blocked.future.done() + or not self._session_is_grantable(blocked) + or self._has_fitting_backfill(blocked, available_slots) + ): + return None + + remaining_backfill = self._backfill_limit - self._backfilled_slots + if request.decode_slots > remaining_backfill: + return None + + lease = self._activate(request, waited_seconds=0.0) + self._backfilled_slots += request.decode_slots + if self._backfilled_slots >= self._backfill_limit: + self._reservation_active = True + return lease + + def _activate(self, request: AdmissionRequest, waited_seconds: float) -> AdmissionLease: + """占用所需槽位并创建对应的准入租约。""" + self._active_slots += request.decode_slots + if request.session_key is not None: + self._active_sessions.add(request.session_key) + return AdmissionLease(self, request, waited_seconds) + + def _release(self, lease: AdmissionLease) -> None: + """归还租约槽位并继续调度等待请求。""" + request = lease.request + self._active_slots -= request.decode_slots + if self._active_slots < 0: + raise RuntimeError("PD admission active slot count became negative") + if request.session_key is not None: + self._active_sessions.discard(request.session_key) + self._drain() + + def _waiting_capacity(self) -> int: + """返回等待队列允许容纳的槽位总数。""" + return self._capacity() * self.policy.waiting_decode_waves + + def _make_queue_room(self, incoming: _WaitingRequest) -> bool: + """必要时淘汰低优先级请求,为新请求腾出队列空间。""" + waiting_capacity = self._waiting_capacity() + if incoming.request.decode_slots > waiting_capacity: + return False + + required_slots = self._queued_slots + incoming.request.decode_slots - waiting_capacity + if required_slots <= 0: + return True + + victims = [] + released_slots = 0 + for priority in reversed(self._PRIORITY_ORDER): + if priority >= incoming.request.priority: + continue + for waiter in reversed(self._queues[priority]): + victims.append(waiter) + released_slots += waiter.request.decode_slots + if released_slots >= required_slots: + break + if released_slots >= required_slots: + break + + if released_slots < required_slots: + return False + + for victim in victims: + self._remove_waiter(victim) + if not victim.future.done(): + victim.future.set_exception(ServerBusyError("Superseded by a higher-priority queued request")) + return True + + def _enqueue(self, waiter: _WaitingRequest) -> None: + """把等待项加入优先级队列和 Session 队列。""" + self._queues[waiter.request.priority].append(waiter) + self._queued_slots += waiter.request.decode_slots + if waiter.request.session_key is not None: + self._session_queues.setdefault(waiter.request.session_key, deque()).append(waiter) + + def _remove_waiter(self, waiter: _WaitingRequest) -> bool: + """从所有索引中移除等待项并归还排队槽位。""" + try: + self._queues[waiter.request.priority].remove(waiter) + except ValueError: + return False + + self._queued_slots -= waiter.request.decode_slots + session_key = waiter.request.session_key + if session_key is not None: + session_queue = self._session_queues[session_key] + session_queue.remove(waiter) + if not session_queue: + self._session_queues.pop(session_key, None) + if waiter is self._blocked_waiter: + # 非正常移除(取消、超时、替换、缩容)放弃已经预扣的 gang + # 服务机会;重置 DRR 状态比跨优先级退款更安全。 + self._clear_backfill_state(reset_deficits=True) + if not self._queues[waiter.request.priority]: + self._deficits[waiter.request.priority] = 0 + return True + + def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[AdmissionLease]: + """取消等待项,或取回已并发发放的租约用于释放。""" + if self._remove_waiter(waiter): + waiter.future.cancel() + self._drain() + return None + if waiter.future.done() and not waiter.future.cancelled(): + try: + result = waiter.future.result() + except BaseException: + return None + if isinstance(result, AdmissionLease): + return result + return None + + async def _wait_for_lease(self, waiter: _WaitingRequest) -> AdmissionLease: + """等待租约、优先级提升后的新截止时间或超时。""" + while True: + deadline_changed_task = asyncio.create_task(waiter.deadline_changed.wait()) + try: + done, _ = await asyncio.wait( + (waiter.future, deadline_changed_task), + timeout=max(0.0, waiter.deadline - self._clock()), + return_when=asyncio.FIRST_COMPLETED, + ) + finally: + if not deadline_changed_task.done(): + deadline_changed_task.cancel() + + if waiter.future in done: + return waiter.future.result() + if deadline_changed_task in done: + waiter.deadline_changed.clear() + continue + raise asyncio.TimeoutError + + def _session_is_grantable(self, waiter: _WaitingRequest) -> bool: + """判断等待项是否满足同 Session 串行和 FIFO 约束。""" + session_key = waiter.request.session_key + if session_key is None: + return True + if session_key in self._active_sessions: + return False + return self._session_queues[session_key][0] is waiter + + def _first_grantable( + self, + priority: AdmissionPriority, + excluded: Optional[_WaitingRequest] = None, + available_slots: Optional[int] = None, + ) -> Optional[_WaitingRequest]: + """按类内 FIFO 返回第一个满足 Session 和可选槽位约束的等待项。""" + for waiter in self._queues[priority]: + if waiter is excluded: + continue + if not self._session_is_grantable(waiter): + continue + if available_slots is not None and waiter.request.decode_slots > available_slots: + continue + return waiter + return None + + def _advance_priority(self) -> None: + """结束当前 DRR 类访问并移到下一优先级。""" + self._priority_index = (self._priority_index + 1) % len(self._PRIORITY_ORDER) + self._priority_visit_started = False + + def _reset_deficits(self) -> None: + """清空按 Decode 槽位计费的 DRR 临时信用。""" + for priority in self._PRIORITY_ORDER: + self._deficits[priority] = 0 + self._priority_index = 0 + self._priority_visit_started = False + + def _select_weighted( + self, + capacity: int, + excluded: Optional[_WaitingRequest] = None, + available_slots: Optional[int] = None, + ) -> Optional[_WaitingRequest]: + """用按槽位计费的 deficit round-robin 选择一个等待项。""" + candidates = { + priority: waiter + for priority in self._PRIORITY_ORDER + if ( + waiter := self._first_grantable( + priority, + excluded=excluded, + available_slots=available_slots, + ) + ) + is not None + } + for priority in self._PRIORITY_ORDER: + # 空类和暂时全部受 Session 串行约束的类都不能积攒无限信用。 + if self._first_grantable(priority) is None: + self._deficits[priority] = 0 + if not candidates: + return None + + only_priority = next(iter(candidates)) if len(candidates) == 1 else None + while True: + priority = self._PRIORITY_ORDER[self._priority_index] + waiter = candidates.get(priority) + if waiter is None: + self._advance_priority() + continue + + if not self._priority_visit_started: + self._deficits[priority] += self.policy.weight(priority) + self._priority_visit_started = True + + # 单一活跃类必须保持 work-conserving;直接补足若干轮 quantum, + # 避免大 gang 仅因 DRR 信用暂时不足而留下 Decode 空槽。 + if only_priority == priority and self._deficits[priority] < waiter.request.decode_slots: + quantum = self.policy.weight(priority) + missing = waiter.request.decode_slots - self._deficits[priority] + visits = (missing + quantum - 1) // quantum + self._deficits[priority] += visits * quantum + + if waiter.request.decode_slots <= self._deficits[priority]: + # 在选择点统一按 choice 槽位扣费。即使 gang 暂时因物理空槽 + # 不足进入 backfill,它的 DRR 服务机会也已经被完整计费。 + self._deficits[priority] -= waiter.request.decode_slots + return waiter + self._advance_priority() + + def _clear_backfill_state(self, reset_deficits: bool = False) -> None: + """清除 gang backfill 或 reservation 的全部临时状态。""" + self._blocked_waiter = None + self._backfilled_slots = 0 + self._backfill_limit = 0 + self._reservation_active = False + if reset_deficits: + self._reset_deficits() + + def _start_backfill(self, waiter: _WaitingRequest, capacity: int) -> None: + """为仅受当前可用槽位阻塞的 gang 启动一波有限 backfill。""" + self._blocked_waiter = waiter + self._backfilled_slots = 0 + self._backfill_limit = capacity + self._reservation_active = False + + def _fail_oversized_waiters(self, capacity: int) -> None: + """容量缩小时失败掉已经不可能原子获得所需槽位的等待项。""" + for priority in self._PRIORITY_ORDER: + for waiter in tuple(self._queues[priority]): + if waiter.request.decode_slots <= capacity: + continue + if self._remove_waiter(waiter) and not waiter.future.done(): + waiter.future.set_exception(ServerBusyError("PD decode capacity fell below queued request size")) + + def _trim_queue_to_capacity(self, capacity: int) -> None: + """容量缩小时按低优先级、同级最新顺序恢复等待队列上限。""" + waiting_capacity = capacity * self.policy.waiting_decode_waves + slots_to_remove = self._queued_slots - waiting_capacity + if slots_to_remove <= 0: + return + + victims = [] + removed_slots = 0 + for priority in reversed(self._PRIORITY_ORDER): + for waiter in reversed(self._queues[priority]): + victims.append(waiter) + removed_slots += waiter.request.decode_slots + if removed_slots >= slots_to_remove: + break + if removed_slots >= slots_to_remove: + break + + for victim in victims: + if self._remove_waiter(victim) and not victim.future.done(): + victim.future.set_exception(ServerBusyError("PD master admission queue capacity shrank")) + + def _grant_waiter(self, waiter: _WaitingRequest) -> bool: + """从队列移除已由 DRR 计费的等待项并原子发放租约。""" + if waiter.future.done(): + self._remove_waiter(waiter) + self._reset_deficits() + return False + + priority = waiter.request.priority + if waiter is self._blocked_waiter: + # 正常兑现 reservation 时保留选择点已经完成的 DRR 扣费。 + self._clear_backfill_state() + if not self._remove_waiter(waiter): + self._reset_deficits() + return False + if not self._queues[priority]: + self._deficits[priority] = 0 + + lease = self._activate( + waiter.request, + waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), + ) + waiter.future.set_result(lease) + return True + + def _has_fitting_backfill(self, blocked: _WaitingRequest, available_slots: int) -> bool: + """判断是否有不含被保护 gang 的请求当前可以填充空槽。""" + return any( + self._first_grantable( + priority, + excluded=blocked, + available_slots=available_slots, + ) + is not None + for priority in self._PRIORITY_ORDER + ) + + def _drain(self) -> None: + """持续发放租约,并为受空槽碎片阻塞的 gang 提供有限 backfill。""" + capacity = self._capacity() + self._fail_oversized_waiters(capacity) + self._trim_queue_to_capacity(capacity) + + while self._active_slots < capacity: + available_slots = capacity - self._active_slots + blocked = self._blocked_waiter + if blocked is not None: + if blocked.future.done(): + self._remove_waiter(blocked) + continue + if not self._session_is_grantable(blocked): + # Session 阻塞不是容量碎片,不能借此获得全局 reservation。 + self._clear_backfill_state(reset_deficits=True) + continue + if blocked.request.decode_slots <= available_slots: + self._grant_waiter(blocked) + continue + if self._reservation_active: + break + + remaining_backfill = self._backfill_limit - self._backfilled_slots + if remaining_backfill <= 0: + self._reservation_active = True + break + backfill_slots = min(available_slots, remaining_backfill) + waiter = self._select_weighted( + capacity, + excluded=blocked, + available_slots=backfill_slots, + ) + if waiter is None: + # 没有可填当前空槽的请求时保留 backfill 机会;稍后到达的小请求 + # 仍可使用本波预算。只有存在物理上可填、但会越过预算的请求时 + # 才立即转入 reservation。 + if self._has_fitting_backfill(blocked, available_slots): + self._reservation_active = True + break + granted_slots = waiter.request.decode_slots + if not self._grant_waiter(waiter): + continue + self._backfilled_slots += granted_slots + if self._backfilled_slots >= self._backfill_limit: + self._reservation_active = True + continue + + waiter = self._select_weighted(capacity) + if waiter is None: + break + if waiter.request.decode_slots > available_slots: + # 此处 waiter 已满足 Session 约束且不超过总容量,唯一阻塞原因 + # 是当前空槽不足,因此可以安全启动 bounded backfill。 + self._start_backfill(waiter, capacity) + continue + self._grant_waiter(waiter) + self._notify_state_change() + + def _notify_state_change(self) -> None: + """通知外部记录最新的准入状态。""" + if self._state_change_callback is not None: + self._state_change_callback(self) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index b422bf770..fba20821d 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -4,6 +4,7 @@ import uvloop import time import datetime +import math import ujson as json import pickle import httpx @@ -12,7 +13,14 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus -from ..pd_io_struct import PD_Client_Obj, PDUpKVStatus, ObjType, PDDecodeNodeInfo +from ..pd_io_struct import ( + PD_Client_Obj, + PDUpKVStatus, + ObjType, + PDDecodeNodeInfo, + PD_MASTER_CAPACITY_EPOCH_KEY, + PD_MASTER_CAPACITY_SHARE_KEY, +) from lightllm.server.core.objs import SamplingParams, StartArgs from ..multimodal_params import MultimodalParams from ..tokenizer import get_tokenizer @@ -26,6 +34,13 @@ 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 .admission import ( + AdmissionPolicy, + AdmissionPriority, + AdmissionRequest, + PDAdmissionController, + SessionTracker, +) logger = init_logger(__name__) @@ -43,6 +58,20 @@ def __init__( self.pd_manager = PDManager(args) + self.admission_policy = AdmissionPolicy() + self.session_tracker = SessionTracker( + ttl_seconds=self.admission_policy.active_session_ttl_seconds, + max_sessions=self.admission_policy.max_tracked_sessions, + ) + self._last_admission_metric_values: Dict[str, int] = {} + self._pending_admission_metric_values: Optional[Dict[str, int]] = None + self._admission_metric_flush_scheduled = False + self.admission_controller = PDAdmissionController( + decode_capacity_provider=self.pd_manager.get_decode_capacity, + policy=self.admission_policy, + state_change_callback=self._record_admission_state, + ) + self.req_id_to_out_inf: Dict[int, ReqStatus] = {} self.infos_queues = None # 这个需要延迟初始化,否则使用的loop不对 self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) @@ -75,10 +104,12 @@ def is_healthy(self): async def register_pd(self, pd_info_json, websocket): self.pd_manager.register_pd(pd_info_json, websocket) + self.admission_controller.on_capacity_change() return async def remove_pd(self, pd_info_json): self.pd_manager.remove_pd(pd_info_json) + self.admission_controller.on_capacity_change() return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -91,6 +122,44 @@ async def update_req_status(self, upkv_status: PDUpKVStatus): pass return + def update_node_load_info(self, load_info: Optional[dict]) -> None: + """更新节点遥测;仅 Decode 容量租约变化时重新驱动准入队列。""" + if self.pd_manager.update_node_load_info(load_info): + self.admission_controller.on_capacity_change() + + def _record_admission_state(self, controller: PDAdmissionController) -> None: + """合并同一事件循环周期内的状态变化,避免重复发送 gauge RPC。""" + self._pending_admission_metric_values = { + "lightllm_pd_master_admission_queue_size": controller.queued_request_count, + "lightllm_pd_master_admission_active_slots": controller.active_slots, + "lightllm_pd_master_admission_decode_capacity": self.pd_manager.get_decode_capacity(), + } + if self._admission_metric_flush_scheduled: + return + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + self._flush_admission_state_metrics() + return + + self._admission_metric_flush_scheduled = True + loop.call_soon(self._flush_admission_state_metrics) + + def _flush_admission_state_metrics(self) -> None: + """只发送相较于上次上报实际发生变化的 admission gauge。""" + self._admission_metric_flush_scheduled = False + metric_values = self._pending_admission_metric_values + self._pending_admission_metric_values = None + if metric_values is None: + return + + for name, value in metric_values.items(): + if self._last_admission_metric_values.get(name) == value: + continue + self.metric_client.gauge_set(name, value) + self._last_admission_metric_values[name] = value + def tokens(self, prompt, multimodal_params, samping_params: SamplingParams, kwargs=None): kwargs = {} if kwargs is None else kwargs prompt_ids = self.tokenizer.encode(prompt, None, **kwargs) @@ -129,21 +198,69 @@ 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: - self.latest_success_infer_time = time.time() + admission_lease = None + running_request_registered = False + session_key = self._get_session_key(request) try: + if not self.args.disable_pd_master_decode_capacity_limit: + admission_request = self._build_admission_request(prompt, sampling_params, session_key) + admission_lease = await self.admission_controller.acquire(admission_request) + self.metric_client.histogram_observe( + "lightllm_request_queue_duration_bucket", admission_lease.waited_seconds + ) + if admission_lease.waited_seconds > 0: + # 排队期间节点和缓存内容可能变化;派发前重新匹配,避免复用过期快照。 + self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) + + was_idle = self.running_request_count == 0 + self.running_request_count += 1 + running_request_registered = True + if was_idle: + self.latest_success_infer_time = time.time() async with aclosing(self._generate(prompt, sampling_params, multimodal_params, request)) as generator: async for result in generator: + if session_key is not None and not self.session_tracker.is_continuation(session_key): + self.session_tracker.mark_success(session_key) + self.admission_controller.promote_session(session_key) yield result finally: - self.running_request_count -= 1 + if running_request_registered: + self.running_request_count -= 1 + if admission_lease is not None: + admission_lease.release() + + def _get_session_key(self, request: Optional[Request]) -> Optional[str]: + """从请求头中提取规范化的 Session 标识。""" + if request is None: + return None + session_key = request.headers.get("X-Session-Id", "").strip() + return session_key or None + + def _build_admission_request( + self, + prompt: Union[str, List[int]], + sampling_params: Optional[SamplingParams], + session_key: Optional[str], + ) -> AdmissionRequest: + """根据会话、缓存估算和 choice 数构造准入请求。""" + estimated_cache_hit_rate = self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) + if estimated_cache_hit_rate is None or not math.isfinite(estimated_cache_hit_rate): + estimated_cache_hit_rate = 0.0 + estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) + + if self.session_tracker.is_continuation(session_key): + priority = AdmissionPriority.CONTINUATION + elif estimated_cache_hit_rate >= self.admission_policy.probable_cache_hit_threshold: + priority = AdmissionPriority.PROBABLE_CACHE_HIT + else: + priority = AdmissionPriority.COLD + + decode_slots = max(1, int(getattr(sampling_params, "n", 1) or 1)) + return AdmissionRequest( + session_key=session_key, + priority=priority, + decode_slots=decode_slots, + ) async def _generate( self, @@ -642,7 +759,7 @@ async def handle_loop(self): for obj in objs: if obj[0] == ObjType.TOKEN_PACKS: token_list, node_load_info = obj[1], obj[2] - self.pd_manager.update_node_load_info(node_load_info) + self.update_node_load_info(node_load_info) for sub_req_id, text, metadata, finish_status in token_list: finish_status: FinishStatus = finish_status @@ -762,6 +879,13 @@ def __init__(self, args: StartArgs): self.selector = create_selector(args.select_p_d_node_strategy, self) return + def get_decode_capacity(self) -> int: + """汇总所有 Decode 节点租给当前 Master 的槽位。""" + return sum( + node.capacity_share if node.capacity_share is not None else node.start_args["running_max_req_size"] + for node in self.decode_nodes + ) + def is_pd_nodes_ready(self): prefill_node_count = len(self.prefill_nodes) decode_node_count = len(self.decode_nodes) @@ -813,7 +937,24 @@ async def check_pd_nodes_health(self): return True def register_pd(self, pd_info_json, websocket): - pd_client = PD_Client_Obj(**pd_info_json) + # Capacity metadata lives in reserved start_args keys so newer P/D + # nodes retain the legacy registration schema for older Masters. Keep + # accepting the short-lived top-level form for branch compatibility. + pd_info = dict(pd_info_json) + start_args = pd_info.get("start_args") or {} + capacity_share = pd_info.pop( + "capacity_share", + start_args.get(PD_MASTER_CAPACITY_SHARE_KEY), + ) + capacity_epoch = pd_info.pop( + "capacity_epoch", + start_args.get(PD_MASTER_CAPACITY_EPOCH_KEY, 0), + ) + pd_client = PD_Client_Obj( + **pd_info, + capacity_share=capacity_share, + capacity_epoch=capacity_epoch, + ) client_max_req_total_len = pd_client.start_args["max_req_total_len"] if client_max_req_total_len != self.args.max_req_total_len: logger.error( @@ -853,19 +994,25 @@ def register_pd(self, pd_info_json, websocket): return def remove_pd(self, pd_info_json): - pd_client = PD_Client_Obj(**pd_info_json) + # Disconnect cleanup only needs the stable legacy identity fields; do + # not reconstruct PD_Client_Obj from a possibly newer registration. + client_ip_port = pd_info_json["client_ip_port"] + mode = pd_info_json.get("mode", "unknown") - self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) - self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] - self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] + self.url_to_pd_nodes.pop(client_ip_port, None) + self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != client_ip_port] + self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != client_ip_port] self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) - logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed") + logger.info(f"mode: {mode} url: {client_ip_port} removed") return - def update_node_load_info(self, load_info: Optional[dict]): - """更新节点负载信息 + def update_node_load_info(self, load_info: Optional[dict]) -> bool: + """更新节点负载信息,并返回有效 Decode 容量份额是否发生变化。 + + capacity_epoch 只用于拒绝旧上报;仅 epoch 变化不会重新驱动 admission。 + load_info: 节点负载信息字典,内容格式如下,可以为 None { "total_token_usage_rate": xxxx, @@ -874,14 +1021,28 @@ def update_node_load_info(self, load_info: Optional[dict]): """ try: if load_info is None: - return + return False client_ip_port = load_info["client_ip_port"] - total_token_usage_rate = load_info["total_token_usage_rate"] pd_client = self.url_to_pd_nodes.get(client_ip_port) - pd_client.run_status.total_token_usage_rate = total_token_usage_rate - except BaseException as e: + if pd_client is None: + return False + pd_client.run_status.total_token_usage_rate = load_info["total_token_usage_rate"] + + capacity_epoch = int(load_info.get("capacity_epoch", pd_client.capacity_epoch)) + if capacity_epoch >= pd_client.capacity_epoch: + fallback_capacity = ( + pd_client.capacity_share + if pd_client.capacity_share is not None + else pd_client.start_args["running_max_req_size"] + ) + capacity_share = max(0, int(load_info.get("capacity_share", fallback_capacity))) + capacity_changed = fallback_capacity != capacity_share + pd_client.capacity_share = capacity_share + pd_client.capacity_epoch = capacity_epoch + return pd_client.mode == "decode" and capacity_changed + except Exception as e: logger.warning(f"udpate node load info failed, load_info: {load_info} error: {str(e)}") - return + return False def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams 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 adf911bc2..6bb9bdf6e 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 @@ -19,13 +19,14 @@ from __future__ import annotations +from contextvars import ContextVar from dataclasses import dataclass from typing import List, Optional from lightllm.server.pd_io_struct import PD_Client_Obj from lightllm.utils.log_utils import init_logger -from .prompt_cache_tree import PromptCacheTree +from .prompt_cache_tree import PromptCacheMatchResult, PromptCacheTree logger = init_logger(__name__) @@ -56,6 +57,18 @@ class CacheAwareConfig: recursion_limit: int = 4000 +@dataclass(frozen=True, slots=True) +class _PromptCacheMatchContext: + policy: "CacheAwarePolicy" + request_text: str + match_result: PromptCacheMatchResult + + +_prompt_cache_match_context: ContextVar[Optional[_PromptCacheMatchContext]] = ContextVar( + "prompt_cache_match_context", default=None +) + + class BalanceRelThresholdController: """根据最近请求的 prompt cache 命中率动态调整负载均衡阈值。""" @@ -156,9 +169,36 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.balance_rel_threshold_controller.append(cache_hit_rate) self.balance_rel_threshold_controller.update_config(self.config) + def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: str) -> float: + """估算可复用缓存的命中率,并在当前异步请求上下文中保留匹配结果。""" + if not workers or not request_text: + _prompt_cache_match_context.set(None) + return 0.0 + + result = self.prompt_cache_tree.prefix_match(request_text) + _prompt_cache_match_context.set( + _PromptCacheMatchContext(policy=self, request_text=request_text, match_result=result) + ) + if result.prefill_node is None or not any(worker.client_ip_port == result.prefill_node for worker in workers): + return 0.0 + if result.input_char_count == 0: + return 0.0 + return min(max(result.matched_char_count / result.input_char_count, 0.0), 1.0) + + def _match_prompt_cache(self, request_text: str) -> PromptCacheMatchResult: + """优先复用当前请求保存的前缀匹配结果。""" + match_context = _prompt_cache_match_context.get() + # 准入和选点共享同一个提示词对象。通过对象身份判断可以避免比较可能很长的字符串, + # ContextVar 则会将快照安全地传递给每个 n-choice 子任务。 + if match_context is not None and match_context.policy is self and match_context.request_text is request_text: + _prompt_cache_match_context.set(None) + return match_context.match_result + _prompt_cache_match_context.set(None) + return self.prompt_cache_tree.prefix_match(request_text) + def _get_cache_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """在指定候选节点中返回达到匹配阈值的 cache 节点。""" - result = self.prompt_cache_tree.prefix_match(request_text) + result = self._match_prompt_cache(request_text) match_rate = 0.0 if result.input_char_count == 0 else result.matched_char_count / result.input_char_count logger.info( 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 5474806b7..eb0f9b570 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 @@ -1,5 +1,5 @@ import random -from typing import Union, List, Tuple, Dict +from typing import Union, List, Tuple, Dict, Optional from lightllm.server.pd_io_struct import PD_Client_Obj from lightllm.server.core.objs import SamplingParams from lightllm.server.multimodal_params import MultimodalParams @@ -30,6 +30,10 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: """记录推理侧返回的 prompt cache 命中率;非 cache-aware 策略无需处理。""" return + def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: + """当选择器无法估算可复用的提示词缓存时返回 None。""" + return None + class RandomSelector(PDSelector): """随机选择器""" @@ -100,3 +104,9 @@ def select_p_d_node( def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.policy.record_prompt_cache_hit_rate(cache_hit_rate) + + def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: + """返回 cache-aware 策略对当前提示词的命中估算。""" + if not isinstance(prompt, str): + return 0.0 + return self.policy.estimate_cache_hit_rate(self.prefill_nodes, prompt) diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d42462c3..5814b11f6 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,6 +32,9 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", + "lightllm_pd_master_admission_queue_size": "Number of requests waiting at the PD master admission queue", + "lightllm_pd_master_admission_active_slots": "Number of decode slots leased by the PD master", + "lightllm_pd_master_admission_decode_capacity": "Decode slot capacity currently assigned to the PD master", } @@ -111,6 +114,9 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") + self.create_gauge("lightllm_pd_master_admission_queue_size") + self.create_gauge("lightllm_pd_master_admission_active_slots") + self.create_gauge("lightllm_pd_master_admission_decode_capacity") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 8a1e2bd42..f60a470c6 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -11,6 +11,13 @@ logger = init_logger(__name__) +# Keep per-Master capacity metadata inside ``start_args`` so a new P/D node can +# still register with an older Master whose PD_Client_Obj rejects unknown +# top-level fields. New Masters extract these reserved protocol keys. +PD_MASTER_CAPACITY_SHARE_KEY = "__pd_master_capacity_share" +PD_MASTER_CAPACITY_EPOCH_KEY = "__pd_master_capacity_epoch" + + # 节点的行为 class NodeRole(enum.Enum): P = "prefill" @@ -57,6 +64,9 @@ class PD_Client_Obj: start_args: object # 节点的启动参数信息,用于做匹配性的校验,防止运行过程中出现问题。 websocket: WebSocket = None # 用于通信的 websocket 连接对象 run_status: _PD_Client_RunStatus = field(default_factory=_PD_Client_RunStatus) + # 节点租给当前 PD Master 的请求槽位;多 Master 之间的份额互不重叠。 + capacity_share: Optional[int] = None + capacity_epoch: int = 0 # cache-aware 选点用:当前派发到该节点且尚未产出首 token 的 prompt 字符数。 dispatched_prompt_chars: int = 0 # 当前派发到该节点且尚未产出首 token 的请求数。 @@ -67,6 +77,10 @@ def __post_init__(self): error_info = f"""mode must in ["prefill", "decode"], but get {self.mode}""" logger.error(error_info) raise ValueError(error_info) + if self.capacity_share is not None and self.capacity_share < 0: + raise ValueError("capacity_share must be non-negative") + if self.capacity_epoch < 0: + raise ValueError("capacity_epoch must be non-negative") return def to_llm_url(self): 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 a3dc268a9..da07759a9 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -15,6 +15,7 @@ def _make_manager(monkeypatch): ) monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) + mgr.args = SimpleNamespace(disable_pd_master_decode_capacity_limit=True) mgr.running_request_count = 0 counter = [0] diff --git a/unit_tests/server/test_pd_admission.py b/unit_tests/server/test_pd_admission.py new file mode 100644 index 000000000..802e293c0 --- /dev/null +++ b/unit_tests/server/test_pd_admission.py @@ -0,0 +1,627 @@ +import asyncio + +import pytest + +from lightllm.server.httpserver_for_pd_master.admission import ( + AdmissionPolicy, + AdmissionPriority, + AdmissionRequest, + PDAdmissionController, + SessionTracker, +) +from lightllm.utils.error_utils import ServerBusyError + + +def _request( + priority=AdmissionPriority.COLD, + session_key=None, + decode_slots=1, +): + return AdmissionRequest( + session_key=session_key, + priority=priority, + decode_slots=decode_slots, + ) + + +def test_admission_waits_before_granting_more_than_decode_capacity(): + async def run(): + controller = PDAdmissionController(lambda: 1) + first = await controller.acquire(_request()) + second_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + assert controller.active_slots == 1 + assert controller.queued_slots == 1 + assert second_task.done() is False + + first.release() + second = await second_task + assert controller.active_slots == 1 + assert controller.queued_slots == 0 + assert second.waited_seconds >= 0 + second.release() + + asyncio.run(run()) + + +def test_admission_prioritizes_continuations_without_starving_lower_classes(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = [await controller.acquire(_request()) for _ in range(3)] + cold_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) + probable_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) + continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + active[0].release() + continuation = await continuation_task + assert probable_task.done() is False + assert cold_task.done() is False + + active[1].release() + probable = await probable_task + active[2].release() + cold = await cold_task + + continuation.release() + probable.release() + cold.release() + + asyncio.run(run()) + + +def test_decode_capacity_one_still_follows_slot_weighted_drr(): + async def run(): + policy = AdmissionPolicy( + continuation_weight=8, + probable_cache_hit_weight=3, + cold_weight=1, + waiting_decode_waves=12, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + acquired = asyncio.Queue() + + async def acquire(label, priority): + lease = await controller.acquire(_request(priority)) + await acquired.put((label, lease)) + + tasks = [ + asyncio.create_task(acquire(f"continuation-{index}", AdmissionPriority.CONTINUATION)) for index in range(8) + ] + tasks.extend( + asyncio.create_task(acquire(f"probable-{index}", AdmissionPriority.PROBABLE_CACHE_HIT)) + for index in range(3) + ) + tasks.append(asyncio.create_task(acquire("cold-0", AdmissionPriority.COLD))) + await asyncio.sleep(0) + + active.release() + expected = ["continuation"] * 8 + ["probable"] * 3 + ["cold"] + for expected_prefix in expected: + label, lease = await asyncio.wait_for(acquired.get(), timeout=1.0) + assert label.startswith(expected_prefix) + lease.release() + + await asyncio.gather(*tasks) + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_higher_priority_request_can_replace_a_queued_cold_request(): + async def run(): + controller = PDAdmissionController(lambda: 1) + active = await controller.acquire(_request()) + cold_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) + await asyncio.sleep(0) + continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + with pytest.raises(ServerBusyError, match="Superseded"): + await cold_task + + active.release() + continuation = await continuation_task + continuation.release() + + asyncio.run(run()) + + +def test_multi_choice_request_acquires_all_slots_atomically(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = await controller.acquire(_request(decode_slots=2)) + multi_choice_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + await asyncio.sleep(0) + + assert controller.active_slots == 2 + assert multi_choice_task.done() is False + + active.release() + multi_choice = await multi_choice_task + assert controller.active_slots == 2 + multi_choice.release() + + asyncio.run(run()) + + +def test_same_priority_small_request_backfills_a_blocked_gang_without_idle_slots(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = [await controller.acquire(_request()) for _ in range(3)] + multi_choice_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + later_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + active[0].release() + later = await later_task + assert multi_choice_task.done() is False + assert controller.active_slots == 3 + assert all(deficit >= 0 for deficit in controller._deficits.values()) + + active[1].release() + assert multi_choice_task.done() is False + + later.release() + multi_choice = await multi_choice_task + assert controller.active_slots == 3 + multi_choice.release() + active[2].release() + + asyncio.run(run()) + + +def test_blocked_gang_keeps_backfill_open_for_a_later_small_request(): + async def run(): + controller = PDAdmissionController( + lambda: 2, + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + await asyncio.sleep(0) + + active[0].release() + assert gang_task.done() is False + assert controller.active_slots == 1 + + small_task = asyncio.create_task(controller.acquire(_request())) + small = await small_task + assert controller.active_slots == 2 + assert gang_task.done() is False + + small.release() + assert gang_task.done() is False + active[1].release() + gang = await gang_task + gang.release() + + asyncio.run(run()) + + +def test_full_queue_of_session_blocked_waiters_cannot_leave_other_sessions_idle(): + async def run(): + controller = PDAdmissionController(lambda: 4) + active_a = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-a")) + queued_a_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-a"))) + for _ in range(4) + ] + await asyncio.sleep(0) + assert controller.queued_slots == 4 + assert controller.active_slots == 1 + + session_b = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-b")) + assert controller.active_slots == 2 + assert controller.queued_slots == 4 + + session_b.release() + active_a.release() + for task in queued_a_tasks: + lease = await task + lease.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_full_gang_queue_still_accepts_a_fitting_request_within_backfill_budget(): + async def run(): + controller = PDAdmissionController(lambda: 4) + active = await controller.acquire(_request(decode_slots=3)) + gang_tasks = [asyncio.create_task(controller.acquire(_request(decode_slots=2))) for _ in range(2)] + await asyncio.sleep(0) + assert controller.queued_slots == 4 + assert controller.active_slots == 3 + + small = await controller.acquire(_request()) + assert controller.active_slots == 4 + assert controller.queued_slots == 4 + assert controller._backfilled_slots == 1 + + small.release() + active.release() + first_gang = await gang_tasks[0] + second_gang = await gang_tasks[1] + first_gang.release() + second_gang.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_gang_reservation_starts_after_one_backfill_wave_and_blocks_later_priority(): + async def run(): + controller = PDAdmissionController( + lambda: 3, + policy=AdmissionPolicy(waiting_decode_waves=5), + ) + active = [await controller.acquire(_request()) for _ in range(3)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=3))) + backfill_tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(3)] + await asyncio.sleep(0) + + backfills = [] + for active_lease, backfill_task in zip(active, backfill_tasks): + active_lease.release() + backfills.append(await backfill_task) + assert gang_task.done() is False + + later_high_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + backfills[0].release() + await asyncio.sleep(0) + assert later_high_task.done() is False + assert gang_task.done() is False + backfills[1].release() + await asyncio.sleep(0) + assert later_high_task.done() is False + assert gang_task.done() is False + + backfills[2].release() + gang = await gang_task + assert later_high_task.done() is False + + gang.release() + later_high = await later_high_task + later_high.release() + + asyncio.run(run()) + + +def test_same_session_is_fifo_while_other_sessions_can_make_progress(): + async def run(): + controller = PDAdmissionController(lambda: 2) + first_a = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="a")) + second_a_task = asyncio.create_task( + controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="a")) + ) + await asyncio.sleep(0) + session_b = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="b")) + + assert second_a_task.done() is False + session_b.release() + assert second_a_task.done() is False + + first_a.release() + second_a = await second_a_task + second_a.release() + + asyncio.run(run()) + + +def test_cancelled_waiter_is_removed_and_does_not_leak_capacity(): + async def run(): + controller = PDAdmissionController(lambda: 1) + active = await controller.acquire(_request()) + waiting_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + waiting_task.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting_task + + assert controller.queued_request_count == 0 + assert controller.queued_slots == 0 + active.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_cancelling_reserved_gang_removes_the_backfill_barrier(): + async def run(): + controller = PDAdmissionController( + lambda: 2, + policy=AdmissionPolicy(waiting_decode_waves=4), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + backfill_tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(2)] + await asyncio.sleep(0) + + backfills = [] + for active_lease, backfill_task in zip(active, backfill_tasks): + active_lease.release() + backfills.append(await backfill_task) + + later_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + assert later_task.done() is False + + gang_task.cancel() + with pytest.raises(asyncio.CancelledError): + await gang_task + + backfills[0].release() + later = await later_task + backfills[1].release() + later.release() + + asyncio.run(run()) + + +def test_wait_timeout_removes_request_from_queue(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=0.01, + probable_cache_hit_max_wait_seconds=0.01, + cold_max_wait_seconds=0.01, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + + with pytest.raises(ServerBusyError, match="wait timed out"): + await controller.acquire(_request()) + + assert controller.queued_request_count == 0 + active.release() + + asyncio.run(run()) + + +def test_reserved_gang_timeout_removes_the_backfill_barrier(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=1.0, + probable_cache_hit_max_wait_seconds=1.0, + cold_max_wait_seconds=0.03, + waiting_decode_waves=4, + ) + controller = PDAdmissionController(lambda: 2, policy=policy) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + await asyncio.sleep(0) + + active[0].release() + first_backfill_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + first_backfill = await first_backfill_task + + second_backfill_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + active[1].release() + second_backfill = await second_backfill_task + + later_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + with pytest.raises(ServerBusyError, match="wait timed out"): + await gang_task + + first_backfill.release() + later = await later_task + second_backfill.release() + later.release() + + asyncio.run(run()) + + +def test_capacity_changes_wake_waiters_without_overcommitting(): + async def run(): + capacity = [1] + controller = PDAdmissionController(lambda: capacity[0]) + first = await controller.acquire(_request()) + second_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + capacity[0] = 2 + controller.on_capacity_change() + second = await second_task + assert controller.active_slots == 2 + + capacity[0] = 1 + controller.on_capacity_change() + third_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + assert third_task.done() is False + + first.release() + await asyncio.sleep(0) + assert third_task.done() is False + second.release() + third = await third_task + third.release() + + asyncio.run(run()) + + +def test_capacity_shrink_fails_an_oversized_blocked_gang_and_clears_state(): + async def run(): + capacity = [3] + controller = PDAdmissionController( + lambda: capacity[0], + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(3)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=3))) + await asyncio.sleep(0) + + active[0].release() + assert gang_task.done() is False + + capacity[0] = 2 + controller.on_capacity_change() + with pytest.raises(ServerBusyError, match="fell below queued request size"): + await gang_task + assert controller.queued_slots == 0 + + small_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + assert small_task.done() is False + active[1].release() + small = await small_task + active[2].release() + small.release() + + asyncio.run(run()) + + +def test_capacity_shrink_trims_queue_to_dynamic_slot_limit_by_priority_and_recency(): + async def run(): + capacity = [4] + policy = AdmissionPolicy(waiting_decode_waves=2) + controller = PDAdmissionController(lambda: capacity[0], policy=policy) + active = await controller.acquire(_request(decode_slots=4)) + + cold_tasks = [asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) for _ in range(3)] + probable_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) for _ in range(2) + ] + continuation_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) for _ in range(3) + ] + await asyncio.sleep(0) + assert controller.queued_slots == 8 + + capacity[0] = 1 + controller.on_capacity_change() + assert controller.queued_slots == capacity[0] * policy.waiting_decode_waves + + for task in cold_tasks + probable_tasks + [continuation_tasks[-1]]: + with pytest.raises(ServerBusyError, match="queue capacity shrank"): + await task + assert continuation_tasks[0].done() is False + assert continuation_tasks[1].done() is False + + active.release() + first = await continuation_tasks[0] + assert continuation_tasks[1].done() is False + first.release() + second = await continuation_tasks[1] + second.release() + + asyncio.run(run()) + + +def test_zero_capacity_fails_all_queued_requests(): + async def run(): + capacity = [2] + controller = PDAdmissionController( + lambda: capacity[0], + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(2)] + await asyncio.sleep(0) + + capacity[0] = 0 + controller.on_capacity_change() + for task in tasks: + with pytest.raises(ServerBusyError, match="fell below queued request size"): + await task + assert controller.queued_slots == 0 + assert controller.queued_request_count == 0 + + for lease in active: + lease.release() + + asyncio.run(run()) + + +def test_session_tracker_requires_a_recent_observed_success(): + now = [0.0] + tracker = SessionTracker(ttl_seconds=10, max_sessions=2, clock=lambda: now[0]) + + assert tracker.is_continuation("session-a") is False + tracker.mark_success("session-a") + assert tracker.is_continuation("session-a") is True + + now[0] = 11 + assert tracker.is_continuation("session-a") is False + + tracker.mark_success("session-a") + tracker.mark_success("session-b") + tracker.mark_success("session-c") + assert tracker.is_continuation("session-a") is False + assert tracker.is_continuation("session-b") is True + assert tracker.is_continuation("session-c") is True + + +def test_successful_session_promotes_its_waiting_requests(): + async def run(): + controller = PDAdmissionController( + lambda: 1, + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = await controller.acquire(_request()) + same_session_task = asyncio.create_task(controller.acquire(_request(session_key="session-a"))) + other_cold_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + controller.promote_session("session-a") + active.release() + promoted = await same_session_task + assert promoted.request.priority == AdmissionPriority.CONTINUATION + assert other_cold_task.done() is False + + promoted.release() + other = await other_cold_task + other.release() + + asyncio.run(run()) + + +def test_session_promotion_extends_the_wait_deadline(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=0.2, + probable_cache_hit_max_wait_seconds=0.1, + cold_max_wait_seconds=0.02, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + waiting_task = asyncio.create_task(controller.acquire(_request(session_key="session-a"))) + await asyncio.sleep(0.005) + + controller.promote_session("session-a") + await asyncio.sleep(0.025) + assert waiting_task.done() is False + + active.release() + promoted = await waiting_task + assert promoted.request.priority == AdmissionPriority.CONTINUATION + promoted.release() + + asyncio.run(run()) + + +def test_requests_remain_fifo_within_the_same_priority_class(): + async def run(): + controller = PDAdmissionController( + lambda: 1, + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = await controller.acquire(_request()) + first_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD, session_key="first"))) + second_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD, session_key="second"))) + await asyncio.sleep(0) + + active.release() + first = await first_task + assert second_task.done() is False + first.release() + second = await second_task + second.release() + + asyncio.run(run()) diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 6bdde7857..1e141110d 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -1,3 +1,4 @@ +import asyncio from types import SimpleNamespace import pytest @@ -73,6 +74,70 @@ 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_estimates_hit_rate_only_for_connected_worker(): + policy = CacheAwarePolicy(CacheAwareConfig(sample_stride=1)) + cached_worker = _worker("10.0.0.1:8000") + prompt = "shared conversation history and a new user turn" + policy.prompt_cache_tree.insert(prompt[:-10], cached_worker.client_ip_port) + + expected_hit_rate = len(prompt[:-10]) / len(prompt) + assert policy.estimate_cache_hit_rate([cached_worker], prompt) == pytest.approx(expected_hit_rate) + assert policy.estimate_cache_hit_rate([_worker("10.0.0.2:8000")], prompt) == 0.0 + + +def test_cache_aware_reuses_admission_match_during_worker_selection(monkeypatch): + policy = CacheAwarePolicy() + cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) + least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100, dispatched_req_num=2) + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) + + prefix_match = policy.prompt_cache_tree.prefix_match + match_call_count = 0 + + def counting_prefix_match(text): + nonlocal match_call_count + match_call_count += 1 + return prefix_match(text) + + monkeypatch.setattr(policy.prompt_cache_tree, "prefix_match", counting_prefix_match) + + async def select_worker(): + return policy.select_worker([cache_worker, least_loaded_worker], prompt) + + async def estimate_then_select_in_child_task(): + policy.estimate_cache_hit_rate([cache_worker, least_loaded_worker], prompt) + return await asyncio.gather(*(asyncio.create_task(select_worker()) for _ in range(2))) + + selected_workers = asyncio.run(estimate_then_select_in_child_task()) + assert selected_workers == [cache_worker, cache_worker] + assert match_call_count == 1 + + policy.select_worker([cache_worker, least_loaded_worker], prompt + " new turn") + assert match_call_count == 2 + + +def test_cache_aware_keeps_reused_matches_isolated_between_requests(): + policy = CacheAwarePolicy() + workers = [ + _worker("10.0.0.1:8000", dispatched_req_num=1), + _worker("10.0.0.2:8000", dispatched_req_num=1), + ] + prompts = ["a" * 1024, "b" * 1024] + for worker, prompt in zip(workers, prompts): + policy.prompt_cache_tree.insert(prompt, worker.client_ip_port) + + async def estimate_then_select(prompt): + policy.estimate_cache_hit_rate(workers, prompt) + await asyncio.sleep(0) + return policy.select_worker(workers, prompt) + + async def run_concurrent_requests(): + return await asyncio.gather(*(estimate_then_select(prompt) for prompt in prompts)) + + assert asyncio.run(run_concurrent_requests()) == workers + + 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) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94..252d7b0c7 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,12 +1,25 @@ import asyncio import json from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, call import pytest from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs +from lightllm.server.httpserver.pd_loop import ( + _allocate_capacity_share, + _build_pd_registration_info, + _update_pd_master_membership, +) +from lightllm.server.httpserver_for_pd_master.admission import ( + AdmissionPolicy, + AdmissionPriority, + PDAdmissionController, + SessionTracker, +) from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.server.pd_io_struct import PD_MASTER_CAPACITY_EPOCH_KEY, PD_MASTER_CAPACITY_SHARE_KEY def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -184,6 +197,457 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True +def test_pd_node_capacity_is_partitioned_without_overlap(): + master_ids = [30, 10, 20] + shares = [_allocate_capacity_share(8, master_ids, node_id) for node_id in sorted(master_ids)] + + assert shares == [3, 3, 2] + assert sum(shares) == 8 + assert _allocate_capacity_share(8, master_ids, 99) == 0 + + +def test_pd_registration_keeps_legacy_top_level_schema(monkeypatch): + from lightllm.server.httpserver import pd_loop + + args = SimpleNamespace(pd_node_id=7, running_max_req_size=8, host="0.0.0.0") + manager = SimpleNamespace( + args=args, + host_ip="10.0.0.7", + pd_mode=SimpleNamespace(value="decode"), + pd_master_ids=(10, 20), + pd_master_capacity_epoch=123, + ) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) + + registration = _build_pd_registration_info(manager, SimpleNamespace(node_id=10)) + + assert set(registration) == {"node_id", "client_ip_port", "mode", "start_args"} + assert registration["start_args"][PD_MASTER_CAPACITY_SHARE_KEY] == 4 + assert registration["start_args"][PD_MASTER_CAPACITY_EPOCH_KEY] == 123 + assert args.host == "0.0.0.0" + + +def test_pd_node_load_info_omits_radix_cache_hot_path(monkeypatch): + from lightllm.server import api_http + from lightllm.server.httpserver import pd_loop + + args = SimpleNamespace(tp=4, dp=2, nnodes=1, running_max_req_size=8) + httpserver_manager = SimpleNamespace( + host_ip="10.0.0.1", + pd_master_ids=(10, 20), + pd_master_capacity_epoch=123, + ) + shared_token_load = SimpleNamespace(get_dynamic_max_load=lambda dp_index: (0.25, 0.75)[dp_index]) + monkeypatch.setattr(api_http.g_objs, "args", args) + monkeypatch.setattr(api_http.g_objs, "httpserver_manager", httpserver_manager) + monkeypatch.setattr(api_http.g_objs, "shared_token_load", shared_token_load) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) + + assert not hasattr(pd_loop, "_get_radix_cache_info") + assert pd_loop._get_load_info(pd_master_node_id=10) == { + "total_token_usage_rate": 0.5, + "client_ip_port": "10.0.0.1:8001", + "capacity_share": 4, + "capacity_epoch": 123, + } + + +def test_pd_master_membership_change_advances_epoch_and_wakes_heartbeats(monkeypatch): + async def run(): + manager = SimpleNamespace() + timestamps = iter([100, 200]) + monkeypatch.setattr("lightllm.server.httpserver.pd_loop.time.time_ns", lambda: next(timestamps)) + + _update_pd_master_membership(manager, {20: object(), 10: object()}) + assert manager.pd_master_ids == (10, 20) + assert manager.pd_master_capacity_epoch == 100 + assert manager.pd_master_membership_changed.is_set() + + manager.pd_master_membership_changed.clear() + _update_pd_master_membership(manager, {10: object(), 20: object()}) + assert manager.pd_master_membership_changed.is_set() is False + + _update_pd_master_membership(manager, {10: object()}) + assert manager.pd_master_capacity_epoch == 200 + assert manager.pd_master_membership_changed.is_set() + + asyncio.run(run()) + + +def test_pd_manager_reports_only_actual_decode_capacity_changes(): + args = StartArgs() + manager = PDManager(args) + client_ip_port = "10.0.0.2:8000" + manager.register_pd( + { + "node_id": 2, + "client_ip_port": client_ip_port, + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 3, + "capacity_epoch": 100, + }, + websocket=object(), + ) + + assert manager.get_decode_capacity() == 3 + + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.25, + "capacity_share": 1, + "capacity_epoch": 99, + } + ) + is False + ) + assert manager.get_decode_capacity() == 3 + + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.5, + "capacity_share": 2, + "capacity_epoch": 101, + # Old nodes may continue sending cache telemetry during a + # rolling upgrade. It must not affect decode admission. + "radix_cache_total_tokens": 700, + "radix_cache_refed_tokens": 200, + "radix_cache_capacity_tokens": 1000, + } + ) + is True + ) + node = manager.decode_nodes[0] + assert manager.get_decode_capacity() == 2 + assert node.run_status.total_token_usage_rate == 0.5 + + # A fresh report and a load-only change are not capacity changes. + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.75, + "capacity_share": 2, + "capacity_epoch": 102, + } + ) + is False + ) + assert manager.get_decode_capacity() == 2 + + +def test_pd_registration_reads_capacity_from_legacy_safe_start_args(): + args = StartArgs() + manager = PDManager(args) + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + PD_MASTER_CAPACITY_SHARE_KEY: 3, + PD_MASTER_CAPACITY_EPOCH_KEY: 100, + }, + }, + websocket=object(), + ) + + node = manager.decode_nodes[0] + assert node.capacity_share == 3 + assert node.capacity_epoch == 100 + assert manager.get_decode_capacity() == 3 + + +def test_pd_disconnect_accepts_transitional_top_level_capacity_fields(): + args = StartArgs() + manager = PDManager(args) + registration = { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 3, + "capacity_epoch": 100, + } + manager.register_pd(registration, websocket=object()) + + manager.remove_pd(registration) + + assert manager.decode_nodes == [] + assert manager.url_to_pd_nodes == {} + + +def test_materializing_decode_fallback_share_is_not_a_capacity_change(): + args = StartArgs() + manager = PDManager(args) + client_ip_port = "10.0.0.2:8000" + manager.register_pd( + { + "node_id": 2, + "client_ip_port": client_ip_port, + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + }, + websocket=object(), + ) + + assert manager.decode_nodes[0].capacity_share is None + assert manager.get_decode_capacity() == 8 + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.5, + "capacity_epoch": 1, + } + ) + is False + ) + assert manager.get_decode_capacity() == 8 + + +def test_prefill_cache_telemetry_does_not_change_decode_capacity(): + args = StartArgs() + manager = PDManager(args) + manager.register_pd( + { + "node_id": 1, + "client_ip_port": "10.0.0.1:8000", + "mode": "prefill", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + "max_image_pixels": args.max_image_pixels, + "disable_image_resize": args.disable_image_resize, + }, + }, + websocket=object(), + ) + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 4, + }, + websocket=object(), + ) + + assert manager.get_decode_capacity() == 4 + assert ( + manager.update_node_load_info( + { + "client_ip_port": "10.0.0.1:8000", + "total_token_usage_rate": 0.25, + "radix_cache_total_tokens": 1000, + "radix_cache_refed_tokens": 100, + "radix_cache_capacity_tokens": 1000, + } + ) + is False + ) + assert manager.get_decode_capacity() == 4 + + +def test_pd_master_redrains_admission_only_for_decode_capacity_changes(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_manager = SimpleNamespace(update_node_load_info=MagicMock(side_effect=[False, True])) + manager.admission_controller = SimpleNamespace(on_capacity_change=MagicMock()) + + unchanged_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 4} + changed_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 2} + + manager.update_node_load_info(unchanged_load) + manager.admission_controller.on_capacity_change.assert_not_called() + + manager.update_node_load_info(changed_load) + manager.admission_controller.on_capacity_change.assert_called_once_with() + assert manager.pd_manager.update_node_load_info.call_args_list == [ + call(unchanged_load), + call(changed_load), + ] + + +def test_token_packs_redrain_admission_only_after_decode_capacity_change(): + from lightllm.server.pd_io_struct import ObjType + + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None) + manager.pd_manager = SimpleNamespace(update_node_load_info=MagicMock(side_effect=[False, True])) + manager.admission_controller = SimpleNamespace(on_capacity_change=MagicMock()) + manager.timer_log = AsyncMock() + manager.infos_queues = None + manager.req_id_to_out_inf = {} + + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + + unchanged_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 4} + changed_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 2} + await manager.put_to_handle_queue((ObjType.TOKEN_PACKS, [], unchanged_load)) + await manager.put_to_handle_queue((ObjType.TOKEN_PACKS, [], changed_load)) + + for _ in range(100): + if manager.pd_manager.update_node_load_info.call_count == 2: + break + await asyncio.sleep(0) + + assert manager.pd_manager.update_node_load_info.call_args_list == [ + call(unchanged_load), + call(changed_load), + ] + manager.admission_controller.on_capacity_change.assert_called_once_with() + finally: + handle_task.cancel() + await asyncio.gather(handle_task, return_exceptions=True) + + asyncio.run(run()) + + +def test_pd_master_deduplicates_unchanged_admission_metrics(monkeypatch): + from lightllm.server.httpserver_for_pd_master import manager as manager_module + + metric_client = SimpleNamespace(gauge_set=MagicMock()) + monkeypatch.setattr(manager_module, "MetricClient", lambda _port: metric_client) + monkeypatch.setattr(manager_module, "ReqIDGenerator", lambda: object()) + monkeypatch.setattr(manager_module, "get_tokenizer", lambda *_args, **_kwargs: object()) + monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) + + manager = HttpServerManagerForPDMaster(StartArgs(max_req_total_len=1024)) + manager.pd_manager.get_decode_capacity = lambda: 4 + controller = SimpleNamespace(queued_request_count=2, active_slots=1) + + async def run(): + manager._record_admission_state(controller) + controller.active_slots = 2 + manager._record_admission_state(controller) + assert metric_client.gauge_set.call_count == 0 + + await asyncio.sleep(0) + assert {metric_call.args for metric_call in metric_client.gauge_set.call_args_list} == { + ("lightllm_pd_master_admission_queue_size", 2), + ("lightllm_pd_master_admission_active_slots", 2), + ("lightllm_pd_master_admission_decode_capacity", 4), + } + first_call_count = metric_client.gauge_set.call_count + + manager._record_admission_state(controller) + await asyncio.sleep(0) + assert metric_client.gauge_set.call_count == first_call_count + + controller.active_slots = 3 + manager._record_admission_state(controller) + await asyncio.sleep(0) + assert metric_client.gauge_set.call_args_list[-1] == call("lightllm_pd_master_admission_active_slots", 3) + assert metric_client.gauge_set.call_count == first_call_count + 1 + + asyncio.run(run()) + + +def test_pd_master_decode_lease_covers_prefill_and_stream_lifecycle(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + manager.metric_client = SimpleNamespace(histogram_observe=lambda *_args: None) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + dispatched_prompts = [] + + async def fake_generate(prompt, *_args): + dispatched_prompts.append(prompt) + yield 1, "prefill", {"prompt_tokens": 16, "prompt_cache_len": 0}, object() + yield 1, "decode", {}, object() + + manager._generate = fake_generate + first = manager.generate("first", None, None, None) + first_prefill = await first.__anext__() + assert first_prefill[1] == "prefill" + assert manager.admission_controller.active_slots == 1 + + second = manager.generate("second", None, None, None) + second_result = asyncio.create_task(second.__anext__()) + await asyncio.sleep(0) + assert second_result.done() is False + assert dispatched_prompts == ["first"] + + first_decode = await first.__anext__() + assert first_decode[1] == "decode" + assert manager.admission_controller.active_slots == 1 + await first.aclose() + + assert (await second_result)[1] == "prefill" + await second.aclose() + assert manager.admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_master_releases_admission_lease_when_post_acquire_setup_fails(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + + def fail_histogram(*_args): + raise RuntimeError("metric enqueue failed") + + manager.metric_client = SimpleNamespace(histogram_observe=fail_histogram) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + + async def fake_generate(*_args): + yield "must not dispatch" + + manager._generate = fake_generate + generator = manager.generate("prompt", SimpleNamespace(n=1), None, None) + with pytest.raises(RuntimeError, match="metric enqueue failed"): + await generator.__anext__() + + assert manager.admission_controller.active_slots == 0 + assert manager.running_request_count == 0 + + asyncio.run(run()) + + def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) @@ -262,12 +726,80 @@ def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): assert HttpServerManagerForPDMaster.is_healthy(manager) is True +def test_pd_master_waits_before_dispatching_beyond_decode_capacity(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + manager.metric_client = SimpleNamespace(histogram_observe=lambda *_args: None) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + dispatched_prompts = [] + + async def fake_generate(prompt, *_args): + dispatched_prompts.append(prompt) + yield prompt + + manager._generate = fake_generate + first = manager.generate("first", None, None, None) + assert await first.__anext__() == "first" + + second = manager.generate("second", None, None, None) + second_result = asyncio.create_task(second.__anext__()) + await asyncio.sleep(0) + assert dispatched_prompts == ["first"] + assert manager.admission_controller.queued_request_count == 1 + + await first.aclose() + assert await second_result == "second" + assert dispatched_prompts == ["first", "second"] + await second.aclose() + assert manager.admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_master_admission_classifies_priority_and_multi_choice_cost(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.75), + ) + manager.admission_policy = AdmissionPolicy() + manager.session_tracker = SessionTracker(ttl_seconds=60, max_sessions=10) + sampling_params = SimpleNamespace(n=3) + + probable = manager._build_admission_request( + "abcdefghij", + sampling_params, + session_key="session-a", + ) + assert probable.priority == AdmissionPriority.PROBABLE_CACHE_HIT + assert probable.decode_slots == 3 + + manager.session_tracker.mark_success("session-a") + continuation = manager._build_admission_request( + "abcdefghij", + sampling_params, + session_key="session-a", + ) + assert continuation.priority == AdmissionPriority.CONTINUATION + + def test_pd_master_restores_request_count_when_preload_fails(): class FailingMultimodalParams: async def verify_and_preload(self, request): raise RuntimeError("preload failed") manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs(disable_pd_master_decode_capacity_limit=True) manager.running_request_count = 0 async def consume_generate(): @@ -282,6 +814,7 @@ async def consume_generate(): def test_pd_master_request_count_covers_async_generator_lifecycle(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs(disable_pd_master_decode_capacity_limit=True) manager.running_request_count = 0 inner_generator_closed = False