跳到主要内容

数据截至 (上游 commit cfacd76a0bdd)

02 · 数据面:DataProto 与 TransferQueue

这一章讲什么: verl 里数据长什么样、怎么在 driver 和 worker 之间流动。这里有一次明显的代际更替:V0 用 DataProto(数据随调用传输),V1 换成 TransferQueue(只传 key)。理解这次更替,才能读懂 V1 代码里满屏的 tq.kv_batch_get


1. 第一代:DataProto

1.1 它要解决的小问题

RL 一个 batch 里混着三类东西:

  • 等长张量input_idsattention_maskadvantages……
  • 不等长/非张量:原始对话 raw_promptuid 字符串、多模态图片对象……
  • 整批共享的元信息temperatureglobal_steps、要不要算熵……

如果只用 dict[str, Tensor],第二三类没地方放;如果全用 dict,切分/拼接/搬设备就要到处写循环。

1.2 结构

DataProto 是个三字段 dataclass(verl/protocol.py:318):

字段类型装什么切分行为
batchTensorDict等长张量,第 0 维是 batch按 dim 0 切
non_tensor_batchdict[str, np.ndarray(dtype=object)]每样本一个 Python 对象np.array_split
meta_infodict整批共享的标量/配置原样复制给每一块

关键约束由 check_consistency()verl/protocol.py:454)在 __post_init__ 里强制:non_tensor_batch 的每个数组长度必须等于 batch 的 batch size。这条检查把「切分时对不齐」这类 bug 挡在了构造阶段,很值得借鉴。

1.3 核心操作

# 示意,非源码:DataProto 的典型用法
batch = DataProto.from_single_dict(batch_dict) # 自动分流张量/非张量
batch.meta_info["temperature"] = 1.0
gen = batch.repeat(repeat_times=8, interleave=True) # GRPO:每题复制 8 份
chunks = gen.chunk(chunks=world_size) # 切给各 rank
merged = DataProto.concat(chunks) # 收回来
merged = merged.union(logprob_output) # 横向合并新字段

重点看 unionverl/protocol.py:781)——RL 数据流的形状是「一个 batch 不断被加上新列」:先有 prompt,加上 response,加上 old_log_prob,加上 ref_log_prob,加上 advantages。union 就是这个「加列」动作,且会检查同名 key 的值必须相等,防止静默覆盖。

repeatinterleave 参数值得注意(verl/protocol.py:971):GRPO 需要同一道题的 n 个采样相邻,这样后面按 uid 分组时天然聚在一起。

1.4 序列化:一个容易忽视的性能点

DataProto.__getstate__verl/protocol.py:377)里做了两件事:

  1. 对 tensordict ≥ 0.5.0,先 contiguous().consolidate() —— 把散落的张量合并成一块连续内存,跨进程传输时是一次拷贝而不是 N 次。
  2. 支持用环境变量 VERL_DATAPROTO_SERIALIZATION_METHOD=numpy 切到 numpy 路径(serialize_tensordictverl/protocol.py:247)。

为什么重要: 在单控制器架构里,每一次 wg.xxx(batch) 都是一次序列化 + 网络传输。一个 512×8 条、每条几千 token 的 batch 有好几个 GB。这就直接引出了第二代。


2. 为什么要换:数据不该跟着控制流跑

把 V0 的一步画出来,问题一目了然:

【V0:数据跟着调用走】

driver worker
│ batch (GB 级) ────────────► compute_log_prob
│ ◄──────────── log_probs │
│ union 进 batch
│ batch (更大了) ───────────► compute_ref_log_prob
│ ◄──────────── ref_log_prob
│ union
│ batch (还在长) ───────────► update_actor

driver 内存里始终握着整个 batch,每步来回搬一次

三个后果:

  • 带宽浪费:同一份 prompt/response 被反复序列化传输。
  • driver 成瓶颈:所有数据过一遍 driver 进程。
  • 没法做异步:生成必须整批做完才能返回 driver,无法「边生成边训练」。

3. 第二代:TransferQueue + KVBatchMeta

3.1 思路

换成传引用。 轨迹一生成出来就写进一个全局 KV 存储,之后 driver 手上只有一串 key;需要真数据的地方(worker 内部)自己按 key 去取。

【V1:只传 key】

┌───────────── TransferQueue (全局 KV 存储) ─────────────┐
│ key: {uid}_{session_id}_{index} │
│ value: prompts / responses / response_mask / ... │
│ tag: status / seq_len / global_steps / ... │
└───▲───────────────▲──────────────────▲────────────────┘
│ 写 │ 读+写 │ 读
AgentLoopWorker 训练 worker 训练 worker

