ARTICLE DETAIL

资讯详情

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

CANN 系列之 RL On-Policy 长尾推理均衡调度引擎:基于 vLLM 的序列级 Rollout Rebalance 原理与实战

CANN 系列之 RL On-Policy 长尾推理均衡调度引擎:基于 vLLM 的序列级 Rollout Rebalance 原理与实战 CANN 系列之 RL On-Policy 长尾推理均衡调度引擎基于 vLLM 的序列级 Rollout Rebalance 原理与实战【免费下载链接】cann-recipes-train本项目针对LLM与多模态模型训练业务中的典型模型、加速算法提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-train导读本文围绕 CANN 训练样例仓库cann-recipes-train中的 RL On-Policy 长尾推理优化方案展开详细剖析一套面向 RLHF/GRPO Rollout 阶段的序列级Request/SEQ 级均衡调度引擎。该方案针对单轮推理同步场景下响应长度长尾分布导致的 DP 组“木桶效应”通过跨 Rank 的请求搬迁与 KV Cache 恢复、结合多档位静态图推理将整组推理档位快速下调从而显著压缩长尾阶段的推理耗时。读完本文你将掌握该引擎的负载均衡算法、数据搬迁与结果还原的完整调用链以及如何通过环境变量与配置文件在 verl Megatron vLLM(Ascend) 训练链路中快速使能该特性。1. 背景Rollout 阶段的“木桶效应”与长尾负载不均在 RLHF 的 Rollout采样推理阶段同一个 DP数据并行组内各 Rank 会并行处理一批 Prompt。由于输入 Prompt 所生成的响应Response长度天然存在长尾分布少数极长的生成任务会拖慢整个 DP 组的进展——这些超长序列长时间占用对应 Rank 的算力与 KV Cache而处理短序列的节点在完成计算后只能进入长时间的闲置等待造成算力浪费。正如关联文档 llm_rl/deepseek/verl_patches/features/rollout_optimize/README.md 所指出的长尾问题优化的本质是 RL 训练系统的负载均衡。本优化针对“单轮推理的同步场景”优化目标明确聚焦于提升进入长尾状态后的推理效率——即在部分 Rollout 提前结束后各 Rank 之间的待处理序列数出现明显不均时对未结束的 Rollout 进行负载均衡的策略分析和重调度从而提升计算资源利用率和长尾状态下的推理吞吐。2. 方案总览与前置依赖本优化的核心目标对应 README.md 的“1.2 解决方案”在 On-Policy 场景中针对部分 Rollout 提前结束导致各 Rank 间负载不均时对未结束的 Rollout 进行负载均衡的策略分析与重调度。方案包含三个关键功能实现Rebalance 条件检测与调度策略生成——周期性检查 DP 组内各 Rank 的剩余请求分布计算最优迁移方案Request(SEQ) 级的数据搬迁与恢复含对应 KV Cache——将被迁移请求的 prompt、已生成 token、logprob 缓存与 KV Cache 完整搬到目标 Rank保证断点续推Rollout 后的结果还原——迁移完成的请求在目标 Rank 生成完毕后将结果通过 AllGather 归还给源 Rank并对多采样分支如 n16重新聚合成原始父请求。2.1 前置依赖多档位静态图推理方案依赖 vllm_ascend 0.9.1 提供的torchair_graph_config能力use_cached_graph与graph_batch_sizes支持提前配置多档位 BatchSize 的图并随着剩余 Seq 减少时自动匹配最小 BS 的图进行推理。这意味着当 DP 组内各 Rank 的剩余请求数因重调度而下降时推理可以切换到更小的档位图单步延迟TPOT随之大幅降低。这一依赖关系是实现性能收益的前提具体配置方式见下文“运行时接入”一节中hook.before(LLM, __init__)的实现。3. 核心实现最大档位最小化均衡算法整个引擎的代码位于 llm_rl/deepseek/verl_patches/features/rollout_optimize/rollout_rebalance.py下称rollout_rebalance.py配套配置类在 llm_rl/deepseek/verl_patches/features/rollout_optimize/config.py。3.1 全局状态感知DP 组内剩余请求分布RolloutRebalanceEngine.get_current_state从本地output_processor.request_states采集当前 Rank 仍在执行的请求 ID 列表sync_group_states通过dist.all_gather_object在dp_group内汇聚所有 Rank 的状态def get_current_state(self): return dict( rankself.rank, req_idslist(self.llm_engine.output_processor.request_states.keys()), ) def sync_group_states(self): if self.world_size 1: return [self.get_current_state()] group_states [None for _ in range(self.world_size)] dist.all_gather_object(group_states, self.get_current_state(), groupself.dp_group) return group_states其中world_size取vllm.distributed.parallel_state的 DP 组规模dp_group来自llm_engine。需要说明的是这里的“剩余请求数”只覆盖当前仍在推进未 finish的请求恰好对应长尾阶段中仍然占用算力的那部分负载。3.2 档位映射函数_get_bs均衡算法依赖“档位BS”这一离散化概念。_get_bs将某个 Rank 的剩余请求数映射为需要执行推理时使用的最小档位——从States.graph_batch_sizes的最小档位开始向上遍历返回第一个不小于请求数的档位若超出所有档位则直接返回原始请求数staticmethod def _get_bs(size): for bs in States.graph_batch_sizes[::-1]: if size bs: return bs return sizeStates.graph_batch_sizes的默认值是[64, 32, 16, 8, 4]在enable_rollout_rebalance中会根据torchair_graph_config的max_batch_size做过滤与去重见第 5 节因此实际档位以运行时计算结果为准。3.3 三层目标的最优迁移策略calc_balancing_tasks是引擎的“决策大脑”其最优策略被明确定义为三层目标主要目标将整个 DP 组所需的最大档位max_bs降至最低次要目标在满足 1 的前提下使迁移的请求数量cost最少补充目标在满足 1 和 2 的前提下使各 Rank 间的数据搬迁流向尽可能均匀避免“多对一”阻塞。算法执行步骤与源码一一对应计算当前最大档位max_bs_before与平均请求数avg_bs从最小档位向大遍历States.graph_batch_sizes寻找满足avg_bs target_bs max_bs_before的最小目标档位max_bs_next找不到则返回空任务无需搬迁以max_bs_next为基准划分捐赠方donorreq_cnt - max_bs_next 0的 Ranksurplus 为超额请求数接收方receiver差额为负的 Rankcapacity 为可接收量生成迁移任务清单循环将 donor 的请求分发给 receiver优先分发给容量最大的 receiver并通过donor_index (donor_index 1) % len(donors)在多个 donor 之间轮转保证流向均匀、避免阻塞。while True: donor_index 0 for receiver in sorted(receivers, keylambda r: r[capacity], reverseTrue): donor donors[donor_index] num_to_move min(donor[surplus], receiver[capacity]) balancing_tasks [dict( from_rankdonor[rank], to_rankreceiver[rank], req_idreq_id ) for req_id in donor[req_ids][:num_to_move]] donor[req_ids] donor[req_ids][num_to_move:] donor[surplus] - num_to_move receiver[capacity] - num_to_move donor_index (donor_index 1) % len(donors) donors [x for x in donors if x[surplus]] if not donors: break receivers [x for x in receivers if x[capacity]] return balancing_tasks该算法本质上是“以最小化整组推理档位为目标”的贪心均衡迁移成本迁移请求数被严格约束在“降档所需的最小数量”因此能最大程度减少通信开销。3.4 触发节奏check与周期检测引擎不会在每个 step 都做重调度而是通过CheckCounter定义于 llm_rl/deepseek/verl_patches/features/rollout_optimize/utils.py控制检测频率。CheckCounter.check()在累计调用达到阈值时返回 True 并重置计数阈值由RolloutRebalanceConfig.check_interval默认 1000 个 step决定def check(self): need_profile self.ProfileCache.enable and self.ProfileCache.counter.check() need_rebalance_check self.rebalance_counter.check() if not (need_profile or need_rebalance_check): return start datetime.datetime.now(tzdatetime.timezone.utc) group_states self.sync_group_states() if need_profile: self.profile(group_states) if not need_rebalance_check: return schedule_tasks self.calc_balancing_tasks(group_states) ... if schedule_tasks: rank_log_info(f[RebalanceScheduleTasks][Cnt{len(schedule_tasks)}]) for schedule_task in schedule_tasks: rank_log_info(f[Rebalance][ReqId{schedule_task[req_id]}] f[Src{schedule_task[from_rank]}][Dst{schedule_task[to_rank]}]) self.all_to_all_v_tasks(schedule_tasks)check()由hook.after(LLMEngine, step)在每个 step 之后触发且仅当达到检测阈值时才会发起 AllGather 与搬迁将调度本身的开销限制在可接受范围。从源码结构看need_profile与need_rebalance_check共用一个 AllGather 结果避免重复通信。4. 序列级数据搬迁元数据 AllToAll KV Cache 点对点传输调度任务生成后all_to_all_v_tasks负责真正执行跨 Rank 的请求搬迁。其设计巧妙地将“轻量元数据”与“重量 KV Cache”分开传输元数据走 AllToAllV每个请求被打包为一个RebalanceRequestTaskget_transfer_dict()产出包含req_id、src_rank、prompt_token_ids、output_token_ids、max_tokens、logprobs_processor_cache、layers_kv_cache_shapes、send_time的字典通过pickle.dumps转成字节流先交换各 Rank 的发送长度dist.all_to_all_single(remote_sizes, local_sizes)让每个 Rank 知道从谁那里要接收多少字节再交换真实数据拼接本地所有目标数据后用第二次all_to_all_single分发接收端按remote_sizes切分并pickle.loads还原为请求任务列表KV Cache 走点对点dist.send/dist.recv元数据交换完成后源 Rank 调用send_kv_caches逐块发送目标 Rank 在load_received_tasks中按layers_kv_cache_shapes分配空张量后dist.recv接收。def all_to_all_v_tasks(self, schedule_tasks): objects_to_send [[] for _ in range(dist.get_world_size())] send_tasks [] for schedule_task in schedule_tasks: if self.rank schedule_task[from_rank]: request_task RebalanceRequestTask(self.llm_engine).load_by_req_id(schedule_task[req_id]) send_tasks.append((request_task, schedule_task[to_rank])) objects_to_send[schedule_task[to_rank]].append(request_task.get_transfer_dict()) request_task.trigger_abort() ... dist.all_to_all_single(remote_sizes, local_sizes) ... dist.all_to_all_single( output_tensor, input_tensor, output_split_sizesremote_sizes.tolist(), input_split_sizeslocal_sizes.tolist(), ) received_tensor torch.split(output_tensor, remote_sizes.tolist()) received_tasks [] for rank_data in received_tensor: received_tasks pickle.loads(rank_data.to(cpu).numpy().tobytes()) self.send_kv_caches(send_tasks) self.load_received_tasks(received_tasks)4.1 源 Rank 侧状态采集与中止RebalanceRequestTask.load_by_req_id在源 Rank 侧完成请求“快照”从model_runner.input_batch的block_table定位该请求对应的 KV Cache block 编号过滤掉 padding 的 0从output_processor.request_states取出prompt_token_ids与 logprob 处理器中的cumulative_logprob和逐 token 的 logprob 缓存从scheduler.requests取出output_token_ids与max_tokens按层采集 KV Cache遍历所有层layer的所有 cache blocktorch.stack成layers_kv_cache_blocks即该请求独占的完整 KV Cache 切片。# request级的kvCache采集 self.layers_kv_cache_blocks [] for cache_block_index in range(len(self.global_kv_caches[0])): self.layers_kv_cache_blocks.append( torch.stack([layer[cache_block_index][request_block_table] for layer in self.global_kv_caches]) )快照完成后调用trigger_abort()将请求从当前 Rank 的引擎中摘除同时调用llm_engine.abort_request与engine_core.abort_requests释放该 Rank 的调度席位。4.2 目标 Rank 侧KV Cache 接收与状态还原load_received_tasks中除接收 KV Cache 外还通过RebalanceRequestTask.load_by_transfer_info还原请求上下文随后调用trigger_load()把请求“塞回”目标 Rank 的引擎。trigger_load是断点续推的关键其调用链包括llm_engine.add_request以src_rank_{src_rank}_{req_id}作为新请求 ID 加入引擎该前缀是后续结果还原的识别标记recover_request_state还原output_processor中的cumulative_logprob、已生成的token_ids与 logprob 列表recover_scheduler_request恢复scheduler.requests中的max_tokens、num_computed_tokens并把请求从waiting队列移回running队列使其继续参与调度recover_model_runner通过scheduler.kv_cache_manager.allocate_slots在目标 Rank 申请新的 KV Cache block重建CachedRequestState最后把搬来的 KV Cache按新 block 编号写回model_runner.kv_caches# kvCache还原 for layer_index, layer_caches in enumerate(self.model_runner.kv_caches): reload_indexes list(range(len(new_block_ids))) for i, cache_block in enumerate(self.layers_kv_cache_blocks): layer_caches[i][new_block_ids] cache_block[layer_index][reload_indexes]至此请求在目标 Rank 上以“已生成部分 token、已计算 KV Cache”的状态无缝续推无需重新 Prefill搬迁对推理语义完全透明。注关联文档描述的前置依赖 vllm_ascend 0.9.1 版本差异会反映在实现细节上例如 qwen3 版本使用vllm_ascend.worker.npu_input_batch.CachedRequestState与vllm_version_is做版本兼容见 llm_rl/qwen3/verl-mindspeed/patches/verl/features/rollout_optimize/rollout_rebalance.py 的recover_model_runner说明该方案在适配不同 vLLM(Ascend) 版本时保持了良好的兼容性设计。5. Rollout 结果还原跨 Rank 输出聚合迁移出去的请求在目标 Rank 上完成后其输出必须回到源 Rank才能与源 Rank 本地未迁移的采样输出一起构成完整的 Response 集如n16的 16 条采样。这一过程由recover完成按前缀分流_split_rebalance_outputs依据request_id是否以src_rank_开头将引擎输出分成“本地输出current_outputs”与“迁移输出rebalance_outputs”并从迁移输出的 request_id 中解析出src_rank与原始req_id同时缓存该输出的prompt_token_ids、token_ids、cumulative_logprob与逐 token logprobAllGather 汇聚迁移输出get_rebalance_outputs通过dist.all_gather_object把各 Rank 完成的迁移输出汇聚到每个 Rank再按src_rank过滤出属于本 Rank 的那部分重建 RequestOutput_build_rebalance_request_output用解析出的字段重新构造RequestOutputfinish_reasonstop、finishedTrue并还原每条采样的CompletionOutput与Logprob父请求聚合引擎初始化时RolloutRebalanceEngine.__init__会先将request_states中各请求的parent_req置空以支持单条 Seq 独立搬迁recover末尾则按 request_id 前缀将同属一个父请求的多条采样重新合并到同一个RequestOutput.outputs中从而还原父子关系。def recover(self, outputs): # 将进行了rebalance迁移的outputs通过all_gather还原到源rank rebalance_outputs, current_outputs self._split_rebalance_outputs(outputs) for rebalance_output in self.get_rebalance_outputs(rebalance_outputs): rank_log_info(f[RecvReqId{rebalance_output[req_id]}], forceTrue) current_outputs.append(self._build_rebalance_request_output(rebalance_output)) ... parent_request.outputs request_output.outputs ... return current_outputs_map.values()在 qwen3 版本的实现中这一还原逻辑通过hook.after(LLM, _run_engine)挂接将States.outputs_cache中缓存的迁移输出合并回最终 outputs与 deepseek 版本的recover调用点不同但语义一致说明该特性在仓库内存在多次演进版本。6. 运行时接入零侵入 Hook 机制与使能入口6.1 使能入口enable_rollout_rebalanceenable_rollout_rebalance是整个特性的总入口见 rollout_rebalance.py 末尾它不修改 vLLM 源码而是基于 utils.py 提供的装饰器式 Hook 工具hook.before/hook.after对关键类方法做运行时增强hook.before(LLM, __init__)在 LLM 初始化前改写additional_config[torchair_graph_config]。首先读取graph_batch_sizes[0]作为max_batch_size然后对配置档位做过滤与降序排序保留max_batch_size为主档位剔除不小于它的档位写入States.graph_batch_sizes当multi_graphTrue时更新graph_batch_sizes_initFalse、use_cached_graphTrue、graph_batch_sizesgraph_batch_sizes从而打开 vllm_ascend 的多档位静态图能力——这正是第 2 节前置依赖的落地位置hook.before(LLM, _run_engine)创建RolloutRebalanceEngine实例传入check_interval启动ProfileCache若profileTrue并缓存首个请求的sampling_params供后续trigger_load复制采样策略hook.after(LLMEngine, step)每个 step 后调用States.rebalance_engine.check()并把src_rank_前缀的迁移输出先缓存到States.outputs_cache避免被当前 Rank 当作普通输出误处理在 qwen3 版本中还有hook.before/after(EngineCore, execute_model_with_error_logging)的 profile 钩子用于统计每个 step 的 Prefill/Decode 数量、最大/最小生成长度与单步耗时。6.2 在 verl 训练链路中启用按关联文档 README.md 的“2.1 初始化配置”在verl/workers/megatron_workers.py的ActorRolloutRefWorker.init_model方法开头追加以下代码通过环境变量ROLLOUT_REBALANCE_ENABLE1使能本特性register(dispatch_modeDispatch.ONE_TO_ALL) def init_model(self): if os.getenv(ROLLOUT_REBALANCE_ENABLE, 0) ! 0: from features.rollout_optimize.rollout_rebalance import enable_rollout_rebalance enable_rollout_rebalance()在 llm_rl/deepseek/verl_patches/workers/megatron_workers.py 中可以看到该调用点已被合入actor_rollout_ref_init_model第 257261 行并在文件末尾通过ActorRolloutRefWorker.init_model actor_rollout_ref_init_model完成方法替换。qwen3 侧的 patchllm_rl/qwen3/verl-mindspeed/patches/verl/0011-verl-feature-enable_rollout_rebalance.patch则采用from patches.verl.features.rollout_optimize import init_rollout_rebalance; init_rollout_rebalance()的方式在文件级直接调用——两种接入方式均以环境变量作为总开关。实际训练脚本中该开关的用法可参考 llm_rl/qwen3/verl-mindspeed/internal/train_grpo_qwen3_resampler_example.shexport ROLLOUT_REBALANCE_ENABLE0 # 0: disable rollout rebalance, 1: enable rollout rebalance6.3 配置项详解可以在 config.py 中直接修改配置或将配置写入 verl 启动的 yaml 中在init_model位置提取后传入enable_rollout_rebalance方法。各配置项含义如下表配置项默认值说明enableTrueRolloutRebalance 特性总开关deepseek 版本为布尔值qwen3 版本为int(os.environ.get(ROLLOUT_REBALANCE_ENABLE, 0))由环境变量驱动check_interval1000间隔多少个 step 进行一次 rebalance 检查控制检测与调度的频率开销multi_graphTrue是否开启多档位编图。若关闭rebalance 依然会按预编图的档位做均衡调度但不会形成明显的性能收益graph_batch_sizes[64, 32, 16, 8, 4]预编图的档位设置运行时会被max_batch_size过滤并降序重排profileTrue是否打印过程中的性能数据profile_interval100profile 打印间隔步长其中multi_graph与graph_batch_sizes直接决定第 2 节的性能前提只有多档位编图开启时负载均衡带来的“剩余 Seq 减少”才能转化为“更低档位图推理”的 TPOT 收益。6.4 运行观测Profile 与日志RolloutRebalanceEngine.profile周期性输出 DP 组各 Rank 的档位映射[BSMap]、剩余序列数[SeqCntMap]、当前最大档位[CurrentMaxBS]、单步耗时与 TPOT[ProfileStepCost]、[ProfileDuration]、[TPOT]并在档位变化时打印[MaxBSChanged: x - y]。迁移过程通过rank_log_info打印[RebalanceScheduleTasks]、[TaskSendKvCache]、[ReceivedTask]含跨 Rank 传输耗时Costxxms、[ReceivedKvCache]、[ReceivedTaskLoaded]等关键日志便于在训练日志中直接观测重调度发生时机与开销。rank_log_info默认只输出 rank_0 的打印forceTrue时所有 Rank 均输出见 utils.py。7. 使能效果与实验配置关联文档README.md 的“1.3 实验结果”在Atlas A3 集群 128 卡环境上的实验表明本方案开启后单轮推理耗时从6200s 左右优化到约 2300s性能收益达 57%62%。实验配置如下模型DeepSeekV3数据集open-r1/OpenR1-Math-220Kdata.train_batch_size512data.max_response_length32768actor_rollout_ref.rollout.n16TP2; DP128。性能收益来源分析收益主要来自单个 step 的 TPOT 性能差距。默认场景下 TPOT 会从 125ms 上升到 200ms而通过使能 Rebalance 并配合多档位编图能在推理长度仅 1~2K 时就快速将推理档位降低让单个 step 的 TPOT 降低到 60ms 量级。在长尾场景下随着剩余序列不断减少性能差距被持续放大——这正是“档位下调 负载均衡”组合拳的核心价值。此外仓库中的配套文档 docs/features/rollout_rebalance.md 记录了该方案在Qwen3 235B / Atlas A3 集群 64 卡 / deepscaler 数据集 / TP4、DP32场景下的另一组实验单轮推理耗时从约 10200s 优化到约 6100s性能收益约 60%TPOT 同样从 125ms→200ms 的区间被压制到 60ms 量级。两组合一可以推断该方案对 DeepSeekV3 与 Qwen3 等大规模 MoE 模型的长尾 Rollout 场景均有显著收益且收益量级与模型规模无强相关。8. 适用前提与注意事项场景限定本方案针对单轮推理的同步场景On-Policy Rollout优化目标聚焦长尾状态下的推理效率在多轮交互或非同步推理场景下的效果需要另行评估版本依赖方案依赖 vllm_ascend 0.9.1 的torchair_graph_config.use_cached_graph与graph_batch_sizes多档位编图能力multi_graphTrue是获得显著性能收益的前提关闭时均衡调度仍生效但 TPOT 不会因降档而改善通信开销重调度涉及跨 Rank 的元数据 AllToAll 与 KV Cache 点对点传输check_interval需要按训练规模权衡调度频率与通信成本从源码日志设计看迁移过程对每次发送/接收/加载都做了毫秒级耗时统计便于在线评估该开销DP 规模sync_group_states、get_rebalance_outputs均在dp_group内做 AllGatherDP 规模越大单次状态同步与结果汇聚的通信量越高适合与 rollout 的n采样数配置共同规划使能方式推荐通过ROLLOUT_REBALANCE_ENABLE1环境变量在 verl 训练脚本中开关参考 megatron_workers.py 与 train_grpo_qwen3_resampler_example.sh避免改动 verl 主流程源码。9. 小结序列级 Rollout Rebalance 引擎为 RL On-Policy 长尾推理提供了一个完整的“检测—决策—搬迁—还原”闭环以 DP 组全局状态感知为基础以最大档位最小化为目标的均衡算法生成最小代价迁移方案借助元数据 AllToAllV KV Cache 点对点传输实现请求级断点续推再通过 AllGather 与父子关系重建完成结果还原最后配合 vllm_ascend 多档位静态图让整组推理档位随剩余序列减少而快速下调。它回答了长尾 Rollout 场景下“如何让每一个 Rank 始终满载运行”的问题是 CANN 平台在 RL 训练工程化方向上的一个高价值样例可在 DeepSeekV3、Qwen3 等大规模模型的 GRPO/DAPO 训练中直接参考使用。进一步阅读仓库内相关源码与文档路径汇总——方案文档llm_rl/deepseek/verl_patches/features/rollout_optimize/README.md核心实现llm_rl/deepseek/verl_patches/features/rollout_optimize/rollout_rebalance.py、llm_rl/deepseek/verl_patches/features/rollout_optimize/config.py、llm_rl/deepseek/verl_patches/features/rollout_optimize/utils.pyverl 接入点llm_rl/deepseek/verl_patches/workers/megatron_workers.pyQwen3 侧演进版本与 patchllm_rl/qwen3/verl-mindspeed/patches/verl/features/rollout_optimize/、llm_rl/qwen3/verl-mindspeed/patches/verl/0011-verl-feature-enable_rollout_rebalance.patch训练脚本示例llm_rl/qwen3/verl-mindspeed/internal/train_grpo_qwen3_resampler_example.shQwen3 235B 实验报告docs/features/rollout_rebalance.md【免费下载链接】cann-recipes-train本项目针对LLM与多模态模型训练业务中的典型模型、加速算法提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-train创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表