跳转至

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 都会被保存并返回,实现了"不浪费已有计算"的设计目标。