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 特别有价值——它允许训练过程根据模型的实时表现动态调整数据采样分布,实现课程学习策略。目前是抽象接口,需要用户自行实现具体的采样逻辑。