You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

多标签图像分类中如何在PyTorch实现类别/样本权重采样?

多标签图像分类任务的类别不平衡处理问题

任务背景

  • 处理严重不平衡的多标签图像分类任务,每张图像对应多个标签,标签采用one-hot编码格式(示例:[1,0,0,0,0,1,0])
  • 训练数据集格式如下:
Image Index Finding Labels
0   00005504_002.png    Pleural_Thickening
1   00003527_002.png    Atelectasis|Pneumonia
2   00018285_000.png    Effusion|Mass
3   00016971_007.png    Emphysema|Mass
4   00014022_071.png    Atelectasis|Consolidation|Pleural_Thickening
  • 曾考虑对过代表现的类别做欠采样,但担心降低模型效果,希望通过类别权重/样本权重的方式解决问题

核心疑问

  1. 在PyTorch中实现带加权采样的自定义损失函数(已知类别权重对one-hot编码标签效果不佳,需采用样本权重或自定义损失)
  2. 多标签场景下如何使用sklearn的compute_sample_weight函数

当前使用的Keras实现代码

resnet50 = ResNet101(input_shape=(256, 256, 3), weights='imagenet', include_top=False)

for layer in resnet50.layers[:-3]:
    layer.trainable = False

x = Flatten()(resnet50.output)
x = Dense(512, activation='relu')(x)
prediction = Dense(13, activation='sigmoid')(x)  

model = Model(inputs=resnet50.input, outputs=prediction)
learning_rate = 0.001
adam_optimizer = Adam(learning_rate=learning_rate)

model.compile(optimizer=adam_optimizer, loss='binary_crossentropy', metrics=['accuracy', AUC(multi_label=True)])
early_stopping = EarlyStopping(monitor='val_auc', patience=5, restore_best_weights=True)
history = model.fit(train_dataset, epochs=100, validation_data=val_dataset, callbacks=[early_stopping])

解决方案

一、PyTorch中实现带加权的自定义损失函数

多标签场景下,常用加权二元交叉熵损失,先针对每个类别计算权重,再对每个样本的所有标签损失加权求和。

  1. 计算类别权重
import torch
import torch.nn as nn

def calculate_class_weights(labels):
    # labels: 形状为(N, num_classes)的one-hot张量
    pos_counts = labels.sum(dim=0)
    total_samples = labels.size(0)
    # 采用total_samples/(num_classes*正样本数)的权重计算方式,用clamp避免除以0
    weights = total_samples / (labels.size(1) * pos_counts.clamp(min=1))
    return weights
  1. 自定义加权二元交叉熵损失
class WeightedMultiLabelBCELoss(nn.Module):
    def __init__(self, class_weights=None):
        super().__init__()
        self.bce_loss = nn.BCEWithLogitsLoss(reduction='none')
        self.class_weights = class_weights

    def forward(self, inputs, targets):
        # inputs: 模型输出,形状(N, num_classes)(未经过sigmoid)
        # targets: one-hot标签,形状(N, num_classes)
        loss = self.bce_loss(inputs, targets)
        if self.class_weights is not None:
            # 将类别权重广播至与loss同形状
            weights = self.class_weights.expand_as(loss)
            # 对每个标签损失加权,再求和每个样本的总损失,最后取平均
            loss = (loss * weights).sum(dim=1).mean()
        else:
            loss = loss.mean()
        return loss
  1. 使用方式
# 假设train_labels是形状为(N,13)的one-hot张量
class_weights = calculate_class_weights(train_labels)
criterion = WeightedMultiLabelBCELoss(class_weights=class_weights)

# 训练循环示例
for images, labels in train_loader:
    outputs = model(images)
    loss = criterion(outputs, labels.float())
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

二、多标签场景下使用sklearn的compute_sample_weight

compute_sample_weight默认适配单标签任务,多标签场景需自定义样本权重逻辑:每个样本的权重取其所有正标签对应类别权重的均值(或求和,可按需调整)。

  1. 计算类别权重
from sklearn.utils.class_weight import compute_class_weight
import numpy as np

# 假设train_labels是形状为(N, num_classes)的numpy数组(one-hot格式)
# 将one-hot转换为每个样本的正标签索引列表
y_multi = [np.where(row == 1)[0] for row in train_labels]

# 统计所有类别并计算类别权重
all_classes = np.arange(train_labels.shape[1])
class_weights = compute_class_weight(class_weight='balanced', classes=all_classes, y=np.concatenate(y_multi))
  1. 计算每个样本的权重
def compute_multi_label_sample_weights(labels, class_weights):
    # labels: 形状为(N, num_classes)的one-hot numpy数组
    # class_weights: 形状为(num_classes,)的类别权重数组
    sample_weights = []
    for row in labels:
        pos_indices = np.where(row == 1)[0]
        if len(pos_indices) == 0:
            # 无正标签的样本权重设为1
            sample_weights.append(1.0)
        else:
            # 取所有正标签类别权重的均值作为样本权重
            sample_weights.append(np.mean(class_weights[pos_indices]))
    return np.array(sample_weights)

sample_weights = compute_multi_label_sample_weights(train_labels, class_weights)
  1. 在PyTorch中应用样本权重
  • 方式一:结合WeightedRandomSampler实现加权采样
from torch.utils.data import WeightedRandomSampler

# 用样本权重作为采样权重,实现对少样本类别的加权采样
sampler = WeightedRandomSampler(weights=sample_weights, num_samples=len(sample_weights), replacement=True)
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, sampler=sampler)
  • 方式二:将样本权重传入损失函数,对样本损失加权
class SampleWeightedBCELoss(nn.Module):
    def __init__(self):
        super().__init__()
        self.bce_loss = nn.BCEWithLogitsLoss(reduction='none')

    def forward(self, inputs, targets, sample_weights):
        # inputs: (N, num_classes)
        # targets: (N, num_classes)
        # sample_weights: (N,),每个样本的权重
        loss = self.bce_loss(inputs, targets)
        # 先求和每个样本的所有标签损失,再乘以样本权重,最后取平均
        loss = (loss.sum(dim=1) * sample_weights).mean()
        return loss

# 训练循环示例(需DataLoader返回样本权重)
for images, labels, weights in train_loader:
    outputs = model(images)
    loss = criterion(outputs, labels.float(), weights.float())
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

内容的提问来源于stack exchange,提问作者Amin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 17:15:22