1. RandomSampler基础:从打乱数据到有放回采样

当你第一次接触PyTorch的数据加载时,RandomSampler可能是最容易被低估的组件之一。这个看似简单的采样器,实际上藏着不少实用技巧。我刚开始用PyTorch时,就经常把它和DataLoader的shuffle参数搞混——后来才发现它们解决的问题完全不同。

RandomSampler的核心功能是对数据源进行随机采样,它有两种工作模式:

  • 无放回采样(replacement=False):这是默认模式,相当于把数据集洗牌后按顺序输出,每个样本只出现一次
  • 有放回采样(replacement=True):就像从袋子里摸球,每次摸完又把球放回去,可能重复抽到相同样本

来看个直观的例子。假设我们有个包含5个样本的数据集:

from torch.utils.data import RandomSampler

# 无放回采样
sampler_false = RandomSampler(range(5), replacement=False)
print("无放回采样:", list(sampler_false))

# 有放回采样
sampler_true = RandomSampler(range(5), replacement=True)
print("有放回采样:", list(sampler_true))

运行几次你会发现,无放回采样的输出总是5个不重复的数字(顺序随机),而有放回采样则可能出现[3,3,1,4,3]这样的结果。这个特性在特定场景下非常有用,比如我们接下来要讨论的数据不平衡问题。

2. 解决类别不平衡:有放回采样的实战应用

去年做一个医学影像分类项目时,我遇到了典型的类别不平衡问题——正常样本是异常样本的20倍。直接训练的结果是模型只会无脑预测"正常",因为这样准确率就能达到95%。这时候RandomSampler的replacement参数就成了救命稻草。

2.1 为什么有放回采样能解决不平衡问题?

想象你有一个装球的袋子:

  • 90个白球(多数类)
  • 10个黑球(少数类)

常规的无放回采样,每轮epoch最多只能见到10个黑球。但如果设置replacement=True,我们可以:

  1. 指定num_samples为100(或其他大于10的值)
  2. 确保每个batch都能包含足够多的黑球样本

具体实现可以这样:

from torch.utils.data import WeightedRandomSampler

# 假设labels是类别标签(0表示多数类,1表示少数类)
class_sample_counts = [900, 100]  # 两类样本数
weights = 1. / torch.tensor(class_sample_counts, dtype=torch.float)
samples_weights = weights[labels]

# 创建采样器(本质是带权重的有放回采样)
sampler = WeightedRandomSampler(
    weights=samples_weights,
    num_samples=len(samples_weights),  # 通常设为多数类的2-3倍
    replacement=True
)

2.2 实际效果对比

在我的项目中,使用有放回采样后模型在验证集上的表现:

指标 无放回采样 有放回采样
准确率 95.2% 93.8%
召回率(少数类) 6.5% 78.3%
F1-score 0.12 0.85

虽然整体准确率略有下降,但对关键少数类的识别能力大幅提升。这正体现了有放回采样的价值——它让模型在训练时"看到"更多稀有样本。

3. 小数据集增强:num_samples的妙用

当数据量不足时,有放回采样还能变相实现数据增强。上个月帮朋友处理一个只有200张图片的分类任务时,我们是这样操作的:

# 假设dataset只有200个样本
sampler = RandomSampler(
    dataset,
    replacement=True,
    num_samples=800  # 相当于每个epoch"看到"800个样本
)

dataloader = DataLoader(
    dataset,
    batch_size=32,
    sampler=sampler
)

这种做法的优势在于:

  1. 不占用额外存储:不需要真的复制图片文件
  2. 随机性更强:每次epoch的样本组合都不同
  3. 可控性强:通过num_samples精确控制"虚拟数据集"大小

不过要注意,设置过大的num_samples可能导致模型对某些样本过拟合。我的经验法则是:

  • 初始值设为原数据量的3-5倍
  • 监控训练loss的波动情况
  • 如果波动过大,适当降低num_samples

4. 进阶技巧:组合使用多种采样策略

在实际项目中,我经常把RandomSampler和其他采样器组合使用。比如这个电商评论情感分析的项目:

from torch.utils.data import ConcatDataset, RandomSampler

# 假设有两个数据源
positive_data = ...  # 正面评论
negative_data = ...  # 负面评论

# 先各自采样,再合并
pos_sampler = RandomSampler(
    positive_data,
    replacement=True,
    num_samples=5000
)
neg_sampler = RandomSampler(
    negative_data,
    replacement=False  # 负面评论数据充足
)

# 合并数据集
combined_dataset = ConcatDataset([positive_data, negative_data])

# 创建DataLoader时使用BatchSampler
batch_sampler = BatchSampler(
    sampler=AlternatingSampler([pos_sampler, neg_sampler]),
    batch_size=64,
    drop_last=False
)

这种组合策略的好处是能对不同类别采用不同的采样策略。在上面的例子中,我们对稀少的正面评论使用有放回采样,对充足的负面评论使用常规采样。

5. 避坑指南:实际使用中的经验分享

在长期使用RandomSampler的过程中,我踩过几个典型的坑:

坑1:内存泄漏 当replacement=True且num_samples设置过大时,可能导致内存暴涨。特别是在处理图像数据时,建议:

  • 先在小型测试集上验证采样逻辑
  • 逐步增加num_samples
  • 使用torch.utils.data.Subset先处理数据子集

坑2:随机性失控 设置generator参数可以保证实验可复现:

generator = torch.Generator().manual_seed(42)
sampler = RandomSampler(
    dataset,
    replacement=True,
    num_samples=1000,
    generator=generator
)

坑3:验证集污染 记住:任何形式的采样(包括有放回采样)只应用于训练集!验证集和测试集必须使用顺序采样:

# 错误的做法
val_sampler = RandomSampler(val_dataset)  # 会导致评估结果不可靠

# 正确的做法
val_sampler = SequentialSampler(val_dataset)

最后分享一个实用技巧:当你不确定采样器是否按预期工作时,可以用这个检查函数:

def check_sampler(sampler, dataset):
    sampled_indices = list(sampler)
    print(f"总采样数: {len(sampled_indices)}")
    print(f"唯一样本数: {len(set(sampled_indices))}")
    print(f"重复样本占比: {(len(sampled_indices)-len(set(sampled_indices)))/len(sampled_indices):.1%}")
    
    # 检查类别分布(适用于分类任务)
    if hasattr(dataset, 'targets'):
        targets = [dataset.targets[i] for i in sampled_indices]
        print("类别分布:", np.bincount(targets))
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