数据截至 (上游 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_ids、attention_mask、advantages…… - 不等长/非张量:原始对话
raw_prompt、uid字符串、多模态图片对象…… - 整批共享的元信息:
temperature、global_steps、要不要算熵……
如果只用 dict[str, Tensor],第二三类没地方放;如果全用 dict,切分/拼接/搬设备就要到处写循环。
1.2 结构
DataProto 是个三字段 dataclass(verl/protocol.py:318):
| 字段 | 类型 | 装什么 | 切分行为 |
|---|---|---|---|
batch | TensorDict | 等长张量,第 0 维是 batch | 按 dim 0 切 |
non_tensor_batch | dict[str, np.ndarray(dtype=object)] | 每样本一个 Python 对象 | np.array_split |
meta_info | dict | 整批共享的标量/配置 | 原样复制给每一块 |
关键约束由 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) # 横向合并新字段
重点看 union(verl/protocol.py:781)——RL 数据流的形状是「一个 batch 不断被加上新列」:先有 prompt,加上 response,加上 old_log_prob,加上 ref_log_prob,加上 advantages。union 就是这个「加列」动作,且会检查同名 key 的值必须相等,防止静默覆盖。
repeat 的 interleave 参数值得注意(verl/protocol.py:971):GRPO 需要同一道题的 n 个采样相邻,这样后面按 uid 分组时天然聚在一起。
1.4 序列化:一个容易忽视的性能点
DataProto.__getstate__(verl/protocol.py:377)里做了两件事:
- 对 tensordict ≥ 0.5.0,先
contiguous().consolidate()—— 把散落的张量合并成一块连续内存,跨进程传输时是一次拷贝而不是 N 次。 - 支持用环境变量
VERL_DATAPROTO_SERIALIZATION_METHOD=numpy切到 numpy 路径(serialize_tensordict,verl/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.txt 里 TransferQueue==0.1.8),不在本仓库内;verl 只写适配层。
3.2 KVBatchMeta 长什么样
driver 手里的 batch 只有三件东西:
| 字段 | 内容 |
|---|---|
partition_id | "train" 或 "val" |
keys | key 列表,格式 {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 里放的是能在不读数据的情况下做决策的信息:status、prompt_len、response_len、seq_len、global_steps、min/max_global_steps(verl/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。总不能把所有方法重写一遍。
答案是 tqbridge(verl/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_collect(verl/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」就不再是「等生成函数返回」,而是「轮询存储直到攒够」。这件事由 ReplayBuffer(verl/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.sample(verl/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_samples、dropped_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_nested(verl/workers/utils/padding.py)和 to_padded_tensor()。原因是:
- 存储时用 nested tensor(
torch.jaggedlayout)——每条序列存自己的真实长度,不浪费空间。 - 计算时转 padded tensor——算 loss 的算子需要规整的
(bs, seqlen)矩阵。
典型转换见 _compute_reward_colocate(verl/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 手里拿的 | 完整张量 batch | KVBatchMeta(key + tag) |
| 数据传输 | 每次 RPC 序列化整批 | worker 自取,driver 只按需拉列 |
| 生成方式 | generate_sequences 阻塞返回整批 | fire-and-forget,写 TQ 后返回 |
| 能否异步 | 否 | 是(partial rollout / off-policy) |
| batch 边界 | 严格「一步一批」 | 由 ReplayBuffer 动态凑 |
| 代码入口 | verl/protocol.py | verl/utils/transferqueue_utils.py + 外部包 |
DataProto 并没有被删掉——V1 里 driver 做优势估计等本地计算时仍会临时构造 DataProto(如 verl/trainer/ppo/v1/trainer_base.py:1594),而且 BatchData(verl/protocol.py:1231)这个适配层让 dispatch 函数同时支持 DataProto 和别的可切分类型。可以理解为:DataProto 从「传输格式」退化成了「本地计算格式」。
7. 代码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 数据协议本体 | verl/protocol.py | DataProto、DataProtoItem、DataProtoFuture |
| 一致性校验 | verl/protocol.py | DataProto.check_consistency |
| 切分/拼接/加列 | verl/protocol.py | chunk、concat、union、repeat、select_idxs |
| 序列化优化 | verl/protocol.py | __getstate__、serialize_tensordict、deserialize_tensordict |
| 可切分类型适配 | verl/protocol.py | BatchData |
| TQ 桥接层 | verl/utils/transferqueue_utils.py | tqbridge、_async_meta_to_realdata、_async_update_meta_with_output |
| 只让该收的 rank 收 | verl/utils/transferqueue_utils.py | _compute_need_collect |
| meta 类型互转 | verl/utils/transferqueue_utils.py | kv_batch_meta2batch_meta、batch_meta2kv_batch_meta |
| 回放缓冲 | verl/trainer/ppo/v1/replay_buffer.py | ReplayBuffer.sample、_sync_metadata_from_transfer_queue、_drop_max_off_policy_samples |
| 轨迹写入 TQ | verl/trainer/ppo/v1/agent_loop_tq.py | AgentLoopWorkerTQ._agent_loop_postprocess、_run_prompt |
| 变长↔定长 | verl/workers/utils/padding.py | response_from_nested、response_to_nested、left_right_2_no_padding |