标签为logits时如何设置WeightedRandomSampler权重解决目标检测类别不平衡问题
问题诊断
现有代码没有生效的核心问题有两个:
- 手动硬编码的类权重容易和实际数据集的类别统计结果不匹配,你给高频类设置的极低位权会导致包含高频类的样本被采样概率过低,反而会打乱正常的数据分布逻辑
- 没有兼容logits格式标签的处理步骤,直接用连续值的logits计算权重会导致结果完全偏离预期
多标签数据集采样权重计算方法
针对标签为logits的多标签数据集,统一按以下逻辑计算采样权重:
- 先对logits做二值转换:用你训练时的分类阈值(通常取0,或者sigmoid后取0.5)把连续logits转成0/1的二值标签,标识样本包含的正类
- 遍历全训练集,统计每个类别的总出现次数,得到长度等于类别数的
class_counts数组 - 每个类的基础权重设置为
1 / class_counts[c],类别的出现频次越低,对应的基础权重越高 - 单样本的采样权重取该样本所有包含的正类的基础权重之和,等价于二值标签和类权重数组的点乘结果
20类图像数据集过采样的正确实现
以下是适配VOC/20类COCO数据集的修正后代码:
import numpy as np from torch.utils.data import WeightedRandomSampler, DataLoader num_classes = 20 class_counts = np.zeros(num_classes) # 第一步:统计全训练集的类别出现频次 for _, label in ds_train: # 标签为logits时先转二值,示例阈值为0,可按需调整 binary_label = (label > 0).astype(float) class_counts += binary_label # 第二步:计算类基础权重,加极小值避免除0 class_weights = 1.0 / (class_counts + 1e-8) # 第三步:计算每个样本的采样权重 sample_weights = np.zeros(len(ds_train)) for idx, (_, label) in enumerate(ds_train): binary_label = (label > 0).astype(float) sample_weights[idx] = np.dot(class_weights, binary_label) # 第四步:初始化采样器 sampler = WeightedRandomSampler( weights=sample_weights, # 可根据需求调整每轮采样的总样本数,默认和原数据集长度一致即可 num_samples=len(sample_weights), replacement=True ) # 传入DataLoader即可,注意设置sampler后不需要再开shuffle train_loader = DataLoader(ds_train, batch_size=32, sampler=sampler)
额外优化建议
- 可以搭配类别损失权重一起使用:在损失函数中给每个类的损失乘以对应的
class_weights,双重缓解类别不平衡带来的过拟合问题 - 不建议给高频类手动设置极低的权重,会导致模型丢失高频类的泛化能力,按出现频次倒数设置权重是兼容性最高的方案
内容的提问来源于stack exchange,提问作者Priya Ravichander
相关产品推荐
相关产品推荐

