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

多标签稀疏输出的神经网络分类:合适损失函数及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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 04:50:07