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 调用方式。