跳转到主内容
趣航编程网 - 趣学编程,启航技术之路!

如何在Python中处理PyTorch中的类别不平衡_通过WeightedRandomSampler采样

WeightedRandomSampler通过调整采样概率解决类别不平衡,它为小类样本赋予更高权重,使其在训练中更频繁出现;需用归一化后的类别倒数计算样本权重并转为torch.double类型,num_samples设为数据集长度且replacement必须为True。 WeightedRandomSampler 为什么能解决类别不平衡 它不修改模型结构或损失函数,而是让 DataLoader 在每次迭代时,按预设权重从训练集里“有倾向地”抽样。比如类别 A 只有 100 个样本、类别 B 有 1000 个,你给 A 的权重设高些,Sampler 就会更频繁地把 A 的样本送进 batch,变相提升小类曝光率。 关键点在于:权重不是直接用样本倒数(
1 / count
),而是要归一化后传给
WeightedRandomSampler
;否则可能触发
RuntimeError: invalid argument 3: invalid weight
。 怎么算每个样本的 weight(别直接用 1/count) 常见错误是为每个类别算一个 weight,然后广播到该类所有样本上——这本身没错,但后续必须展平成和 dataset 长度一致的一维 tensor,且 dtype 必须是
torch.double
torch.float64
float32
会报错)。 先统计每个类别的样本数:
class_counts = np.bincount(labels)
再算每个类的权重:
class_weights = 1. / class_counts
映射到每个样本:
samples_weight = np.array([class_weights[label] for label in labels])
转成 tensor 并指定 dtype:
samples_weight = torch.from_numpy(samples_weight).double()
初始化 WeightedRandomSampler 的三个关键参数
WeightedRandomSampler
构造时接受三个参数:
weights
num_samples
replacement
。最容易出问题的是后两个: 立即学习 “ Python免费学习笔记(深入) ”; Python 3.14.3 微软官方的 Python 扩展,是 VS Code 安装量最高的扩展(209M+)。集成 IntelliSense(通过 Pylance)、调试(通过 Python Debugger)、代码检查、格式化、重构和单元测试等功能。支持 Jupyter Notebook、虚拟环境管理和多 Python 版本切换。 下载
num_samples
一般设为原训练集长度(
len(dataset)
),这样每个 epoch 的 batch 数不变;设小了会漏样本,设大了会重复采样
replacement=True
必须为 True —— 否则权重无效,Sampler 退化为普通随机采样 如果误传
replacement=False
,不会报错,但 loss 下降缓慢、小类指标几乎不动,调试时极难定位 和 DataLoader 配合时的隐藏坑 WeightedRandomSampler 和
shuffle=True
不能共存,否则会报
ValueError: sampler option is mutually exclusive with shuffle
。这不是 bug,是设计使然:Sampler 本身已控制顺序,再 shuffle 会破坏权重逻辑。 正确写法是显式关闭 shuffle,并把 Sampler 传进去:
train_sampler = WeightedRandomSampler(weights=samples_weight, num_samples=len(dataset), replacement=True) train_loader = DataLoader(dataset, batch_size=32, sampler=train_sampler, num_workers=4)
另外,如果你用了
DistributedSampler
(多卡训练),就不能再套
WeightedRandomSampler
—— 它们不兼容。此时得改用 per-GPU 的加权逻辑,或在损失层用
weight
参数替代。 权重计算本身不耗时,但一旦 dataset 很大、label 分布极偏(比如 99% 背景 + 1% 目标),
samples_weight
张量会占用额外内存,建议用
torch.tensor(..., dtype=torch.float64)
显式声明,避免默认 int 类型引发隐式转换失败。

相关文章