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 等)变得透明。