跳转至

main_generation_server.py — 这是一个独立的文本生成服务器脚本

文件概述

模块路径: verl.trainer.main_generation_server

这是一个独立的文本生成服务器脚本。它启动多个 vLLM/SGLang 推理副本,通过 OpenAI 兼容的 API 接口对给定的 prompt 数据集批量生成回复,并将结果保存为 parquet 文件。

在训练流程中的位置

这个文件独立于训练流程,是一个纯推理工具。常见用途: - 训练后使用最终模型批量生成回复 - 为离线评估 (main_eval.py) 准备生成数据 - 大规模推理任务

关键代码讲解

1. 启动推理服务器

async def start_server(config):
    tp_size = config.actor_rollout_ref.rollout.tensor_model_parallel_size
    num_replicas = (config.trainer.n_gpus_per_node * config.trainer.nnodes) // tp_size

    rollout_server_class = get_rollout_replica_class(config.actor_rollout_ref.rollout.name)
    rollout_servers = [
        rollout_server_class(
            replica_rank=replica_rank,
            config=rollout_config,
            model_config=model_config,
            gpus_per_node=config.trainer.n_gpus_per_node,
        )
        for replica_rank in range(num_replicas)
    ]
    await asyncio.gather(*[server.init_standalone() for server in rollout_servers])
    return server_handles, server_addresses

根据 GPU 数量和张量并行大小,计算可以启动多少个推理副本。例如,8 GPU + tp_size=2 = 4 个副本。

2. 异步请求提交

async def submit_request(server_address, **chat_complete_request):
    timeout = aiohttp.ClientTimeout(total=None)
    session = aiohttp.ClientSession(timeout=timeout)
    async with session.post(
        url=f"http://{server_address}/v1/chat/completions",
        headers={"Authorization": "Bearer token-abc123", ...},
        json=chat_complete_request,
    ) as resp:
        data = await resp.json()
        return ChatCompletion(**data)

使用 aiohttp(而非 OpenAI SDK)发送 HTTP 请求,因为在大量并发请求时 aiohttp 不会出现挂起问题。

3. 按副本分配请求

async def generate(server_addresses, model_path, n_samples, sampling_params, chat_numpy):
    num_replicas = len(server_addresses)
    chat_sub_array = np.array_split(chat_numpy, num_replicas)  # 均匀分配
    results = await asyncio.gather(*[
        generate_per_replica(server_addresses[i], model_path, n_samples, sampling_params, chat_sub_array[i])
        for i in range(num_replicas)
    ])
    return results

将所有 prompt 均匀分配到各推理副本上,然后并行生成。

4. 主流程

@hydra.main(config_path="config", config_name="ppo_trainer", version_base=None)
def main(config):
    ray.init(...)

    sampling_params = {
        "temperature": config.actor_rollout_ref.rollout.temperature,
        "top_p": config.actor_rollout_ref.rollout.top_p,
        "max_tokens": config.actor_rollout_ref.rollout.response_length,
    }

    # 读取数据集
    dataset = pd.concat([pd.read_parquet(f) for f in train_files], ...)
    chat_numpy = np.array(chat_lst)

    # 启动服务器 + 生成
    server_handles, server_addresses = asyncio.run(start_server(config))
    gen_results = asyncio.run(generate(server_addresses, ..., chat_numpy))

    # 整理结果并保存
    results = np.reshape(results, (-1, n_samples))
    dataset["responses"] = results.tolist()
    dataset.to_parquet(config.data.output_path)

核心类/函数列表

名称 类型 作用
start_server(config) async function 启动多个推理副本
submit_request(addr, **req) async function 向单个副本提交一个请求
generate_per_replica(...) async function 一个副本处理一批 prompt
generate(...) async function 协调所有副本并行生成
main(config) function Hydra 入口,完整的生成流程

数据流和调用关系

parquet 数据集 (包含 prompt)
      |
      v
main() --> 读取数据集 --> np.array(chat_lst)
      |
      +-- start_server() --> 启动 N 个推理副本
      |
      +-- generate() --> 将 prompt 均匀分配到副本
      |       |
      |       +-- generate_per_replica() x N
      |               |
      |               +-- submit_request() x (prompt数 * n_samples)
      |
      +-- 整理结果 --> reshape(-1, n_samples)
      |
      +-- dataset.to_parquet(output_path)

小结

main_generation_server.py 是一个高效的批量推理工具,通过多副本异步并发的方式最大化 GPU 利用率。它使用 OpenAI 兼容的 API 接口,使得切换不同的推理后端(vLLM、SGLang 等)变得透明。