Prhub

#6354 [trainer,data] fix: Support merging `extra_info` for `main_ppo_sync` & update TQ dependency

原始 PR 作者 0oshowero0 合并时间 2026-05-18 08:43 文件变更 11 提交数 4 评论 4 代码增减 +141 / -61

执行摘要

修复同步 Trainer 中 extra_info 合并问题并升级 TransferQueue

由于 #6270 为每个 rank 添加了独立的性能指标(如 perf/max_memory_allocated_gb)到 extra_info,导致 BatchMeta.concat 在遇到不同 rank 下同一 key 的值不同时触发 ValueError。PR希望修复此问题,使同步 Trainer 能正常使用 per-rank metrics。

值得精读,特别是 BatchData.concat 的手动合并策略展示了如何在三方库升级(TQ v0.1.7)后处理兼容性问题。同时可学习清理调试代码和配置扩展的最佳实践。

讨论亮点

Review 中 gemini-code-assist 指出手动合并 extra_info 可能导致列表长度不一致问题,且 TQ v0.1.7 已原生支持合并,建议直接利用新版本能力。但作者最终保留了手动合并(改为 plain dict + setdefault),并清理了调试打印。wuxibin89 要求清理 setup.py 中残留的 TransferQueue 项和 main_ppo_sync.py 中的调试打印,均已执行。

实现拆解

  1. 修改 verl/protocol.pyBatchData.concat 处理 BatchMeta 的逻辑:在调用 BatchMeta.concat 之前,遍历所有 BatchMeta,将其 extra_info 收集到 merged_extra_info(dict of lists),然后清空每个 BatchMetaextra_info 以通过 BatchMeta.concat 的检测,最后将合并后的 merged_extra_info 赋给结果 KVBatchMetaextra_info
  2. 清理 verl/trainer/main_ppo_sync.py:移除所有 _compute_*_compute_advantage 中的调试时间打印和 import,使代码简洁;同时添加 try-finally 块确保 TransferQueue 在运行结束时正常关闭。
  3. 依赖管理变更:从 setup.py 中删除可选的 transferqueue extra,直接将 TransferQueue==0.1.7 加入 install_requires(写入 requirements.txtrequirements-npu.txt),使其成为默认依赖。
  4. 配置更新:在 verl/trainer/config/ppo_trainer.yaml 以及所有 _generated_* yaml 中扩展 transfer_queue 配置,新增 metrics(支持 Prometheus 指标导出)和 MooncakeStore(RDMA 兼容后端)子配置。
  5. 文档更新:docs/data/transfer_queue.md 同步更新以反映 v0.1.7 的新功能。
文件 模块 状态 重要度
verl/protocol.py 数据协议 modified 5.67
verl/trainer/main_ppo_sync.py 同步训练 modified 6.74
setup.py 安装配置 modified 4.89
verl/trainer/config/ppo_trainer.yaml 训练配置 modified 4.92
requirements.txt 依赖列表 modified 1.75

关键符号

BatchData.concat _compute_old_log_prob _compute_ref_log_prob _compute_values _compute_advantage

关键源码片段

verl/protocol.py core-logic

核心修复:在 BatchData.concat 中手动合并 extra_info,解决 per-rank metrics 导致 concat 失败的 bug。

def concat(self):
    '''Concat the wrapped list of data items into a single result.'''
    data = self._data
    if not data:
        raise ValueError('Cannot concatenate an empty list of data items.')
    sample = data[0]
    if isinstance(sample, ray.ObjectRef):
        return DataProtoFuture.concat(data)
    if isinstance(sample, TensorDict):
        from verl.utils.tensordict_utils import concat_tensordict
        return concat_tensordict(data)
    if isinstance(sample, BatchMeta):
        # 手动合并每个 rank 的 extra_info,避免原生 concat 因值不一致而失败
        merged_extra_info = {}
        for meta in data:
            for k, v in meta.extra_info.items():
                merged_extra_info.setdefault(k, []).append(v)
            meta.extra_info = {} # 清空以绕过原生检查
        batch_meta = BatchMeta.concat(data)
        batch_meta.extra_info = merged_extra_info # 恢复合并后的信息
        from verl.utils.transferqueue_utils import batch_meta2kv_batch_meta
        return batch_meta2kv_batch_meta(batch_meta)
    return type(sample).concat(data)
