尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

RL训练瓶颈在推理?独立扩展推理服务的架构实践

RL训练瓶颈在推理?独立扩展推理服务的架构实践 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/stoken/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], dtypetorch.long) with torch.no_grad(): output self.model.generate( prompt, max_new_tokensrequest.max_new_tokens, use_cacheTrue, return_dict_in_generateTrue, output_scoresTrue, ) output_ids output.sequences[0, prompt.shape[1]:] logprob_sum self._compute_logprob_sum(output, output_ids) return GenerateResponse( request_idrequest.request_id, output_idsoutput_ids.tolist(), logprob_sumlogprob_sum, num_generated_tokenslen(output_ids), )这里的关键是use_cacheTrue。没有 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_sizeserver.max_batch_size - len(batch), timeoutdeadline - 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 环境下可以基于自定义指标做 HPAapiVersion: 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_usageKV Cache 使用率推理引擎指标接近上限时请求排队或 max length 受限这套指标要同时看不能只看一个。比如 GPU 利用率高可能只是 batch 打得满但队列还在增长说明算力本身不够。再比如 GPU 利用率低可能不是请求少而是单请求生成太慢batch 合并没有生效。5.2 从训练日志倒推瓶颈训练日志通常能直接告诉我们瓶颈在哪。一个典型的表现是step 100: train_time0.8s, wait_data_time320s, samples256 step 101: train_time0.9s, wait_data_time315s, samples256wait_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 cachingrollout 成功但训练迟迟收不到数据缓冲层或消息队列成为新瓶颈检查 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 的瓶颈在推理以及为什么要单独扩展它。
返回列表