Skip to content
2 changes: 1 addition & 1 deletion lightllm/server/api_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 2 additions & 0 deletions lightllm/server/api_http_pd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
105 changes: 87 additions & 18 deletions lightllm/server/httpserver/pd_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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。
Expand Down Expand Up @@ -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

Expand All @@ -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
Loading
Loading