执行摘要
- 一句话:显式路由PD warmup到每个DP rank,修复覆盖不全问题
- 推荐动作:值得精读。设计上使用显式路由到每个DP rank、并通过
asyncio.gather并发控制,是典型的批量初始化模式。建议关注return_exceptions讨论,根据实际容错需求决定是否添加。
功能与动机
根因是bootstrap-room affinity在total_tokens等显式负载均衡策略下不被使用,startup warmup不能保证每个rank都执行dry run(PR body描述)。某个rank首次被真实流量接触时才触发延迟的JIT或初始化成本,引入rank skew。
实现拆解
- 新增异步预热函数:在
http_server.py中新增_send_disaggregation_warmup_requests,接收server_args、url、headers、ssl_verify和timeout参数,返回每个rank的状态码列表。
- 构造每个rank的显式路由请求:内部
send_request协程为单个dp_rank构建JSON体,其中routed_dp_rank和bootstrap_room均设置为该rank值,bootstrap_host固定为FAKE_BOOTSTRAP_HOST,输入为标量input_ids。
- 并发发送所有请求:函数内创建
aiohttp.ClientSession,使用asyncio.gather并发执行所有rank的send_request,所有请求共享同一个session。
- 修改预热入口:在
_execute_server_warmup中,将原来单次requests.post批次请求替换为调用asyncio.run(_send_disaggregation_warmup_requests(...)),检查返回的状态码列表,仅当全为200时才标记预热成功。
- 新增单元测试:新建
test_http_server_warmup.py,通过mock aiohttp.ClientSession验证函数并发发送N个请求、每个请求的routed_dp_rank覆盖0~dp_size-1,且所有响应状态码200。
关键文件:
python/sglang/srt/entrypoints/http_server.py(模块 预热逻辑;类别 source;类型 core-logic;符号 _send_disaggregation_warmup_requests, send_request, _execute_server_warmup): 核心修改文件,新增异步并发warmup函数并修改预热入口。
test/registered/unit/entrypoints/test_http_server_warmup.py(模块 预热测试;类别 test;类型 test-coverage;符号 TestDisaggregationServerWarmup, test_sends_concurrent_scalar_request_to_each_dp_rank, Response, Session): 新增测试验证并发发送到每个rank的正确性。
关键符号:_send_disaggregation_warmup_requests, send_request, _execute_server_warmup
关键源码片段
python/sglang/srt/entrypoints/http_server.py
核心修改文件,新增异步并发warmup函数并修改预热入口。
async def _send_disaggregation_warmup_requests(
server_args: ServerArgs,
url: str,
headers: Dict[str, str],
ssl_verify: Union[bool, str],
timeout: int,
) -> List[int]:
"""发送异步并发 warmup 请求到每个 DP rank。"""
# 根据 ssl_verify 创建 SSL 上下文
ssl_context = (
ssl_verify
if isinstance(ssl_verify, bool)
else ssl.create_default_context(cafile=ssl_verify)
)
# 内部函数:为单个 dp_rank 发送一个 /generate 请求
async def send_request(session: aiohttp.ClientSession, dp_rank: int) -> int:
# 构造请求体,显式指定 routed_dp_rank 和 bootstrap_room
json_data = {
"sampling_params": {"temperature": 0.0, "max_new_tokens": 8, "ignore_eos": True},
"bootstrap_host": FAKE_BOOTSTRAP_HOST,
"bootstrap_room": dp_rank, # 作为 rank 唯一标识
"input_ids": [10, 11, 12, 13],
"routed_dp_rank": dp_rank, # 显式路由到此 rank
}
async with session.post(url + "/generate", json=json_data, ssl=ssl_context) as response:
await response.read()
return response.status
# 使用共享的 aiohttp ClientSession 并发发送所有请求
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(total=timeout),
headers=headers,
) as session:
# 注意:此处未设置 return_exceptions=True,若某个请求失败会直接抛出异常
return await asyncio.gather(
*(send_request(session, dp_rank) for dp_rank in range(server_args.dp_size))
)
test/registered/unit/entrypoints/test_http_server_warmup.py
新增测试验证并发发送到每个rank的正确性。
class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase):
async def test_sends_concurrent_scalar_request_to_each_dp_rank(self):
server_args = SimpleNamespace(dp_size=4)
all_started = asyncio.Event()
calls = []
sessions = []
class Response:
# mock 响应:状态码 200,__aenter__ 屏障确保并发
status = 200
async def __aenter__(self):
# 当所有 dp_size 个调用都到达后释放屏障
if len(calls) == server_args.dp_size:
all_started.set()
await asyncio.wait_for(all_started.wait(), timeout=5)
return self
async def __aexit__(self, *args):
pass
async def read(self):
return b""
class Session:
# mock ClientSession,记录每个调用及其 kwargs
def __init__(self, **kwargs):
self.kwargs = kwargs
sessions.append(self)
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
def post(self, *args, **kwargs):
calls.append((args, kwargs))
return Response()
with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session):
status_codes = await _send_disaggregation_warmup_requests(
server_args=server_args,
url="http://localhost:30000",
headers={"Authorization": "Bearer token"},
ssl_verify=False,
timeout=123,
)
# 验证所有 rank 都收到 200
self.assertEqual(status_codes, [200] * server_args.dp_size)
# 验证每个 rank 的请求体中 routed_dp_rank 唯一且覆盖 0..dp_size-1
calls_by_rank = {kwargs["json"]["routed_dp_rank"]: (args, kwargs) for args, kwargs in calls}
self.assertEqual(set(calls_by_rank), set(range(server_args.dp_size)))
评论区精华
审阅者ShangmingCai建议在asyncio.gather中添加return_exceptions=True(line 2055),以避免单个rank请求失败时整个预热抛出异常而中断。该建议未被作者采纳,最终提交未包含该参数。讨论状态未关闭,仍属未解决事项。
- 建议在asyncio.gather中添加return_exceptions=True (design): 未采纳,代码保持无return_exceptions,单个rank失败将直接中断预热。
风险与影响
- 风险:
- 启动时间增加:现在发送dp_size个并发请求,但预热超时长达1800秒(默认),实际并发开销小,风险可控。
- 单个rank失败风险:未设置
return_exceptions=True,若某个rank的预热请求失败(如网络抖动),asyncio.gather会直接抛出异常,导致服务器就绪失败。但这也是一种严格保障,确保所有rank都预热成功。
- 依赖aiohttp:新增导入
aiohttp,但该库已是项目常用依赖,不需额外安装,风险极低。
- 影响:影响范围:仅改动disaggregated预热路径,不涉及正常请求路由。对用户:DP-Attention部署的PD服务器在启动时可靠地预热所有rank,消除首次流量时的rank skew。对系统:启动时多出dp_size个并发HTTP请求,但总耗时相近,因为原有批次请求本身需要等待所有rank处理。对团队:代码更清晰,可维护性改善。
- 风险标记:单个DP rank失败会阻止服务器就绪, 启动时间线性增长(dp_size个并发请求)
关联脉络
参与讨论