跳转至

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 与对齐

def pad_dataloader_output(td, max_length, pad_value=0):
    """对 TensorDict 中的序列进行 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 各组件间流动的关键。