跳转至

vllm_async_server.py — 实现了支持取消/恢复的 vLLM 推理服务器

文件路径: verl/experimental/fully_async_policy/vllm_rollout/vllm_async_server.py

文件概述

实现了支持取消/恢复的 vLLM 推理服务器。与 sglang_async_server.py 功能对等,但针对 vLLM 推理引擎。

关键代码讲解

1. vLLMHttpServerForPartial

class vLLMHttpServerForPartial(vLLMHttpServer):
    def __init__(self, ...):
        super().__init__(...)
        self.paused = False
        self.lock = asyncio.Lock()
        self.cancel_event: dict[str, asyncio.Event] = {}
        self.req_output: dict[str, Optional[RequestOutput]] = {}

2. generate_for_partial

与 SGLang 版本结构相同,但使用 vLLM 的 API:

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, pend = await asyncio.wait(
        [generation_handle, cancel_handle],
        return_when=asyncio.FIRST_COMPLETED
    )

    async with self.lock:
        if self.req_output[request_id] is None:
            return [], [], True

        token_ids = self.req_output[request_id].outputs[0].token_ids
        log_probs = []
        for i, x in enumerate(self.req_output[request_id].outputs[0].logprobs):
            token_id = token_ids[i]
            log_probs.append(x[token_id].logprob)

        is_cancel = generation_handle not in done
    return token_ids, log_probs, is_cancel

3. _generate_step 差异

vLLM 版本使用 SamplingParams 和 TokensPrompt:

async def _generate_step(self, prompt_ids, sampling_params, request_id, ...):
    prompt_ids = normalize_token_ids(prompt_ids)
    sampling_params["logprobs"] = 1
    sampling_params = SamplingParams(max_tokens=max_tokens, **sampling_params)
    prompt = TokensPrompt(prompt_token_ids=prompt_ids, multi_modal_data=multi_modal_data)
    generator = self.engine.generate(prompt=prompt, sampling_params=sampling_params, request_id=request_id)
    async for output in generator:
        self.req_output[request_id] = output

4. FullyAsyncvLLMReplica

class FullyAsyncvLLMReplica(vLLMReplica):
    def __init__(self, ...):
        super().__init__(...)
        self.server_class = ray.remote(vLLMHttpServerForPartial)

    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])

核心类/函数列表

名称 类型 说明
vLLMHttpServerForPartial 类 支持取消的 vLLM HTTP 服务器
FullyAsyncvLLMReplica 类 全异步 vLLM 副本管理器

与其他模块的关系

  • 继承自 verl.workers.rollout.vllm_rollout.vllm_async_server 中的 vLLMHttpServer 和 vLLMReplica
  • 与 sglang_async_server.py 功能对等
  • 被 FullyAsyncAgentLoopManager 根据配置选择使用

小结

vllm_async_server.py 是 SGLang 版本的 vLLM 对应实现。两者共享相同的取消/恢复架构(asyncio.Event 竞争模式),区别仅在于底层推理引擎的 API 调用方式。