│ 只有 KVBatchMeta(keys, tags)
driver ─────────────┘
driver 从头到尾没碰过真数据

TransferQueue 是外部依赖requirements.txtTransferQueue==0.1.8),不在本仓库内;verl 只写适配层。

3.2 KVBatchMeta 长什么样

driver 手里的 batch 只有三件东西:

字段内容
partition_id"train""val"
keyskey 列表,格式 {uid}_{session_id}_{index}
tags每个 key 的元数据字典

key 的三段式含义(verl/trainer/ppo/v1/agent_loop_tq.py:177-181 的注释):

  • uid —— dataset 里一道题的唯一 id。
  • session_id —— GRPO 的第几个采样(0..n-1)。
  • index —— 一个 agent loop 可能产出多条输出,这是第几条。

tag 里放的是能在不读数据的情况下做决策的信息statusprompt_lenresponse_lenseq_lenglobal_stepsmin/max_global_stepsverl/trainer/ppo/v1/agent_loop_tq.py:205-220)。

妙在哪: driver 做负载均衡只需要 seq_len,做 staleness 过滤只需要 global_steps。这两件事都不用读真数据——tag 就是为「driver 能做的决策」量身定制的投影

3.3 tqbridge:让老方法自动适配新数据面

问题来了:worker 上那些 @register 的方法签名是 def compute_log_prob(self, data: TensorDict),而 driver 现在传的是 KVBatchMeta。总不能把所有方法重写一遍。

答案是 tqbridgeverl/utils/transferqueue_utils.py:347),它被塞在 register 装饰器的最里层

def decorator(func):
func = tqbridge(dispatch_mode=dispatch_mode)(func)
...

verl/single_controller/base/decorator.py:425

所以每个注册方法都自动获得了这层转换:

worker 进程内:

收到 KVBatchMeta


① _find_meta(*args, **kwargs) 找出参数里的 meta
│ 找不到 → 原样调用(TQ 未启用时的直通路径)

② _async_meta_to_realdata(meta)
tq_client.async_get_data(meta) → 真正的 TensorDict
再把 meta.extra_info 里的标量塞成非张量字段


③ 调用原函数 func(tensordict)


④ _async_update_meta_with_output(output, meta)
把输出张量写回 TQ,返回更新后的 meta

还有一个省带宽的优化:_compute_need_collectverl/utils/transferqueue_utils.py:210)会去问 worker「按这个 mesh,我这个 rank 是不是负责收集的那个」。不是的话直接返回空 meta,避免 TP 组里 8 个 rank 各写一份一模一样的结果

3.4 driver 侧长什么样

于是 V1 trainer 里的典型片段变成这样(verl/trainer/ppo/v1/trainer_base.py:1540_compute_ref_log_prob):

output = self.ref_policy_wg.compute_ref_log_prob(batch) # batch 是 KVBatchMeta
data = tq.kv_batch_get(keys=batch.keys, partition_id=batch.partition_id,
select_fields=["log_probs", "response_mask"])
data["ref_log_prob"] = response_from_nested(data.pop("log_probs"), data["response_mask"])
tq.kv_batch_put(keys=batch.keys, partition_id=batch.partition_id,
fields=data.select("ref_log_prob"))

注意 select_fields —— driver 只把它真正需要的那几列拉下来,而不是整个 batch。这是 V0 做不到的。


4. ReplayBuffer:从 KV 存储里凑出一个 batch

有了全局存储,「取一个 batch」就不再是「等生成函数返回」,而是「轮询存储直到攒够」。这件事由 ReplayBufferverl/trainer/ppo/v1/replay_buffer.py:63)做。

4.1 GRPO 组的状态机

GRPO 要求同一道题的 n 个采样一起参与优势计算(组内减均值)。所以不能按条采样,得按组。verl 为此在 TQ 里额外存了「题级」的状态标记,key 就是裸 uid

pending ──► running ──┬──► finished (n 个 session 全部成功)
└──► failure (至少一个失败)

只有 finished / failure 的题,它的轨迹才可以被采样

状态转移点很清晰:_add_batch_to_generate 写 pending(verl/trainer/ppo/v1/trainer_base.py:1117),AgentLoopWorkerTQ._run_prompt 开头写 running、asyncio.gather 成功后写 finished、异常写 failure(verl/trainer/ppo/v1/agent_loop_tq.py:111-148)。

4.2 采样逻辑

ReplayBuffer.sampleverl/trainer/ppo/v1/replay_buffer.py:185):

