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 策略和变长序列的批处理方式。