verl/trainer/main_ppo_sync.py core-logic

清理调试打印并添加 TransferQueue 关闭保护,提高代码质量和健壮性。

def _compute_old_log_prob(self, batch: KVBatchMeta, metrics: dict) -> KVBatchMeta:
    # 旁路模式直接使用 rollout_log_probs
    rollout_corr_config = self.config.algorithm.get('rollout_correction', None)
    bypass_recomputing_logprobs = rollout_corr_config and rollout_corr_config.get('bypass_mode', False)
    if bypass_recomputing_logprobs:
        data = tq.kv_batch_get(keys=batch.keys, partition_id=batch.partition_id,
                               select_fields=['rollout_log_probs'])
        data['old_log_probs'] = data.pop('rollout_log_probs')
        tq.kv_batch_put(keys=batch.keys, partition_id=batch.partition_id, fields=data)
        return
​
    batch.extra_info.update({
        'calculate_entropy': True,
        'compute_loss': False,
        'temperature': self.config.actor_rollout_ref.rollout.temperature,
    })
    output = self.actor_rollout_wg.compute_log_prob(batch)
    assert len(output) == len(batch)
​
    fields = ['entropy', 'log_probs', 'response_mask']
    if self.config.actor_rollout_ref.rollout.calculate_log_probs:
        fields.extend(['responses', 'rollout_log_probs'])
    data = tq.kv_batch_get(keys=batch.keys, partition_id=batch.partition_id,
                           select_fields=fields)
    # 已移除 debug 时间打印
​
    data['old_log_probs'] = response_from_nested(data.pop('log_probs'), data['response_mask'])
    data['entropy'] = response_from_nested(data.pop('entropy'), data['response_mask'])
    batch = tq.kv_batch_put(keys=batch.keys, partition_id=batch.partition_id,
                            fields=data.select('old_log_probs', 'entropy'))
    # 已移除 debug 时间打印
​
    data = DataProto(batch=data.to_padded_tensor())
    # ... 后续指标计算 ...
    return batch

评论区精华

手动合并 extra_info 的风险与冗余 设计

gemini-code-assist 指出手动合并可能导致列表长度不一致(若 key 缺失),且 TQ v0.1.7 已原生支持合并,建议移除手动逻辑。liziniu 评论认为这是正确的修复点,但建议使用 plain dict 避免副作用。

结论:作者保留了手动合并但改用 plain dict + setdefault,未完全采用 TQ 原生能力,但通过了 review。 · 已解决

清理 setup.py 中的 TransferQueue 条目 style

wuxibin89 在 review 中评论要求清理 setup.py 中残留的 TransferQueue 条目。

结论:作者在最终提交中删除了 TRANSFERQUEUE_REQUIRES 和相关 extras。 · 已解决

清理 main_ppo_sync.py 中的调试打印 style

wuxibin89 要求清理该文件中的 debug print。

结论:作者移除了所有 _compute_* 函数中的时间打印。 · 已解决

风险与影响

  1. verl/protocol.py 中手动合并 extra_info:如果不同 rank 的 extra_info 键集合不一致,合并后的列表长度将不同,可能导致下游索引错误。但同步 Trainer 中各 rank 配置相同,此风险较低。
  2. 将 TransferQueue 变为默认依赖:可能增加非 NPU 环境的安装负担或依赖冲突,且失去可选性。
  3. 无新增测试覆盖,回归依赖现有 CI。

对用户:修复了一个导致同步 Trainer 崩溃的问题,允许 per-rank metrics 正常收集。TransferQueue 自动安装,无需手动添加 extra。配置新增 metricsMooncakeStore,为可观测性和高性能后端提供基础。对系统:新增依赖可能略微增加安装体积。对团队:代码更干净,配置更丰富。

手动合并 extra_info 的索引风险 无新增测试覆盖

关联 Issue

#6270 [worker] feat: support log memory in engine worker

完整报告

参与讨论