跳转至

sampler.py — 定义了两个抽象采样器接口

文件路径: verl/experimental/dataset/sampler.py

文件概述

定义了两个抽象采样器接口,为 verl 的训练数据采样提供可扩展的框架。特别是 AbstractCurriculumSampler 支持课程学习——根据训练进展动态调整数据采样策略。

关键代码讲解

1. AbstractSampler - 自定义采样器基类

class AbstractSampler(Sampler[int]):
    """自定义采样器的抽象接口"""

    @abstractmethod
    def __init__(
        self,
        data_source: Sized,       # 数据源(如 Dataset)
        data_config: DictConfig,  # 数据配置
    ):
        pass

继承自 PyTorch 的 Sampler[int],要求子类实现 __init__(以及 __iter__ 和 __len__ 等 Sampler 标准方法)。

2. AbstractCurriculumSampler - 课程学习采样器

class AbstractCurriculumSampler(AbstractSampler):
    """课程学习采样器的实验性接口"""

    @abstractmethod
    def update(self, batch: DataProto) -> None:
        pass

课程学习的核心思想:在训练过程中,根据模型的当前能力动态调整训练数据的难度。例如: - 初始阶段使用简单样本 - 模型变强后逐渐引入困难样本

update() 方法在每个训练 batch 结束后被调用,接收当前 batch 的数据(包含模型表现指标),采样器可以据此调整下一轮的采样策略。

使用场景示例

class DifficultySampler(AbstractCurriculumSampler):
    def __init__(self, data_source, data_config):
        self.difficulty_weights = [1.0] * len(data_source)  # 初始均匀权重

    def update(self, batch):
        # 根据模型在当前 batch 上的表现调整权重
        for idx, reward in zip(batch.indices, batch.rewards):
            if reward < 0.5:  # 表现差的样本增加权重
                self.difficulty_weights[idx] *= 1.1

    def __iter__(self):
        return iter(torch.multinomial(torch.tensor(self.difficulty_weights), ...))

核心类/函数列表

名称 类型 说明
AbstractSampler 抽象基类 自定义采样器接口
AbstractCurriculumSampler 抽象基类 课程学习采样器接口(含 update 方法)

与其他模块的关系

  • 继承自 PyTorch 的 torch.utils.data.Sampler
  • 使用 verl.DataProto 作为 update() 的输入参数
  • 可在 verl.trainer 的训练循环中集成使用

小结

sampler.py 定义了实验性的采样器接口。AbstractCurriculumSampler 特别有价值——它允许训练过程根据模型的实时表现动态调整数据采样分布,实现课程学习策略。目前是抽象接口,需要用户自行实现具体的采样逻辑。