默认dataloader使用randomsampler,仅随机打乱索引而不考虑类别分布,导致多数类样本在batch中占比过高;weightedrandomsampler通过为样本分配权重(如类别倒数归一化后值)、设num_samples为数据集长度且replacement=true,使少数类被采样概率提升,从而平衡每epoch中各类曝光率。

为什么直接用DataLoader会采样不均衡?
默认情况下,DataLoader 使用 RandomSampler,它对每个 epoch 随机打乱整个数据集索引——但不会考虑类别分布。如果训练集里猫占 80%、狗占 20%,那每次 batch 里大概率还是猫多狗少,模型根本学不好少数类。
用WeightedRandomSampler强制按类别权重采样
这是最常用也最直接的解法:给每个样本分配一个权重,让 DataLoader 在采样时按权重概率选择。关键不是“让每类数量相等”,而是“让每类在每个 epoch 中被采样的总概率接近相等”。
- 先统计每类样本数,比如
class_counts = [1200, 300](猫、狗) - 计算每个类别的倒数权重:
class_weights = [1/1200, 1/300] - 再为每个样本分配对应类别的权重:
samples_weight = [class_weights[label] for label in dataset.targets] - 传入
WeightedRandomSampler(weights=samples_weight, num_samples=len(dataset), replacement=True)
注意 replacement=True 必须设为 True,否则当少数类样本数远小于 batch_size 时,根本凑不够一个 batch;num_samples 通常设为原数据集长度,保证每个 epoch 的 batch 数不变。
Sampler 和 collate_fn 冲突吗?
不会冲突,但容易误用。如果你自定义了 collate_fn,它只负责把一批样本拼成 tensor,不影响采样逻辑。真正要检查的是:别在 collate_fn 里偷偷做重采样或过滤——那会和 WeightedRandomSampler 的行为叠加,导致实际分布失控。
- 常见错误:在
collate_fn里用random.sample截断 batch,破坏了 sampler 的权重意图 - 正确做法:所有采样逻辑只交给
Sampler,collate_fn只做 padding / stacking / type 转换 - 验证方式:打印几个 batch 的
labels统计值,看是否接近 1:1(或你设定的目标比例)
验证 sampler 是否真起作用?
别只信文档,动手验证最可靠。在训练循环前加一段测试代码:
loader = DataLoader(dataset, batch_size=32, sampler=weighted_sampler)
batch = next(iter(loader))
print("Label distribution:", torch.bincount(batch[1], minlength=len(class_counts)))
跑几次,观察输出是否稳定在目标比例附近。如果仍严重偏斜,大概率是 samples_weight 构造错了——比如用了类别 ID 当索引却没对齐 class_weights 列表顺序,或者 dataset.targets 返回的是字符串而非整数。
权重采样本身不改变数据集长度,也不做数据增强,它只是调整“看到谁”的概率。真正难的是怎么让模型在采样均衡后依然不过拟合少数类——那得靠 loss 加权、Focal Loss 或者专门的评估指标,不是 sampler 能解决的。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











