RL 项目的训练进度经常不是卡在反向传播上,而是卡在推理上。这里的推理,指的是策略模型在环境交互过程中不断生成动作或输出文本,也就是 rollout 阶段。很多团队一开始把训练和推理放在同一批 GPU 上跑,结果发现训练 step 本身很快,但每一轮迭代都要等很久,最后定位到的问题不是模型不收敛,而是推理吞吐跟不上数据生产速度。这个问题的正确解法不是简单加大 batch,而是把推理从训练链路里拆出来,作为一个独立服务单独扩展。
这篇文章围绕“独立扩展 RL 推理”这条主线展开:先讲清楚推理为什么是 RL 的瓶颈,再给出训练与推理混跑的问题分析,然后拆解推理服务化、动态批处理、KV Cache 等具体落地手段,最后补充监控指标、排查路径、学习环境与生产环境的差异,以及一份可以直接套用的扩展检查清单。
1. 先看懂 RL 训练链路里的推理环节
1.1 强化学习的循环:策略、环境、奖励
强化学习解决的问题,是一个智能体(Agent)如何通过和环境交互来学习最优策略。每一轮交互可以拆成四步:
- 策略模型根据当前状态生成一个动作。
- 环境接收动作,返回新的状态和奖励。
- 系统把状态、动作、奖励记录下来,形成样本。
- 训练模块用这批样本更新策略模型。
这四步循环往复。传统 RL 里这个循环运行在仿真环境中,动作空间可能是离散的按键或连续的力矩。在 LLM 对齐场景里,这个循环变成了:策略模型根据 prompt 生成一段回答,规则模型或奖励模型对回答打分,然后系统用 PPO 这类算法更新策略。
关键点在于:第 1 步和第 4 步在计算性质上完全不同。第 1 步是推理,它决定了环境交互的速度;第 4 步是训练,它决定了策略更新的速度。两者消耗的都是 GPU,但负载特征、延迟要求和扩展方式都不一样。
1.2 推理在 RL 中到底做了什么
在 RL 训练管道里,推理不是一个独立环节,而是数据生产的源头。可以这样理解:
- 训练需要样本,样本来自环境交互。
- 环境交互需要策略模型输出动作或文本。
- 每次输出都是一次推理调用。
在传统 RL 中,一次推理可能只输出一个离散动作,计算量很小。但在 LLM 对齐、RLHF(基于人类反馈的强化学习)、智能体(Agent)任务中,情况完全不同。策略模型需要在每个 prompt 下生成几百甚至上千个 token。自回归生成是 token-by-token 的,每生成一个 token 都要跑一遍模型的前向计算,而且后一个 token 依赖前一个 token 的结果,无法像训练那样直接并行。
所以 RL 里的推理,本质上是一次高吞吐、长序列、带依赖关系的批量生成任务。它消耗的算力往往不比训练 step 少,甚至更多。把这一层当成“调用一次模型接口”来理解,会严重低估它的资源配置需求。
1.3 为什么推理会成为 RL 扩展的瓶颈
推理成为瓶颈,有三个直接原因。
第一个原因是串行依赖。训练更新可以等数据攒够后再做,但 rollout 生成必须等策略模型当前版本给出输出。如果策略每更新一步就换一个版本,那么 rollout 和训练天然是串行关系:训练完,才能拿新模型去生成下一批数据。
第二个原因是生成成本高。对于一次生成 512 个 token 的请求,推理引擎实际执行的 forward 次数是 512 次。同样的模型,训练时一个 step 可以把 128 条样本作为 batch 一次算完,推理时却很难把 128 条长度不同的请求高效合并。自回归生成让 GPU 的算力利用率天然低于训练。
第三个原因是扩展方向不匹配。训练集群适合水平扩展训练并行度,比如数据并行、张量并行、流水线并行,目标是缩短单次 step 时间。推理集群的扩展目标则是提升每秒生成的 token 数,同时不把单请求延迟拖得太高。两个优化目标放到同一批机器上,往往会互相干扰。
一个简单的估算就能看出问题。假设每轮需要生成 128 条长度 512 的样本,推理引擎单卡吞吐是 2000 token/s,那么一轮 rollout 需要128 * 512 / 2000 = 32.8秒。如果训练 step 只需要 1 秒,那么每轮迭代的等待时间绝大部分消耗在推理上。这时候把训练 batch 翻倍,节省的时间远不如把推理吞吐翻倍来得明显。
2. 训练与推理混跑为什么不可持续
2.1 资源竞争:算力、显存和带宽
训练和推理混跑最容易出的问题,是显存冲突。训练时模型参数、优化器状态、梯度、激活值会把显存占满,尤其是 Adam 优化器需要额外保存一阶和二阶动量,显存开销往往是模型参数的好几倍。推理服务部署在同一张卡上时,要么被迫缩小 batch,要么直接 OOM。
算力分配同样有问题。训练 step 是周期性的、可预测的,一批数据算完才进入下一批。推理请求则是突发性的,环境并行度提高时,请求会集中到达。混跑时,训练抢占算力会让推理延迟飙升,推理突发又会打断训练的稳定节奏,最终两边都变慢,而且很难通过调参解决。
另一个容易忽略的是显存带宽。自回归生成每个 token 都要把模型参数从显存读到计算单元,对显存带宽的消耗很高。训练的大矩阵乘法同样依赖高带宽。两者混跑时,即使显存容量没有打满,带宽也可能成为新的瓶颈。
2.2 弹性差异:训练是长任务,推理是高频短请求
训练和推理的生命周期特征差异很大。一个训练任务通常持续几小时甚至几天,期间 GPU 利用率应该保持稳定,不适合频繁扩缩容。推理则不同,rollout 的请求量会随着环境并行度、样本数量、生成长度的变化而波动。
如果两者在同一个集群里,扩缩容策略会互相打架:为了让推理跟上突发负载而扩容,会把训练节点挤掉;为了确保训练节点不被抢占,推理又无法及时扩容。结果是系统变得难以预测,每次改动都像在调整跷跷板。
正确的做法是允许两者独立扩缩容。训练集群按训练任务数量和数据并行度控制;推理集群按队列积压、请求吞吐、GPU 利用率等指标控制。这样训练慢不会拖垮推理,推理突发也不会打断训练。
2.3 一个可量化对比:训练与推理的负载特征
| 维度 | 训练 | 推理 |
|---|---|---|
| 负载类型 | 持续、稳态、可预测 | 突发、波动、与采样并行度相关 |
| 优化目标 | 缩短 step 时间、提升收敛效率 | 提升 token 吞吐、控制生成延迟 |
| GPU 显存 | 参数 + 优化器态 + 梯度 + 激活值 | 参数 + KV Cache + 中间激活 |
| 单位 | step/s | token/s、请求/s |
| 弹性需求 | 低,短时间扩缩容收益小 | 高,需要应对 rollout 批次波动 |
| 故障影响 | 训练中断,浪费已算步数 | 数据断供,训练等待空转 |
这张表说明,训练和推理虽然用同样型号的 GPU,但它们对资源管理系统的要求完全不同。把两者拆开,不是为了增加架构复杂度,而是让每一类负载都能按自己的规律扩展。
3. 把推理独立出来:架构分解
3.1 分离训练集群与推理集群
独立扩展推理的第一步,是在物理或逻辑上把训练和推理分开。物理隔离适合生产环境,推理服务独占一批 GPU,不参与训练调度。逻辑隔离则适合资源有限的环境,使用 Kubernetes 的节点池、资源配额或调度标签,让训练 Pod 和推理 Pod 互不抢占。
拆分后的数据流必须明确:
- 训练端在策略更新完成后,把最新模型权重发布到模型仓库。
- 推理集群加载新版本模型,开始服务。
- 环境交互客户端向推理集群发起 rollout 请求。
- 推理集群返回生成结果,客户端把样本写入缓冲存储。
- 训练端从缓冲存储异步消费样本,计算奖励、做策略更新。
这个流程里,训练和推理之间不再是直接调用,而是通过“模型仓库”和“样本缓冲”解耦。好处是训练端不必等待推理完成,推理端也不必依赖训练端释放资源。
3.2 推理服务化:用接口屏蔽生成细节
推理独立扩展的前提,是把它变成一个可以水平调用的服务。常见的做法是提供 gRPC 接口,因为 gRPC 对流式返回、压测、多语言客户端支持都更友好。HTTP 接口也可以,但长文本生成场景下需要自己处理连接超时和流式响应。
一个最小接口只需要三个能力:接收状态或 prompt、返回动作或文本、附带生成过程的关键信息,比如对数概率、生成长度、请求 ID。这些信息在 RL 训练中是必需的,PPO 计算 advantage 时需要动作的对数概率,样本回放时需要知道每个样本来自哪个策略版本。
接口定义不要过于复杂。先保证一次请求能拿到完整结果,再考虑流式返回。RL 场景里,多数情况下需要的是完整序列,而不是边生成边消费。
3.3 数据流与回传:样本缓冲是关键
推理服务和训练端之间,建议至少隔一层缓冲,避免直接同步调用。理由很简单:rollout 生成速度不稳定,训练消费速度也不稳定,直接同步调用会把两边的抖动互相传递。
缓冲层可以用 Redis Stream、Kafka 或对象存储加索引的方式实现。数据规模不大时,Redis Stream 足够;数据量大、需要回放历史样本时,Kafka 或对象存储更合适。每个样本至少要包含:
- prompt 或状态序列。
- 生成的动作或文本。
- 动作的对数概率。
- 策略版本号。
- 请求 ID 和时间戳。
奖励计算可以放在缓冲之后。奖励模型本身也要跑推理,如果奖励模型和策略模型是同一个模型,可以复用推理集群;如果是独立模型,建议单独部署,避免与 rollout 抢资源。
4. 独立扩展推理的具体做法
4.1 最小推理服务示例:从 proto 到服务端
下面用一个最小 gRPC 示例说明推理服务怎么落地。先定义接口协议:
syntax = "proto3"; package rl_inference; service PolicyInference { rpc Generate(GenerateRequest) returns (GenerateResponse); } message GenerateRequest { string request_id = 1; repeated int32 prompt_ids = 2; int32 max_new_tokens = 3; } message GenerateResponse { string request_id = 1; repeated int32 output_ids = 2; float logprob_sum = 3; int32 num_generated_tokens = 4; }这个协议里,prompt_ids是已经编码好的 token 序列,max_new_tokens限制生成长度,logprob_sum返回整个生成序列的对数概率之和,训练端可以直接用。编码放在客户端做还是服务端做,取决于 tokenizer 和模型是否一起部署。生产环境建议把 tokenizer 和模型放在同一份镜像里,服务端直接接收原始文本,客户端逻辑更简单。
服务端实现非常直接:
class PolicyInferenceService(PolicyInferenceServicer): def __init__(self, model): self.model = model def Generate(self, request, context): prompt = torch.tensor([request.prompt_ids], dtype=torch.long) with torch.no_grad(): output = self.model.generate( prompt, max_new_tokens=request.max_new_tokens, use_cache=True, return_dict_in_generate=True, output_scores=True, ) output_ids = output.sequences[0, prompt.shape[1]:] logprob_sum = self._compute_logprob_sum(output, output_ids) return GenerateResponse( request_id=request.request_id, output_ids=output_ids.tolist(), logprob_sum=logprob_sum, num_generated_tokens=len(output_ids), )这里的关键是use_cache=True。没有 KV Cache,自回归生成每一步都要重新计算之前所有 token 的键值,耗时随序列长度平方增长。打开缓存后,每步只计算新 token 的键值,生成成本从平方级降到线性级。后面的 4.3 小节还会展开讲 KV Cache 和连续批处理。
4.2 动态批处理:让请求排队合并
推理吞吐低,很大程度是因为请求到达不均匀。环境并行的每个 worker 完成上一轮生成的时间不同,发起新请求的时间也不同。如果来一个请求就生成一次,GPU 一直处于小 batch 状态,算力利用率很低。
动态批处理(continuous batching)的思路是:不让单个请求独占一次前向计算,而是把短时间内到达的请求攒起来,凑成较大的 batch 一起算。下面是一段调度伪代码:
def batch_loop(server): batch = [] deadline = None while True: if deadline is None: batch.append(server.queue.get()) deadline = time.time() + server.max_wait_ms / 1000 else: batch.extend(server.queue.get_many( max_batch_size=server.max_batch_size - len(batch), timeout=deadline - time.time(), )) if len(batch) >= server.max_batch_size or time.time() >= deadline: yield batch batch = [] deadline = Nonemax_wait_ms控制请求最多等多久,max_batch_size控制一次最多合并多少个请求。这两个参数决定了吞吐和延迟的平衡:
max_wait_ms调大,能凑到更大 batch,吞吐更高,但单请求延迟变大。max_batch_size调大,单次计算更满,但如果超过了显存或模型并行限制,会直接 OOM。- 如果队列长期为空,说明到达速度小于处理速度,不需要扩容。
- 如果队列长期积压,即使
max_batch_size已经打满,说明需要增加推理实例。
推荐做法是监控队列积压和平均 batch 大小,再决定调参数还是扩实例。不要只凭直觉把等待时间调大,否则 rollout 延迟会拖慢整个训练循环。
4.3 用推理引擎接管生成优化
手写model.generate在验证原型时没问题,但生产环境建议直接使用专门的推理引擎,比如 vLLM、TGI 这类开源方案。原因有三个:
- 它们实现了 PagedAttention 或类似技术,KV Cache 不再是一整块连续显存,而是按页分配,显存利用率显著提升。
- 它们原生支持 continuous batching,不需要自己写队列调度。
- 它们对长序列、并发请求、流式返回做了大量优化,自己做很难达到同等效果。
下面是一份 vLLM 风格配置示例,用于说明关键参数:
inference: engine: vllm model: /models/policy_v1 max_model_len: 8192 gpu_memory_utilization: 0.85 tensor_parallel_size: 4 max_num_batched_tokens: 4096 max_num_seqs: 64 enable_prefix_caching: true| 参数 | 含义 | 调大影响 | 调小影响 |
|---|---|---|---|
max_model_len | 允许的最大序列长度 | 能处理更长样本,但挤占 KV Cache 空间 | 短样本友好,显存更宽裕,长样本会被截断报错 |
gpu_memory_utilization | 推理引擎最多使用显存比例 | 可用 KV Cache 更多,吞吐更高 | 留出显存给其他进程,但增加 OOM 风险 |
tensor_parallel_size | 张量并行的 GPU 数量 | 单请求延迟更低,可服务更大模型 | 多卡之间通信开销增加,小模型可能不划算 |
max_num_seqs | 一次最多并发处理的序列数 | 并发能力更强,显存压力增大 | 并发低,单请求排队时间变长 |
enable_prefix_caching | 是否缓存相同前缀的计算结果 | 重复 prompt 场景节省大量算力 | 关闭后每次都要重算,吞吐下降 |
在 RL 场景里,同一个 prompt 可能被反复用于生成多个候选回答。开启 prefix caching 后,共享前缀的 KV Cache 可以复用,能明显提升有效吞吐。如果环境状态本身前缀变化不大,这个优化尤其值得打开。
4.4 独立扩缩容:按队列和吞吐驱动
推理集群独立部署之后,扩缩容策略要围绕两个核心指标设计:队列积压长度和推理吞吐。
队列积压表示当前请求处理速度跟不上到达速度。如果积压持续增长,说明需要扩容;如果积压长期为零且 GPU 利用率不高,说明可以缩容。在 Kubernetes 环境下,可以基于自定义指标做 HPA:
apiVersion: autoscaling/v2 kind: HorizontalPodAutoscaler metadata: name: rl-inference-hpa spec: scaleTargetRef: apiVersion: apps/v1 kind: Deployment name: rl-inference minReplicas: 2 maxReplicas: 16 metrics: - type: Pods pods: metric: name: inference_queue_depth target: type: AverageValue averageValue: 32这个配置的意思是,当每个推理 Pod 的平均队列积压超过 32 时,自动扩容;积压降下来后,自动缩容。缩容前要注意处理优雅退出,让正在生成的请求跑完,否则会导致样本丢失。
训练端也要配合调整。不要每次都等上一批 rollout 完全结束才开始下一轮训练,而是采用流水线方式:训练当前批次的同时,让推理集群预生成下一批样本。只要推理吞吐和训练速度大致匹配,训练循环就不会空转等待。
5. 验证与排查:怎么知道推理是瓶颈
5.1 需要长期监控的指标
| 指标 | 含义 | 推荐监控方式 | 瓶颈提示 |
|---|---|---|---|
rollout_throughput | 每秒生成的 token 数 | 推理服务计数器 | 低于模型规模预期,说明批次或并行度不够 |
rollout_latency_p50_p95 | 单个请求生成耗时 | 请求日志分位数 | p95 远高于 p50,说明有长尾请求或排队 |
inference_gpu_util | 推理 GPU 利用率 | DCGM 或 Prometheus GPU exporter | 长期低于 60% 时,批次策略或并发不足 |
queue_depth | 推理队列积压 | 服务端指标 | 持续增长时,需要扩容或提升吞吐 |
training_wait_time | 训练端等待样本的时间 | 训练日志 | 占每轮耗时比例超过 50%,推理是主瓶颈 |
kv_cache_usage | KV Cache 使用率 | 推理引擎指标 | 接近上限时,请求排队或 max length 受限 |
这套指标要同时看,不能只看一个。比如 GPU 利用率高,可能只是 batch 打得满,但队列还在增长,说明算力本身不够。再比如 GPU 利用率低,可能不是请求少,而是单请求生成太慢,batch 合并没有生效。
5.2 从训练日志倒推瓶颈
训练日志通常能直接告诉我们瓶颈在哪。一个典型的表现是:
step 100: train_time=0.8s, wait_data_time=320s, samples=256 step 101: train_time=0.9s, wait_data_time=315s, samples=256wait_data_time远大于train_time,说明训练端大部分时间在等 rollout 数据。这时候看推理服务日志,如果显示 batch 一直凑不满、平均 batch 远小于max_batch_size,说明是请求到达模式问题,应该调整动态批处理参数。如果 batch 已经很大但生成耗时还是很高,说明是算力或 KV Cache 限制,应该扩容或调大gpu_memory_utilization。
还有一种情况是训练端和推理端实际没有问题,但中间的数据传输层慢了。特征是推理日志显示请求已完成,返回时间很短,但训练端收到数据的时间比预期晚很多。此时要检查样本缓冲、消息队列或对象存储的写入延迟。
5.3 常见问题与排查表
| 问题现象 | 可能原因 | 检查方式 | 处理建议 |
|---|---|---|---|
| 训练 step 很快,但每轮迭代耗时长 | rollout 生成是主瓶颈 | 对比train_time和wait_data_time | 独立扩展推理集群,提升推理吞吐 |
| 推理 GPU 利用率长期偏低 | batch 太小或请求到达不均匀 | 查看平均 batch 大小和请求到达曲线 | 开启动态批处理,适当调大max_wait_ms |
| 推理与训练混跑时 OOM | 显存被训练进程占满 | 查看 GPU 显存监控、检查进程显存占用 | 物理隔离训练和推理节点,分配独立显存 |
| 单请求生成很慢 | 未开 KV Cache 或gpu_memory_utilization太低 | 检查推理引擎配置和生成日志 | 打开use_cache,调大 KV Cache 空间 |
| 请求排队但 GPU 利用率不高 | 单请求本身计算密集,batch 无法合并 | 查看队列长度和序列长度分布 | 启用连续批处理和 prefix caching |
| rollout 成功但训练迟迟收不到数据 | 缓冲层或消息队列成为新瓶颈 | 检查 Kafka/Redis 的写入延迟和消费 lag | 优化序列化,扩大缓冲分区或换存储 |
| 长样本生成被截断或报错 | max_model_len设置过小 | 检查错误日志中的长度字段 | 按训练数据的最大序列长度合理设置 |
| 扩缩容后请求仍堆积 | 扩容依赖的指标不敏感或冷却时间过长 | 查看 HPA 指标历史和扩容事件 | 缩短指标采集周期,改用队列深度触发扩容 |
每个问题都要按“现象 -> 可能原因 -> 检查方式 -> 处理建议”的顺序排查,而不是看到异常就直接改参数。最常见的错误是批处理参数、显存参数、扩缩容策略一起改,出了问题根本没法判断是哪一步引起的。
6. 落地建议与扩展方向
6.1 学习环境与生产环境的差异
学习环境和生产环境的资源条件不同,落地方式不能照搬同一个模板。
学习环境建议先在一台多卡机器上把流程跑通:训练脚本、推理服务、样本缓冲都部署在同一机架上,用逻辑隔离代替物理隔离。重点是验证接口协议、数据格式和指标监控是否完整,不要急着追求吞吐。此时即使推理和训练混跑,只要清楚瓶颈在哪,就能继续开发。
生产环境至少需要考虑这些额外保障:
- 推理集群独立部署,训练和推理使用不同的节点池。
- 模型版本管理:策略更新后,推理服务如何平滑切换到新版本,如何回滚。
- 样本缓冲必须持久化,不能因为推理实例重启而丢失数据。
- 扩缩容要有上下限,避免 HPA 抖动导致频繁重启。
- 推理服务的监控、告警和日志必须完整,至少要覆盖队列深度、GPU 利用率、生成延迟和错误率。
- 训练端要处理推理集群不可用的情况,比如超时重试、降级使用旧批次数据。
生产环境的优先事项不是把单次生成做到最快,而是让整个数据生产链路稳定、可观测、可回滚。推理扩展做得好不好,最终要看训练循环是否稳定,而不是某一秒的吞吐峰值。
6.2 推理扩展检查清单
在把 RL 训练任务接入新的推理架构之前,可以按这份清单逐项确认:
- 确认策略模型版本在训练端和推理端一致,避免用旧模型生成数据训练新模型。
- 确认推理接口返回了训练所需的对数概率和策略版本号。
- 确认 KV Cache 已开启,
max_model_len能覆盖最长样本。 - 确认动态批处理参数与请求到达模式匹配,队列不会无限积压。
- 确认推理集群的 GPU 利用率、队列深度、生成延迟都有监控。
- 确认训练端等待数据的时间有日志记录,能识别每轮的瓶颈。
- 确认训练和推理的扩缩容策略互不影响。
- 确认推理实例重启、发布、回滚不会导致批次数据丢失。
- 确认环境并行度提高时,推理集群能通过扩容跟上。
- 确认长序列、大批量请求不会触发超时或 OOM。
6.3 后续可以深入的方向
独立扩展推理是一个架构起点,不是终点。沿着这个方向继续深入,可以考虑以下几个层次。
第一个层次是推理本身提速。规格化解码(speculative decoding)可以用一个小模型先草拟多个 token,再由大模型一次性验证,能明显降低生成延迟。这类技术适合单请求延迟敏感的 RL 场景。
第二个层次是跨任务资源调度。真实项目里往往同时跑多个 RL 任务,每个任务的模型版本、序列长度、吞吐要求都不一样。此时需要一个统一调度层,按优先级把不同任务的推理请求分配到共享推理集群,而不是给每个任务固定一批 GPU。
第三个层次是环境交互的分布式化。RL 的推理只是数据生产链路的一环,环境仿真的并行度、奖励模型的计算量、样本回放存储的吞吐都会成为新的瓶颈。把推理独立出来之后,下一步通常是优化整个数据生产管道,而不是只盯着模型生成这一环。
对新手来说,最有价值的练习不是一开始就搭完整架构,而是先在自己的训练脚本里打印train_time和wait_data_time,真实感受到推理等待占了多少比例。这个简单动作,比任何架构图都更能帮助你理解为什么 RL 的瓶颈在推理,以及为什么要单独扩展它。