跳转至

dataset_utils.py — 数据集工具

文件路径: verl/utils/dataset/dataset_utils.py

文件概述

提供数据集相关的通用工具类,包括 padding 模式枚举和自定义 collate 函数。

核心类

DatasetPadMode — Padding 模式

class DatasetPadMode(str, Enum):
    RIGHT = "right"          # 右侧 padding
    LEFT_RIGHT = "left_right"  # 左右双侧 padding
    NO_PADDING = "no_padding"  # 不 padding(变长序列)

SFTTensorCollator — 自定义 Collate 函数

class SFTTensorCollator:
    def __init__(self, pad_mode=DatasetPadMode.LEFT_RIGHT):
        self.pad_mode = pad_mode

    def __call__(self, batch):
        if self.pad_mode == DatasetPadMode.NO_PADDING:
            return self.collate_variable_batch(batch)  # 使用 NestedTensor
        else:
            return default_collate(batch)  # 标准 collate

    def collate_variable_batch(self, batch):
        """变长序列使用 NestedTensor 打包"""
        for key in tensor_keys:
            if isinstance(batch[0][key], torch.Tensor):
                tensors = [item[key] for item in batch]
                final_batch[key] = torch.nested.as_nested_tensor(tensors, layout=torch.jagged)

NO_PADDING 模式使用 PyTorch 的 NestedTensor(jagged tensor),避免不必要的 padding 浪费。

与其他模块的关系

  • 被 multiturn_sft_dataset.py 使用 DatasetPadMode
  • SFTTensorCollator 作为 DataLoader 的 collate_fn

小结

数据集的通用工具,定义了 padding 策略和变长序列的批处理方式。