执行摘要
- 一句话:修复同步 Trainer 中 extra_info 合并问题并升级 TransferQueue
- 推荐动作:值得精读,特别是
BatchData.concat 的手动合并策略展示了如何在三方库升级(TQ v0.1.7)后处理兼容性问题。同时可学习清理调试代码和配置扩展的最佳实践。
功能与动机
由于 #6270 为每个 rank 添加了独立的性能指标(如 perf/max_memory_allocated_gb)到 extra_info,导致 BatchMeta.concat 在遇到不同 rank 下同一 key 的值不同时触发 ValueError。PR希望修复此问题,使同步 Trainer 能正常使用 per-rank metrics。
实现拆解
- 修改
verl/protocol.py 中 BatchData.concat 处理 BatchMeta 的逻辑:在调用 BatchMeta.concat 之前,遍历所有 BatchMeta,将其 extra_info 收集到 merged_extra_info(dict of lists),然后清空每个 BatchMeta 的 extra_info 以通过 BatchMeta.concat 的检测,最后将合并后的 merged_extra_info 赋给结果 KVBatchMeta 的 extra_info。
- 清理
verl/trainer/main_ppo_sync.py:移除所有 _compute_* 和 _compute_advantage 中的调试时间打印和 import,使代码简洁;同时添加 try-finally 块确保 TransferQueue 在运行结束时正常关闭。
- 依赖管理变更:从
setup.py 中删除可选的 transferqueue extra,直接将 TransferQueue==0.1.7 加入 install_requires(写入 requirements.txt 和 requirements-npu.txt),使其成为默认依赖。
- 配置更新:在
verl/trainer/config/ppo_trainer.yaml 以及所有 _generated_* yaml 中扩展 transfer_queue 配置,新增 metrics(支持 Prometheus 指标导出)和 MooncakeStore(RDMA 兼容后端)子配置。
- 文档更新:
docs/data/transfer_queue.md 同步更新以反映 v0.1.7 的新功能。
关键文件:
verl/protocol.py(模块 数据协议;类别 source;类型 core-logic;符号 BatchData.concat): 核心修复:在 BatchData.concat 中手动合并 extra_info,解决 per-rank metrics 导致 concat 失败的 bug。
verl/trainer/main_ppo_sync.py(模块 同步训练;类别 source;类型 core-logic;符号 _compute_old_log_prob, _compute_ref_log_prob, _compute_values, _compute_advantage): 清理调试打印并添加 TransferQueue 关闭保护,提高代码质量和健壮性。
setup.py(模块 安装配置;类别 source;类型 configuration): 将 TransferQueue 从可选依赖变为默认依赖,影响安装行为。
verl/trainer/config/ppo_trainer.yaml(模块 训练配置;类别 config;类型 configuration): 扩展 TransferQueue 配置以支持 metrics 和 MooncakeStore,为可观测性和高性能后端奠定基础。
requirements.txt(模块 依赖列表;类别 infra;类型 configuration): 添加 TransferQueue==0.1.7 作为默认依赖。
关键符号:BatchData.concat, _compute_old_log_prob, _compute_ref_log_prob, _compute_values, _compute_advantage
关键源码片段
verl/protocol.py
核心修复:在 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
清理调试打印并添加 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
评论区精华
Review 中 gemini-code-assist 指出手动合并 extra_info 可能导致列表长度不一致问题,且 TQ v0.1.7 已原生支持合并,建议直接利用新版本能力。但作者最终保留了手动合并(改为 plain dict + setdefault),并清理了调试打印。wuxibin89 要求清理 setup.py 中残留的 TransferQueue 项和 main_ppo_sync.py 中的调试打印,均已执行。
- 手动合并 extra_info 的风险与冗余 (design): 作者保留了手动合并但改用 plain dict + setdefault,未完全采用 TQ 原生能力,但通过了 review。
- 清理 setup.py 中的 TransferQueue 条目 (style): 作者在最终提交中删除了 TRANSFERQUEUE_REQUIRES 和相关 extras。
- 清理 main_ppo_sync.py 中的调试打印 (style): 作者移除了所有 compute* 函数中的时间打印。
风险与影响
- 风险:
verl/protocol.py 中手动合并 extra_info:如果不同 rank 的 extra_info 键集合不一致,合并后的列表长度将不同,可能导致下游索引错误。但同步 Trainer 中各 rank 配置相同,此风险较低。
- 将 TransferQueue 变为默认依赖:可能增加非 NPU 环境的安装负担或依赖冲突,且失去可选性。
- 无新增测试覆盖,回归依赖现有 CI。
- 影响:对用户:修复了一个导致同步 Trainer 崩溃的问题,允许 per-rank metrics 正常收集。TransferQueue 自动安装,无需手动添加 extra。配置新增 metrics 和 MooncakeStore,为可观测性和高性能后端提供基础。对系统:新增依赖可能略微增加安装体积。对团队:代码更干净,配置更丰富。
- 风险标记:手动合并 extra_info 的索引风险, 无新增测试覆盖
关联脉络
- PR #6270 [worker] feat: support log memory in engine worker: #6270 为 per-rank 添加了 metrics 到 extra_info,导致同步 Trainer 的 concat 失败,此 PR 修复该问题。
参与讨论