跳转至

protocol.py — 数据传输协议

模块路径: verl/protocol.py


文件概述

protocol.py 是 verl 框架中最核心的文件之一,定义了模块间数据传输的标准协议。在强化学习训练中,数据需要在多个组件之间传递:

Prompt 数据 → Actor(生成回复)→ Critic(评估价值)→ Reward Model(打分)→ 训练更新

每个组件产生和消费的数据格式不同(有 tensor、有文本、有元信息),DataProto 就是统一封装这些数据的"快递包裹"。

本文件定义的核心类: - DataProto:主要数据容器,封装 tensor 数据、非 tensor 数据和元信息 - DataProtoItem:DataProto 中单个样本的表示 - DataProtoFuture:DataProto 的异步版本,支持 Ray 分布式场景 - BatchData:统一的批数据操作包装器


关键代码讲解

1. DataProto 的三层数据结构

@dataclass
class DataProto:
    batch: TensorDict = None              # tensor 数据(如 input_ids, attention_mask)
    non_tensor_batch: dict = field(default_factory=dict)  # 非 tensor 数据(如原始文本)
    meta_info: dict = field(default_factory=dict)          # 元信息(如配置参数)

这是 DataProto 最核心的设计——将数据分为三层:

层 类型 存储内容 特点
batch TensorDict PyTorch tensor,如 token IDs、attention mask 支持 GPU 操作,有 batch 维度
non_tensor_batch dict[str, np.ndarray] 非 tensor 数据,如原始文本字符串 每个值是 numpy 数组,第 0 维是 batch 维
meta_info dict 元信息,如超参数、配置 不参与 batch 操作,在 split/concat 时共享

为什么需要三层?因为强化学习训练中的数据很复杂: - input_ids(tensor)需要在 GPU 上做计算 - 原始 prompt 文本(非 tensor)需要传给 reward model 做评估 - 训练的超参数(元信息)需要跟着数据一起传递

2. 一致性检查

def check_consistency(self):
    if self.batch is not None:
        assert len(self.batch.batch_size) == 1, "only support num_batch_dims=1"

    if self.non_tensor_batch is not None:
        for key, val in self.non_tensor_batch.items():
            assert isinstance(val, np.ndarray)

    if self.batch is not None and self.non_tensor_batch is not None and len(self.non_tensor_batch) != 0:
        batch_size = self.batch.batch_size[0]
        for key, val in self.non_tensor_batch.items():
            assert val.shape[0] == batch_size

在创建 DataProto 时自动检查: 1. batch 只支持一维 batch(即 [batch_size, ...]) 2. non_tensor_batch 中的值必须是 numpy 数组 3. 如果同时有 batch 和 non_tensor_batch,它们的第 0 维大小必须一致

这个检查确保了数据的一致性——你不会把 32 条 tensor 数据和 64 条文本数据混在一起。

3. 创建 DataProto 的工厂方法

从字典创建

@classmethod
def from_dict(cls, tensors=None, non_tensors=None, meta_info=None,
              num_batch_dims=1, auto_padding=False):
    # 检查所有 tensor 的 batch size 一致
    batch_size = None
    for key, tensor in tensors.items():
        if batch_size is None:
            batch_size = tensor.shape[:num_batch_dims]
        else:
            current_batch = tensor.shape[:num_batch_dims]
            assert batch_size == current_batch

    # 将非 numpy 类型自动转换为 numpy 数组
    for key, val in non_tensors.items():
        if not isinstance(val, np.ndarray):
            non_tensors[key] = np.array(val, dtype=object)

    tensor_dict = TensorDict(source=tensors, batch_size=batch_size) if tensors else None
    return cls(batch=tensor_dict, non_tensor_batch=non_tensors, meta_info=meta_info)

使用示例:

data = DataProto.from_dict(
    tensors={"input_ids": torch.tensor([[1,2,3], [4,5,6]])},
    non_tensors={"text": np.array(["hello", "world"], dtype=object)},
    meta_info={"task": "rlhf"}
)

从混合字典创建