① 从 TQ 拉一遍全量元数据(只有 tag,不含数据)
② while 攒够的题数 < batch_size: sleep 2s 再拉一次
③ 按 global_steps 从小到大排序 → 优先取最老的题(减少 staleness)
④ 取前 batch_size 道题,把它们名下所有轨迹 key 收集起来
⑤ 丢弃过期样本(drop 策略)

第 ③ 步那行注释很实在:Prioritize sampling the oldest prompts (smallest global_steps first) to reduce staleness

4.3 staleness 控制:drop 还是 wait

异步 RL 的核心风险是:某条轨迹是用 10 步之前的旧权重生成的,拿来更新现在的模型会不稳。verl 给了两条策略(trainer.v1.sampler.max_off_policy_strategy):

策略行为代价
drop超过阈值的轨迹直接丢掉并从 TQ 清除浪费算力,但不阻塞
wait阻塞采样,等所有濒临超期的轨迹跑完不浪费样本,但会卡住训练

判据都是同一个式子(verl/trainer/ppo/v1/replay_buffer.py:147:162):

staleness = (当前 global_steps - 轨迹诞生的 global_steps + 1) / parameter_sync_step

除以 parameter_sync_step 是因为异步模式下不是每步都同步权重——衡量陈旧程度的单位应该是「模型版本数」而不是「训练步数」。默认阈值 8(verl/trainer/config/ppo_trainer.yaml:251)。

丢弃时还会记一组指标:training/off_policy/dropped_samplesdropped_samples_staleness/{mean,max,min}——这类可观测性对调异步训练至关重要。

源码里留了个诚实的 TODO:「是否应该在某个 session 超期时丢掉整个 GRPO 组?」(verl/trainer/ppo/v1/replay_buffer.py:172)。目前是按条丢,可能让某些组的采样数变少。


5. 关键细节:nested tensor 存变长序列

V1 里反复出现 response_from_nested / response_to_nestedverl/workers/utils/padding.py)和 to_padded_tensor()。原因是:

  • 存储时用 nested tensortorch.jagged layout)——每条序列存自己的真实长度,不浪费空间。
  • 计算时转 padded tensor——算 loss 的算子需要规整的 (bs, seqlen) 矩阵。

典型转换见 _compute_reward_colocateverl/trainer/ppo/v1/trainer_base.py:1382-1416):先 offsets().diff() 拿到每条的长度,to_padded_tensor 补齐算 RM 分数,再用 torch.nested.as_nested_tensor 按长度切回去写入 TQ。

这个来回不是冗余:它让「存储层按真实长度计费」和「计算层要求规整形状」两个矛盾需求各得其所。


6. 两代对比

维度DataProto(V0)TransferQueue(V1)
driver 手里拿的完整张量 batchKVBatchMeta(key + tag)
数据传输每次 RPC 序列化整批worker 自取,driver 只按需拉列
生成方式generate_sequences 阻塞返回整批fire-and-forget,写 TQ 后返回
能否异步是(partial rollout / off-policy)
batch 边界严格「一步一批」由 ReplayBuffer 动态凑
代码入口verl/protocol.pyverl/utils/transferqueue_utils.py + 外部包

DataProto 并没有被删掉——V1 里 driver 做优势估计等本地计算时仍会临时构造 DataProto(如 verl/trainer/ppo/v1/trainer_base.py:1594),而且 BatchDataverl/protocol.py:1231)这个适配层让 dispatch 函数同时支持 DataProto 和别的可切分类型。可以理解为:DataProto 从「传输格式」退化成了「本地计算格式」


7. 代码地图

主题文件路径符号名
数据协议本体verl/protocol.pyDataProtoDataProtoItemDataProtoFuture
一致性校验verl/protocol.pyDataProto.check_consistency
切分/拼接/加列verl/protocol.pychunkconcatunionrepeatselect_idxs
序列化优化verl/protocol.py__getstate__serialize_tensordictdeserialize_tensordict
可切分类型适配verl/protocol.pyBatchData
TQ 桥接层verl/utils/transferqueue_utils.pytqbridge_async_meta_to_realdata_async_update_meta_with_output
只让该收的 rank 收verl/utils/transferqueue_utils.py_compute_need_collect
meta 类型互转verl/utils/transferqueue_utils.pykv_batch_meta2batch_metabatch_meta2kv_batch_meta
回放缓冲verl/trainer/ppo/v1/replay_buffer.pyReplayBuffer.sample_sync_metadata_from_transfer_queue_drop_max_off_policy_samples
轨迹写入 TQverl/trainer/ppo/v1/agent_loop_tq.pyAgentLoopWorkerTQ._agent_loop_postprocess_run_prompt
变长↔定长verl/workers/utils/padding.pyresponse_from_nestedresponse_to_nestedleft_right_2_no_padding