sglang_async_server.py — 实现了支持取消/恢复的 SGLang 推理服务器¶
文件路径: verl/experimental/fully_async_policy/sglang_rollout/sglang_async_server.py
文件概述¶
实现了支持取消/恢复的 SGLang 推理服务器。在全异步训练的部分回滚场景中,推理服务器需要支持在生成过程中被取消,并能返回已生成的部分结果。
关键代码讲解¶
1. SGLangHttpServerForPartial¶
class SGLangHttpServerForPartial(SGLangHttpServer):
def __init__(self, ...):
super().__init__(...)
self.paused = False
self.lock = asyncio.Lock()
self.cancel_event: dict[str, asyncio.Event] = {} # 每个 request 一个取消事件
self.req_output: dict[str, Optional[dict]] = {} # 每个 request 的输出
2. generate_for_partial - 支持取消的生成¶
async def generate_for_partial(self, prompt_ids, sampling_params, request_id, ...):
async with self.lock:
if self.paused:
return [], [], True # 已暂停,立即返回取消
self.cancel_event[request_id] = asyncio.Event()
cancel_handle = asyncio.create_task(self.cancel_event[request_id].wait())
generation_handle = asyncio.create_task(
self._generate_step(prompt_ids, sampling_params, request_id, ...)
)
# 等待生成完成或取消信号(谁先到用谁)
done, pending = await asyncio.wait(
[generation_handle, cancel_handle],
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel() # 取消未完成的任务
# 提取结果
async with self.lock:
output = self.req_output.get(request_id)
if output is None:
return [], [], True # 没有输出,说明被取消
# 提取 token IDs 和 log probs
token_ids, log_probs = [], []
for log_prob, token_id, _ in output["meta_info"]["output_token_logprobs"]:
token_ids.append(int(token_id))
log_probs.append(float(log_prob))
is_cancel = generation_handle not in done
return token_ids, log_probs, is_cancel
3. cancel / resume¶
async def cancel(self):
async with self.lock:
self.paused = True
for request_id in self.cancel_event:
self.cancel_event[request_id].set() # 触发所有取消事件
async def resume(self):
async with self.lock:
self.paused = False
4. FullyAsyncSGLangReplica¶
class FullyAsyncSGLangReplica(SGLangReplica):
def __init__(self, ...):
super().__init__(...)
self.server_class = ray.remote(SGLangHttpServerForPartial) # 使用支持取消的服务器
async def cancel(self):
await asyncio.gather(*[server.cancel.remote() for server in self.servers])
async def resume(self):
await asyncio.gather(*[server.resume.remote() for server in self.servers])
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
SGLangHttpServerForPartial |
类 | 支持取消的 SGLang HTTP 服务器 |
FullyAsyncSGLangReplica |
类 | 全异步 SGLang 副本管理器 |
与其他模块的关系¶
- 继承自
verl.workers.rollout.sglang_rollout.async_sglang_server中的SGLangHttpServer和SGLangReplica - 被
FullyAsyncAgentLoopManager使用 generate_for_partial被FullyAsyncLLMServerManager.generate_for_partial调用
小结¶
通过 asyncio.Event + asyncio.wait 的竞争模式,实现了推理生成和取消信号的竞争执行。无论生成是否完成,已经产出的 token 都会被保存并返回,实现了"不浪费已有计算"的设计目标。