@classmethod
def from_single_dict(cls, data: dict, meta_info=None, auto_padding=False):
    tensors = {}
    non_tensors = {}
    for key, val in data.items():
        if isinstance(val, torch.Tensor):
            tensors[key] = val
        elif isinstance(val, np.ndarray):
            non_tensors[key] = val
    return cls.from_dict(tensors=tensors, non_tensors=non_tensors, meta_info=meta_info)

更方便的创建方式——自动按类型分拣 tensor 和非 tensor。

4. 索引和切片

def __getitem__(self, item):
    # slice → 返回 DataProto
    if isinstance(item, slice):
        return self.slice(item.start, item.stop, item.step)
    # list/ndarray/Tensor → 返回 DataProto
    elif isinstance(item, list | np.ndarray | torch.Tensor):
        return self.select_idxs(item)
    # 单个 int → 返回 DataProtoItem
    elif isinstance(item, int | np.integer):
        tensor_data = self.batch[item] if self.batch is not None else None
        non_tensor_data = {key: val[item] for key, val in self.non_tensor_batch.items()}
        return DataProtoItem(batch=tensor_data, non_tensor_batch=non_tensor_data, meta_info=self.meta_info)

DataProto 支持丰富的索引方式:

data[0]       # 单个样本 → DataProtoItem
data[0:10]    # 切片 → DataProto
data[[0,3,5]] # 索引列表 → DataProto

5. 分块和合并 — 分布式训练的基础

chunk(分块)

def chunk(self, chunks: int) -> list["DataProto"]:
    batch_lst = self.batch.chunk(chunks=chunks, dim=0)
    # ... 同时对 non_tensor_batch 做相同的分块
    output = []
    for i in range(chunks):
        output.append(type(self)(batch=batch_lst[i],
                                 non_tensor_batch=non_tensor_batch_lst[i],
                                 meta_info=self.meta_info))
    return output

将数据沿 batch 维度均分成 chunks 份。这在分布式训练中非常重要——需要把数据分给多个 worker。

concat(合并)

@staticmethod
def concat(data: list["DataProto"]) -> "DataProto":
    batch_lst = [batch.batch for batch in data]
    new_batch = torch.cat(batch_lst, dim=0)

    non_tensor_batch = list_of_dict_to_dict_of_list([d.non_tensor_batch for d in data])
    for key, val in non_tensor_batch.items():
        non_tensor_batch[key] = np.concatenate(val, axis=0)

    return DataProto(batch=new_batch, non_tensor_batch=non_tensor_batch, meta_info=merged_meta_info)

将多个 DataProto 沿 batch 维度合并为一个。这是 chunk 的逆操作,用于收集多个 worker 的结果。

6. 合并两个 DataProto(字段维度)

def union(self, other: "DataProto") -> "DataProto":
    self.batch = union_tensor_dict(self.batch, other.batch)
    self.non_tensor_batch = union_numpy_dict(self.non_tensor_batch, other.non_tensor_batch)
    self.meta_info = union_two_dict(self.meta_info, other.meta_info)
    return self

注意 union 和 concat 的区别: - concat:沿 batch 维度拼接(增加样本数) - union:合并不同的字段(增加每个样本的信息)

例如,Actor 生成了 response_ids,Critic 计算了 values,可以用 union 把它们合并到一起。

7. 创建数据迭代器

def make_iterator(self, mini_batch_size, epochs, seed=None, dataloader_kwargs=None):
    train_dataloader = DataLoader(
        dataset=self,
        batch_size=mini_batch_size,
        collate_fn=collate_fn,
        generator=generator,
        **dataloader_kwargs
    )

    def get_data():
        for _ in range(epochs):
            for d in train_dataloader:
                d.meta_info = self.meta_info
                yield d

    return iter(get_data())

将 DataProto 变成可迭代的 mini-batch 迭代器。在 PPO 训练中,一批经验数据通常需要被多个 epoch 的 mini-batch 遍历。这个方法直接复用了 PyTorch 的 DataLoader。

8. 序列化与反序列化

def __getstate__(self):
    if os.getenv("VERL_DATAPROTO_SERIALIZATION_METHOD") == "numpy":
        batch = serialize_tensordict(self.batch)
        return (batch, self.non_tensor_batch, self.meta_info)
    else:
        import io
        buffer = io.BytesIO()
        torch.save(batch, buffer)
        return buffer.getvalue(), self.non_tensor_batch, self.meta_info

