多标签稀疏输出的神经网络分类:合适损失函数及PyTorch/TensorFlow实现
多标签稀疏分类任务的合适损失函数及实现
一、推荐的损失函数
- 带类别权重的二元交叉熵(Weighted BCE):针对稀疏标签,给正样本(1)分配对应权重(通常按负/正样本比例计算),既提升正样本的损失占比,又通过合理权重避免模型偏向全1输出。
- Focal Loss:降低易分类样本(如大量负样本)的损失权重,让模型聚焦于难分类的正样本,天然适配稀疏场景,不会轻易出现全1输出的极端情况。
- Dice Loss:基于样本重叠度计算损失,对不平衡数据鲁棒,能有效规避全0或全1的无效输出,适合稀疏标签任务。
二、PyTorch实现示例
1. 带类别权重的BCE
import torch import torch.nn.functional as F def weighted_bce_loss(preds, targets): # 计算正负样本比例作为正样本权重 pos_weight = torch.sum(targets == 0) / torch.sum(targets == 1) # 用内置带权重的BCE(输入为logits时无需手动sigmoid) loss = F.binary_cross_entropy_with_logits(preds, targets.float(), pos_weight=pos_weight) return loss
2. Focal Loss
class FocalLoss(torch.nn.Module): def __init__(self, alpha=1, gamma=2, reduction='mean'): super(FocalLoss, self).__init__() self.alpha = alpha # 平衡正负样本的权重系数 self.gamma = gamma # 调节易分类样本的损失衰减程度 self.reduction = reduction def forward(self, preds, targets): BCE_loss = F.binary_cross_entropy_with_logits(preds, targets.float(), reduction='none') pt = torch.exp(-BCE_loss) # 预测正确的概率 focal_loss = self.alpha * (1 - pt)**self.gamma * BCE_loss if self.reduction == 'mean': return torch.mean(focal_loss) elif self.reduction == 'sum': return torch.sum(focal_loss) else: return focal_loss
3. Dice Loss
class DiceLoss(torch.nn.Module): def __init__(self, smooth=1e-6): super(DiceLoss, self).__init__() self.smooth = smooth # 避免分母为0的平滑项 def forward(self, preds, targets): preds = torch.sigmoid(preds) # 将logits转为0-1概率 intersection = torch.sum(preds * targets) union = torch.sum(preds) + torch.sum(targets) dice = (2. * intersection + self.smooth) / (union + self.smooth) return 1 - dice # 损失为1-Dice系数
三、TensorFlow实现示例
1. 带类别权重的BCE
import tensorflow as tf def weighted_bce_loss(y_true, y_pred): # 计算正负样本比例作为正样本权重 pos_weight = tf.reduce_sum(tf.cast(y_true == 0, tf.float32)) / tf.reduce_sum(tf.cast(y_true == 1, tf.float32)) loss = tf.keras.losses.binary_crossentropy(y_true, y_pred, from_logits=True, pos_weight=pos_weight) return tf.reduce_mean(loss)
2. Focal Loss
class FocalLoss(tf.keras.losses.Loss): def __init__(self, alpha=1., gamma=2., reduction='mean'): super().__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def call(self, y_true, y_pred): bce_loss = tf.keras.losses.binary_crossentropy(y_true, y_pred, from_logits=True) pt = tf.exp(-bce_loss) focal_loss = self.alpha * tf.pow(1 - pt, self.gamma) * bce_loss if self.reduction == 'mean': return tf.reduce_mean(focal_loss) elif self.reduction == 'sum': return tf.reduce_sum(focal_loss) else: return focal_loss
3. Dice Loss
class DiceLoss(tf.keras.losses.Loss): def __init__(self, smooth=1e-6): super().__init__() self.smooth = smooth def call(self, y_true, y_pred): y_pred = tf.math.sigmoid(y_pred) intersection = tf.reduce_sum(y_true * y_pred) union = tf.reduce_sum(y_true) + tf.reduce_sum(y_pred) dice = (2. * intersection + self.smooth) / (union + self.smooth) return 1 - dice
内容的提问来源于stack exchange,提问作者ShoutOutAndCalculate
相关产品推荐
相关产品推荐

