今天聊一个在 RL 工程化里很容易被低估的问题:推理(Inference)才是强化学习训练流水线里最容易被卡死的环节。
很多人跑过 PPO、GRPO 这类在线策略强化学习后会有一个共同感受:策略模型参数更新本身并不慢,真正把训练节奏拖垮的是 rollout 阶段——每更新一次策略,就要用最新权重重新生成一大批轨迹。自回归生成是逐 token 解码,一次 rollout 要串行跑几十上百步,再加上 RL 训练本来就需要大量采样来估计优势函数,推理开销会迅速膨胀到训练开销的几倍甚至十几倍。
问题的关键不在于“多买几张卡硬算”,而在于架构上是否把推理当作一个可以独立扩展的能力来对待。训练和推理的算力需求、显存特征、弹性伸缩方式完全不同,放在同一个进程里相互耦合,只会让两边都迁就短板。更合理的做法,是把策略模型的推理单独做成服务化组件,让它独立扩缩容,再通过接口向训练循环提供采样结果。
这篇文章会从瓶颈产生的原因讲起,给出训练/推理解耦的架构思路,并提供一套可落地的部署、调用、性能观察和问题排查方案。如果你正在做大模型 RL、RLHF、推理时扩展(inference-time scaling),或者只是觉得自己的 RL 训练跑得不够快,这篇文章可以直接参考。
1. 核心能力速览
| 维度 | 说明 |
|---|---|
| 核心问题 | RL 训练中策略模型推理(采样 / rollout)成为训练吞吐瓶颈 |
| 解决思路 | 将推理从训练循环中拆出,作为独立服务单独扩展 |
| 技术方向 | vLLM、TensorRT-LLM 等推理服务化,采样任务异步化 |
| 硬件需求 | 按模型规模配置推理节点,显存需高于模型权重最低需求,具体以实测为准 |
| 扩展方式 | 推理服务水平扩展 + 请求队列 + 批量采样 |
| 适用框架 | PPO、GRPO、DPO 等需要在线采样的 RL 流程 |
| 接口形态 | OpenAI 兼容接口 / 自定义 rollout 接口 |
| 适用场景 | 大模型 RL 训练、RLHF、多轮推理、推理时搜索扩展 |
| 主要收益 | 训练和推理互不阻塞,采样吞吐可独立扩容,训练迭代节奏更稳定 |
这里不涉及某个特定开源项目的安装教程,而是讲一种工程范式:RL 训练应该像调用外部能力一样调用推理服务。如果你已经把 vLLM 部署过一遍,下面很多内容可以直接迁移到自己的训练框架里。
2. 适用场景与使用边界
2.1 适合什么场景
- 在线策略算法:PPO、GRPO、Reinforce++ 这类算法要求每一步都用“当前策略”采样,旧参数产生的数据无法直接复用,推理被反复调用,瓶颈效应最明显。
- 高采样量任务:奖励模型不可导、需要通过生成多条候选轨迹来估计优势的任务,采样量越大,独立扩展推理的价值越大。
- 长序列推理:代码生成、数学推理、Agent 多步决策等任务,生成序列长,自回归解码时间长,推理耗时占训练总耗时的比例更高。
- 推理时计算扩展:需要做多路径采样、树搜索、最佳-of-N 采样等策略,本质上是把推理当作可扩展计算资源来使用。
2.2 不适合什么场景
- 小模型、小规模实验:单卡就能跑完整个 RL 循环时,强行拆出推理服务反而增加网络开销和部署复杂度。
- 实时交互调试阶段:当你在频繁修改奖励函数和训练逻辑时,一体化进程更容易定位问题。
- 对隐私和网络隔离要求极高、无法在训练节点之外启动独立服务的内网环境,需要先解决网络连通性再做拆分。
2.3 使用边界与合规提醒
- 使用的策略基座模型需要确认开源授权范围,商用环境需要单独核对模型 License。
- 通过 API 对外提供推理服务时,建议限制访问来源、设置鉴权,避免未授权调用。
- 如果采样数据包含用户隐私、肖像、版权素材,需要先完成脱敏和授权确认。
- RoLLout 数据会进入训练集,生成内容的版权和有害信息过滤需要在数据管线中单独处理。
3. 为什么推理会成为 RL 的瓶颈
3.1 在线策略的“新鲜度”要求
监督学习里,一批训练数据可以反复使用多个 epoch。但 PPO 这类 on-policy 算法要求采样数据来自“当前策略”,策略每更新一次,旧数据的行为分布就偏离了,继续使用会带来严重的偏差。也就是说,每个训练 step 都要拿最新权重重新跑一遍大规模生成,推理任务永远在线、永远不重复。
3.2 自回归生成的固有开销
生成一条长度为 L 的轨迹,至少需要串行执行 L 次前向传播,每次只产生一个 token。这个开销远高于训练阶段一次 forward-backward 对固定长度序列的处理。RL 还要通过 temperature、top_p 等采样参数生成多条候选来估计收益分布,实际 token 产量是训练数据消耗量的大好几倍。
3.3 训练和推理的特征冲突
训练阶段追求高吞吐,适合大批次、长序列、梯度累积;推理阶段追求低延迟和高并发,需要 continuous batching、PagedAttention 这类动态调度。把两者挤在同一批 GPU 上,显存分配和调度策略会互相干扰。训练节点需要预留优化器状态和梯度显存,推理节点则需要大量 KV Cache,混合部署容易出现“训练不满载、推理不够用”的尴尬局面。
3.4 一个简单的数量级感
假设一个 7B 模型,训练一个 batch 的前向反向耗时是 T,生成 128 条、每条 512 token 的轨迹,自回归解码的耗时往往数倍于 T。当训练步数达到几千甚至几万时,rollout 总耗时会在整体时间里占据绝对主导。这时候盲目增加训练卡数并没有用,因为瓶颈根本不在训练侧。
4. 架构设计:把推理从训练循环里拆出去
4.1 整体架构
推荐的架构是“训练端-推理服务端-数据缓冲”三层:
训练节点 推理服务节点 │ │ ├─ 更新策略参数 ──────────► 加载最新权重 │ │ │◄────────── rollout 请求 ─┤ │ │ ├─ 接收采样结果 │ │ │ ├─ 计算 reward / advantage │ │ │ └─ 更新参数 ──────────────►实际部署时不需要真的在训练节点和推理节点之间搬权重文件,可以共用共享存储(NFS、对象存储、分布式文件系统),训练端保存 checkpoint 后,推理服务动态加载或重启加载新版本。
4.2 为什么要独立扩展
推理服务可以单独按采样量扩容。训练阶段峰值采样时,启动更多推理副本;奖励计算或参数更新阶段,缩容推理副本。训练循环不再等待推理完成,而是通过异步请求-响应模式把 rollout 和训练重叠起来。
4.3 最小可用组件
- 推理服务:vLLM 或其他支持 continuous batching 的推理框架,监听一个 HTTP/gRPC 端口。
- 请求客户端:训练循环里封装一个 rollout client,把 prompt 列表发送给推理服务。
- 经验缓冲队列:采样结果先写入内存或磁盘队列,训练循环按需消费。
- 策略版本管理:每次参数更新后记录策略版本号,确保采样结果对应当前版本。
5. 环境准备与前置条件
5.1 节点规划
| 节点 | 职责 | 建议配置 |
|---|---|---|
| 训练节点 | 策略参数更新、奖励计算 | GPU 显存需容纳模型权重 + 优化器状态 + 梯度 |
| 推理节点 | 加载策略模型、批量生成 | GPU 显存需容纳模型权重 + KV Cache,可多机多卡 |
| 共享存储 | 权重文件、采样数据、日志 | 训练节点与推理节点均可访问 |
5.2 软件依赖
- Python 3.10+;
- PyTorch(训练框架使用,版本按训练脚本要求);
- vLLM 或 TensorRT-LLM(推理服务端);
- 可选 Ray 或 Celery(分布式采样调度);
- 模型文件:需要是推理框架支持的格式,例如 Hugging Face 格式或 TensorRT-LLM 的 Engine 格式。
5.3 检查清单
- GPU 驱动与 CUDA 版本是否匹配;
- 推理服务端和训练端是否在同一内网,端口是否互通;
- 磁盘是否有足够空间存放 rollout 数据;
- 模型文件路径是否对推理服务可见;
- 推理服务请求超时时间是否设置合理。
6. 部署步骤:用 vLLM 把推理服务跑起来
下面以 vLLM 为例演示如何把 RL 策略模型变成可独立扩展的推理服务。如果你的项目用的是其他推理框架,逻辑完全一致,只是启动参数不同。
6.1 启动推理服务
python -m vllm.entrypoints.openai.api_server \ --model /data/models/policy-model \ --tensor-parallel-size 1 \ --max-model-len 4096 \ --port 8000 \ --host 0.0.0.0 \ --gpu-memory-utilization 0.9 \ --enforce-eager参数说明:
--model:策略模型路径。--tensor-parallel-size:张量并行卡数,模型较大时按需调整。--max-model-len:最大序列长度,会影响 KV Cache 显存占用。--gpu-memory-utilization:控制推理进程最多使用多少显存,避免 OOM。--enforce-eager:禁用 CUDA Graph,便于调试显存问题。
启动后可以用 curl 验证服务是否正常:
curl http://127.0.0.1:8000/v1/completions \ -H "Content-Type: application/json" \ -d '{ "model": "/data/models/policy-model", "prompt": "Hello", "max_tokens": 16, "n": 1 }'6.2 在训练循环里调用推理服务
用 OpenAI 兼容接口的话,可以用 openai 库或直接 requests 调用。RL 采样阶段和普通对话补全唯一区别是:一个采样步要生成大量候选,最好通过n参数批量返回,减少 HTTP 往返次数。
import requests def generate_samples(prompt, n=8, max_tokens=512, temperature=0.8): url = "http://127.0.0.1:8000/v1/completions" payload = { "model": "/data/models/policy-model", "prompt": prompt, "max_tokens": max_tokens, "temperature": temperature, "n": n, "stop": ["<|endoftext|>"] } response = requests.post(url, json=payload, timeout=120) response.raise_for_status() data = response.json() return [item["text"] for item in data["choices"]]6.3 把 rollout 和训练解耦的伪代码
下面是一个简化版的 GRPO/PPO 训练循环结构,重点展示采样和训练如何分离:
class RolloutClient: """封装推理服务调用""" def __init__(self, endpoint: str): self.endpoint = endpoint def sample(self, prompts, n=1, temperature=0.7): # 实际实现可以改成批量异步请求 results = [] for prompt in prompts: results.append(generate_samples(prompt, n=n, temperature=temperature)) return results class RLTrainer: def __init__(self, policy_model, rollout_client, reward_func, lr=1e-6): self.policy = policy_model self.rollout_client = rollout_client self.reward_func = reward_func self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=lr) def train_step(self, prompts, num_samples=8): # 1. 采样 samples = self.rollout_client.sample(prompts, n=num_samples) # 2. 构造轨迹和奖励 trajectories = build_trajectories(prompts, samples) rewards = self.reward_func(trajectories) # 3. 计算优势(简化:GRPO 用组内相对奖励) advantages = compute_group_advantages(rewards, num_samples) # 4. 策略更新 loss = self.compute_policy_loss(trajectories, advantages) self.optimizer.zero_grad() loss.backward() self.optimizer.step() def compute_policy_loss(self, trajectories, advantages): # 根据具体 RL 算法实现,这里不展开 pass这里的关键点是RolloutClient和训练逻辑不再共享进程,训练端不需要加载推理模型,也不需要在反向传播时等待解码完成。
6.4 双缓冲:让采样和训练重叠
最简单的优化是采样阶段和训练阶段串行交替,但这种方式仍然有空闲。工程上常用双缓冲:当前 batch 在训练时,下一个 batch 的采样已经在推理服务上并发执行。
from concurrent.futures import ThreadPoolExecutor import time def train_with_double_buffer(trainer, prompts_pool, num_steps=100, num_workers=4): executor = ThreadPoolExecutor(max_workers=num_workers) # 预取第一批 future = executor.submit(trainer.rollout_client.sample, prompts_pool[0]) for step in range(num_steps): prompts = prompts_pool[step] # 提交下一批采样任务 next_future = executor.submit(trainer.rollout_client.sample, prompts_pool[step + 1]) \ if step + 1 < num_steps else None # 等待当前 batch 采样完成 samples = future.result() # 训练当前 batch trainer.train_step_from_samples(prompts, samples) # 切换到下一个 future future = next_future print(f"step {step} done at {time.time():.2f}")双缓冲能把采样延迟从训练关键路径上部分隐藏掉。推理服务越稳定,训练循环的等待就越少。
7. 功能测试与效果验证
推理服务部署完,不能只看“能返回结果”就认为链路通了。下面是一套针对 RL 采样场景的验证流程。
7.1 基础连通性测试
测试目标:确认推理服务可以正常处理补全请求。
操作步骤:
- 启动 vLLM 服务;
- 用 curl 或 Python 发送一个简单 prompt;
- 检查返回结果是否包含完整文本、是否触发 stop。
判断标准:
- HTTP 200;
- 返回文本非空;
- 服务日志无 CUDA OOM 报错。
常见失败:
- 模型路径错误导致启动失败;
- 端口被占用;
- 显存不足导致进程退出。
7.2 批量采样测试
测试目标:验证 RL 场景最常用的“一次返回多条候选”能力。
import time prompt = "The capital of France is" start = time.time() samples = generate_samples(prompt, n=16, max_tokens=64, temperature=0.9) elapsed = time.time() - start print(f"elapsed: {elapsed:.2f}s") print(f"num samples: {len(samples)}") for i, s in enumerate(samples[:3]): print(f"sample {i}: {s[:50]}")判断标准:
- 返回条数等于
n; - 候选之间有明显多样性;
- 耗时在可接受范围,如果慢得离谱,需要观察显存和批处理配置。
7.3 并发压力测试
测试目标:模拟训练循环高并发请求时服务是否稳定。
from concurrent.futures import ThreadPoolExecutor, as_completed prompts = [ "Write a python function to compute fibonacci numbers.", "Explain quantum entanglement in simple terms.", "Write a short story about a robot learning to paint.", ] * 20 def worker(prompt): return generate_samples(prompt, n=4, max_tokens=128) with ThreadPoolExecutor(max_workers=16) as executor: futures = [executor.submit(worker, p) for p in prompts] for idx, fut in enumerate(as_completed(futures)): fut.result() if idx % 20 == 0: print(f"completed {idx + 1} requests")判断标准:
- 所有请求都成功返回;
- 无连接超时;
- 服务端日志无异常。
7.4 训练端到端验证
测试目标:确认训练循环能从推理服务拿到数据并完成参数更新。
操作步骤:
- 用一个极小模型(如 1B 以下)跑 20 个 training step;
- 每一步记录采样耗时、训练耗时、loss 值;
- 对比拆分解耦前后的整体耗时变化。
判断标准:
- loss 正常下降;
- 训练耗时不被采样请求阻塞;
- 采样和训练的时间线有重叠。
8. 接口 API 与批量任务
8.1 API 形态选择
| 形态 | 优点 | 缺点 | 适用 |
|---|---|---|---|
| OpenAI 兼容 HTTP | 生态成熟、工具多 | 文本协议有一定开销 | 快速接入、实验验证 |
| gRPC | 低延迟、强类型 | 需要生成 proto 客户端 | 大规模生产环境 |
| 内部 SDK 直连 | 延迟最低 | 耦合度太高,扩展性差 | 不推荐用于独立扩展 |
RL 训练如果采样量巨大,建议先用 OpenAI 兼容接口完成验证,后续再视性能瓶颈决定是否切换到 gRPC。
8.2 批量任务队列
RL 采样本质上是大批量任务。最简单的批量方式是用线程池并发请求,但更稳妥的做法是引入队列:
- 训练端把需要采样的 prompt 写入 Redis / Kafka / 本地队列;
- 一组 Worker 进程消费队列,调用推理服务;
- 采样结果写回结果队列;
- 训练端从结果队列读取数据。
import queue import threading from typing import List import requests class RolloutWorker(threading.Thread): def __init__(self, task_queue, result_queue, endpoint): super().__init__() self.task_queue = task_queue self.result_queue = result_queue self.endpoint = endpoint self.daemon = True def run(self): while True: item = self.task_queue.get() if item is None: break try: samples = generate_samples(item["prompt"], n=item.get("n", 8)) self.result_queue.put({"prompt": item["prompt"], "samples": samples}) except Exception as e: self.result_queue.put({"prompt": item["prompt"], "error": str(e)}) finally: self.task_queue.task_done() def submit_rollout_tasks(prompts, task_queue, result_queue, n=8): for prompt in prompts: task_queue.put({"prompt": prompt, "n": n})批量任务要特别注意失败重试。如果一个请求超时,应该记录日志后重新放入队列,而不是让训练循环整体卡死。
9. 资源占用与性能观察
9.1 关键指标
| 指标 | 含义 | 查看方式 |
|---|---|---|
| GPU 利用率 | 推理卡是否满载 | nvidia-smi |
| 显存占用 | 权重 + KV Cache 占用量 | nvidia-smi/ vLLM metrics |
| 吞吐量 | tokens/s | vLLM/metrics |
| TTFT | 首 token 延迟 | vLLM metrics |
| TPOT | 每 token 解码耗时 | vLLM metrics |
| KV Cache 使用率 | 是否接近上限 | vLLM/metrics |
9.2 影响性能的主要因素
max_model_len:越大 KV Cache 占用越多,但能处理的单条请求越长;n采样数量:一次请求生成多条候选能提高整体吞吐,但不是线性提升;- 并发请求数:vLLM 使用 continuous batching,并发过低时 GPU 无法充分饱和;
- 温度、top_p:对性能影响不大,但会改变生成分布;
- GPU 显存利用率:设置过低会浪费显存,过高容易 OOM。
9.3 降低显存占用的通用方法
- 减小
max_model_len; - 降低
gpu-memory-utilization; - 使用更高吞吐的推理框架,比如开启 chunked prefill;
- 对超长序列做截断或分块处理;
- 必要时使用量化(如 AWQ、GPTQ)。
9.4 从时间线找瓶颈
在训练脚本里记录每个阶段的耗时:
import time import json def log_time(step, phase, duration): with open("timeline.jsonl", "a") as f: f.write(json.dumps({"step": step, "phase": phase, "duration": duration}) + "\n")如果采样耗时明显高于训练耗时,优先扩容推理节点或增加推理服务副本;如果训练耗时高于采样耗时,则瓶颈在训练侧,推理服务可以缩容。
10. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 推理服务启动后请求超时 | 模型加载失败、显存不足 | 查看服务日志、nvidia-smi | 降低并发数、减小max_model_len、调整gpu-memory-utilization |
| 训练端连不上推理服务 | 网络不通、端口未监听 | curl http://<ip>:8000/v1/models | 检查服务监听地址和防火墙规则 |
| 采样结果空洞、长度过短 | 设置了过强的 stop 条件 | 查看返回内容和日志 | 去掉或放宽 stop 条件,检查 EOS token |
| 采样结果几乎一样 | temperature 过低或 top_p 过小 | 打印不同请求的文本分布 | 适当提高 temperature、降低 top_p |
| CUDA OOM | KV Cache 设置过大 | 查看 vLLM 日志 | 降低max_num_seqs、gpu-memory-utilization |
| 批量任务中途卡住 | 队列积压、某个请求长时间不返回 | 看队列长度、请求日志 | 增加 Worker、设置单请求超时 |
| 策略模型更新后推理服务还在用旧权重 | 没有做版本管理 | 检查服务加载时间 | 重启推理服务或支持动态加载权重 |
| 整体速度反而变慢 | 规模太小,网络开销占主导 | 对比拆分前后耗时 | 小规模实验不拆服务,先直接跑通 |
11. 最佳实践与使用建议
11.1 先小规模验证链路
不要一开始就在 70B 模型上做完整拆解。先用 1B 或 3B 模型,跑通 vLLM 服务、训练循环、采样结果回流这条链路,再逐步放大。小模型能暴露大部分架构问题,且调试成本低。
11.2 把策略版本做成显式参数
训练过程中参数不断更新,旧权重生成的 rollout 数据不能直接丢弃也不能无脑混合。建议在训练数据中记录策略版本号,后续分析时能准确回溯。
11.3 推理服务安全边界
推理服务如果暴露到内网,务必增加鉴权。最简单的方案是在请求头里加 token,服务端校验后再放行。不要让训练集群之外的机器随意调用采样接口。
11.4 采样数据合规
RL 训练中生成的大量文本可能包含版权内容、隐私信息或有害内容。建议在数据落盘前做内容过滤,必要的时候引入人工抽检。
11.5 保留一套最小可运行配置
无论实验还是生产,都保留一组最小的可复现配置:模型路径、推理参数、训练超参、端口号。出现问题时可以快速回到已知可用状态。
11.6 监控和告警
对推理服务的吞吐、错误率和排队时间做监控。RL 训练时间很长,等到训练卡住了再人工介入,浪费的算力成本很高。建议至少设置以下告警:
- 推理服务错误率超过阈值;
- 采样队列积压超过 N 条;
- 推理服务 GPU 显存使用率超过 95%;
- 训练 step 平均耗时异常上涨。
12. 总结与下一步
“RL Is Bottlenecked by Inference. Scale It Independently” 这句话的本质是:RL 训练里推理不是训练循环的附属品,而是一个独立的核心计算环节。把它拆出来做成可独立扩展的推理服务,能解决采样吞吐不足、训练推理相互干扰、显存调度冲突等一系列问题。
最值得先试的事情,是用你当前正在用的策略模型启动一个 vLLM 服务,用 OpenAI 兼容接口跑一批批量采样请求,观察吞吐量和延迟,再把这个服务接到训练循环里做双缓冲。你会发现训练节奏比原来更容易控制,至少采样耗时不再和参数更新强耦合。
最容易踩的坑有两个:一是小规模实验也强行拆分,网络开销反而拖慢整体;二是推理服务更新权重不及时,导致采样数据来自过期策略。前者用规模阈值判断,后者用版本管理和重启策略解决。
下一步可以考虑两个方向:一是把采样任务做成真正的异步队列,用 Redis 或 Kafka 解耦训练端和推理端;二是在推理服务的基础上做 inference-time scaling,例如多路径采样、树搜索、best-of-n 重排,让推理能力直接服务于策略质量的提升。两者方向不同,但都建立在“推理可独立扩展”这个前提下。
如果你正在搭 RL 训练流水线,建议把这套拆分思路先放进设计文档,哪怕第一版不拆,也要给推理服务留出独立扩展的接口。等到采样量上来,再回头改架构,代价会大得多。