__getstate__ 和 __setstate__ 控制 Python pickle 序列化的行为。这在 Ray 分布式通信中至关重要——数据需要在不同进程间传输,而传输前必须序列化为字节流。

verl 提供了两种序列化方式: - 默认方式:用 torch.save 序列化 TensorDict - numpy 方式:将 tensor 转为 numpy 数组后序列化(通过环境变量 VERL_DATAPROTO_SERIALIZATION_METHOD=numpy 启用)

9. DataProtoFuture — 异步数据传输

@dataclass
class DataProtoFuture:
    collect_fn: Callable            # 如何收集多个 future 的结果
    futures: list[ray.ObjectRef]    # Ray 的异步引用列表
    dispatch_fn: Callable = None    # 如何分发结果

    def get(self):
        output = ray.get(self.futures)  # 等待所有 future 完成
        output = DataProto.concat(output)  # 合并结果
        if self.dispatch_fn is not None:
            output = self.dispatch_fn(output)  # 分发
        return output

DataProtoFuture 是 DataProto 的"延迟版本"。在分布式场景中,driver 进程不需要立即获取 worker 的计算结果,而是持有一个 future 引用,在需要时再 .get() 获取实际数据。这样可以实现异步执行,提高吞吐量。

10. BatchData — 统一批操作接口

class BatchData:
    def __init__(self, data):
        self._data = data

    def chunk(self, chunks: int):
        data = self._data
        if isinstance(data, TensorDict):
            raw_chunks = chunk_tensordict(data, chunks)
            return tuple(contiguous(val).consolidate() for val in raw_chunks)
        return data.chunk(chunks=chunks)

    def concat(self):
        data = self._data
        sample = data[0]
        if isinstance(sample, ray.ObjectRef):
            return DataProtoFuture.concat(data)
        if isinstance(sample, TensorDict):
            return concat_tensordict(data)
        return type(sample).concat(data)

BatchData 是一个适配器,它将 DataProto、TensorDict、DataProtoFuture 等不同类型的数据统一到相同的 chunk/concat 接口下,调用方不需要关心数据的具体类型。

11. 辅助函数

pad_dataproto_to_divisor — 填充对齐

def pad_dataproto_to_divisor(data: "DataProto", size_divisor: int):
    if len(data) % size_divisor != 0:
        pad_size = size_divisor - len(data) % size_divisor
        # 用已有数据的副本来填充
        data_padded = DataProto.concat([data] + padding_protos)
    else:
        pad_size = 0
        data_padded = data
    return data_padded, pad_size

在分布式训练中,数据需要均匀分配给各个 worker。如果数据量不能被 worker 数整除,就需要填充。这个函数用已有数据的副本来填充,确保 batch 大小可以被 size_divisor 整除。

union_tensor_dict — 合并两个 TensorDict

def union_tensor_dict(tensor_dict1: TensorDict, tensor_dict2: TensorDict) -> TensorDict:
    assert tensor_dict1.batch_size == tensor_dict2.batch_size
    for key in tensor_dict2.keys():
        if key not in tensor_dict1.keys():
            tensor_dict1[key] = tensor_dict2[key]
        else:
            assert tensor_dict1[key].equal(tensor_dict2[key])
    return tensor_dict1

将两个 batch size 相同的 TensorDict 合并。如果有重复的 key,检查它们的值是否相同。

all_gather_data_proto — 分布式收集

def all_gather_data_proto(data: DataProto, process_group):
    group_size = torch.distributed.get_world_size(group=process_group)
    data = data.to(get_device_id())
    data.batch = allgather_dict_tensors(data.batch.contiguous(), size=group_size,
                                         group=process_group, dim=0)
    all_non_tensor_batch = [None for _ in range(group_size)]
    torch.distributed.all_gather_object(all_non_tensor_batch, data.non_tensor_batch,
                                         group=process_group)
    data.non_tensor_batch = {k: np.concatenate([d[k] for d in all_non_tensor_batch])
                              for k in data.non_tensor_batch}

在分布式训练中,每个 worker 有自己的数据分片。all_gather_data_proto 将所有 worker 的数据收集到一起。tensor 部分用 allgather_dict_tensors(基于 NCCL),非 tensor 部分用 all_gather_object(基于 Python 对象序列化)。


