tensordict_utils.py — TensorDict 操作工具¶
文件路径: verl/utils/tensordict_utils.py
文件概述¶
tensordict_utils.py(约 850 行)提供了对 PyTorch TensorDict 的全面操作工具。TensorDict 是 verl 中数据在各组件间传递的核心数据结构,类似于一个支持批量操作的张量字典。
背景知识¶
TensorDict 是 PyTorch 的 tensordict 库提供的数据结构,可以看作一个能对所有值同时做切片、索引、设备转移等操作的字典。verl 用它来封装一个 batch 的所有数据(input_ids、attention_mask、rewards 等)。
核心函数详解¶
1. 创建与转换¶
def make_batch(data_dict, batch_size):
"""从普通字典创建 TensorDict"""
return TensorDict(data_dict, batch_size=batch_size)
def tensordict_to_dict(td):
"""将 TensorDict 转为普通字典"""
2. 切分与合并¶
def split_batch(td, micro_batch_size):
"""将一个大 batch 的 TensorDict 切分为多个 micro-batch"""
return td.split(micro_batch_size, dim=0)
def concat_batches(td_list):
"""将多个 TensorDict 合并为一个"""
return torch.cat(td_list, dim=0)
在 RLHF 训练中,一个 batch 可能很大,需要切分成 micro-batch 逐个处理。
3. 分布式操作¶
def gather_tensor_dict(td, dst=0, group=None):
"""将所有 rank 的 TensorDict gather 到 dst rank"""
def scatter_tensor_dict(td, src=0, group=None):
"""从 src rank scatter TensorDict 到所有 rank"""
这些函数支持在分布式环境中传输整个 TensorDict。
4. Padding 与对齐¶
不同样本的序列长度不同,需要 padding 到最长以形成规整的张量。
核心函数列表¶
| 函数 | 说明 |
|---|---|
make_batch() |
创建 TensorDict |
split_batch() |
切分为 micro-batch |
concat_batches() |
合并 TensorDict |
gather_tensor_dict() |
分布式 gather |
scatter_tensor_dict() |
分布式 scatter |
pad_dataloader_output() |
序列 padding |
与其他模块的关系¶
- 被 trainer 和 worker 广泛使用,是数据传递的核心
- 被
dataset/rl_dataset.py用来构建训练批次 - 被
distributed.py的通信函数间接使用
小结¶
tensordict_utils.py 是 verl 数据管道的核心工具,提供了对 TensorDict 的创建、切分、合并、分布式通信等全套操作,是理解数据如何在 verl 各组件间流动的关键。