protocol.py — 数据传输协议¶
模块路径: verl/protocol.py
文件概述¶
protocol.py 是 verl 框架中最核心的文件之一,定义了模块间数据传输的标准协议。在强化学习训练中,数据需要在多个组件之间传递:
每个组件产生和消费的数据格式不同(有 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 支持丰富的索引方式:
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 类。要记住的关键点:
- 三层数据结构:
batch(tensor)+non_tensor_batch(numpy)+meta_info(元信息) - batch 操作:
chunk分片、concat合并、union合并字段、make_iterator创建迭代器 - 分布式支持:自定义序列化(
__getstate__/__setstate__)、DataProtoFuture异步执行、all_gather数据收集 - 丰富的索引:支持 int、slice、list、ndarray、Tensor 索引
DataProto 的设计哲学是:"让数据在模块间传输时,格式统一、操作简单、分布式友好"。理解了 DataProto,就理解了 verl 中数据是如何流动的。