核心类/函数列表

名称 类型 作用
DataProto dataclass 核心数据容器,封装 tensor、非 tensor 和元信息
DataProtoItem dataclass 单个样本的数据表示
DataProtoFuture dataclass DataProto 的异步版本,用于 Ray 分布式
DataProtoConfig 类 DataProto 的全局配置(如是否自动填充)
BatchData 类 统一的批数据操作包装器
pad_dataproto_to_divisor 函数 将 DataProto 填充到可整除的大小
unpad_dataproto 函数 移除填充部分
union_tensor_dict 函数 合并两个 TensorDict
collate_fn 函数 DataLoader 的数据整理函数
all_gather_data_proto 函数 分布式 all-gather 操作
fold_batch_dim / unfold_batch_dim 函数 batch 维度的折叠和展开
serialize_tensordict / deserialize_tensordict 函数 TensorDict 的序列化和反序列化

与其他模块的关系

verl/protocol.py (DataProto)
│
├── 被 verl/__init__.py 导入并暴露为公开 API
│
├── 被训练流程使用
│   ├── Actor 生成数据 → 封装为 DataProto
│   ├── Critic 评估 → 从 DataProto 读取,结果 union 回去
│   ├── Reward Model → 从 DataProto 读取 prompt 和 response
│   └── PPO Trainer → 用 make_iterator 遍历 mini-batch
│
├── 被分布式系统使用
│   ├── Ray worker 之间传输 DataProto(序列化/反序列化)
│   ├── DataProtoFuture 支持异步执行
│   └── all_gather_data_proto 支持数据收集
│
└── 依赖
    ├── tensordict (TensorDict) — tensor 数据的底层存储
    ├── torch — tensor 操作
    ├── numpy — 非 tensor 数据存储
    ├── ray — 分布式计算
    └── verl/utils/ — 工具函数

数据流示意图

下图展示了 DataProto 在 RLHF 训练中的典型数据流:

┌─────────────────────────────────────────────────────────────────────┐
│                         训练循环                                     │
│                                                                     │
│  1. 准备数据                                                         │
│     prompts (DataProto)                                             │
│       ├── batch: {input_ids, attention_mask}                        │
│       └── non_tensor_batch: {raw_text}                              │
│                    │                                                │
│                    ▼                                                │
│  2. Actor 生成     chunk(n) 分给 n 个 worker                         │
│     ┌──────┐ ┌──────┐ ┌──────┐                                     │
│     │worker│ │worker│ │worker│                                      │
│     │  0   │ │  1   │ │  2   │                                      │
│     └──┬───┘ └──┬───┘ └──┬───┘                                     │
│        └────┬───┘────────┘                                          │
│             ▼ concat() 合并结果                                      │
│     responses (DataProto)                                           │
│       ├── batch: {input_ids, response_ids, attention_mask}          │
│       └── non_tensor_batch: {raw_text, generated_text}              │
│                    │                                                │
│                    ▼                                                │
│  3. Critic/Reward 评估                                               │
│     values, rewards → union() 合并到同一个 DataProto                  │
│                    │                                                │
│                    ▼                                                │
│  4. PPO 训练                                                         │
│     make_iterator(mini_batch_size, epochs)                          │
│     逐 mini-batch 更新模型参数                                        │
│                                                                     │
└─────────────────────────────────────────────────────────────────────┘

小结

protocol.py 定义了 verl 的数据传输协议,核心是 DataProto 类。要记住的关键点:

  1. 三层数据结构:batch(tensor)+ non_tensor_batch(numpy)+ meta_info(元信息)
  2. batch 操作:chunk 分片、concat 合并、union 合并字段、make_iterator 创建迭代器
  3. 分布式支持:自定义序列化(__getstate__/__setstate__)、DataProtoFuture 异步执行、all_gather 数据收集
  4. 丰富的索引:支持 int、slice、list、ndarray、Tensor 索引

DataProto 的设计哲学是:"让数据在模块间传输时,格式统一、操作简单、分布式友好"。理解了 DataProto,就理解了 verl 中数据是如何流动的。