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

求PyTorch中多类别图像分割的可靠Focal Loss实现

多类别图像分割的PyTorch Focal Loss实现方案

官方现状

PyTorch目前没有像nn.BCELoss那样内置的Focal Loss实现,需要自定义实现适配多类别分割场景。

支持形状不匹配的可靠实现

以下是适配多类别图像分割、兼容目标张量(类别索引格式)与预测张量(通道对应类别)形状差异的Focal Loss实现:

import torch
import torch.nn as nn
import torch.nn.functional as F

class FocalLoss(nn.Module):
    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):
        super().__init__()
        self.alpha = alpha
        self.gamma = gamma
        self.reduction = reduction

        # 处理alpha参数:支持单值(全局类别权重)或列表(逐类别权重)
        if isinstance(alpha, list):
            self.alpha = torch.tensor(alpha)
        elif alpha is not None:
            self.alpha = torch.tensor([alpha])

    def forward(self, inputs, targets):
        # inputs形状: [B, C, H, W],模型输出的logits(未经过softmax)
        # targets形状: [B, H, W],每个像素为类别索引(0到C-1)

        # 将目标张量转换为one-hot编码,匹配inputs的设备与维度顺序
        targets_onehot = F.one_hot(targets, num_classes=inputs.size(1)).permute(0, 3, 1, 2).float()
        targets_onehot = targets_onehot.to(inputs.device)

        # 计算类别概率与目标类别的概率值
        probs = F.softmax(inputs, dim=1)
        pt = torch.sum(targets_onehot * probs, dim=1)
        # 计算基础交叉熵损失
        ce_loss = F.cross_entropy(inputs, targets, reduction='none')

        # 计算focal权重,降低易分类样本的损失占比
        focal_weight = (1 - pt) ** self.gamma

        # 应用alpha类别权重(若配置)
        if self.alpha is not None:
            self.alpha = self.alpha.to(inputs.device)
            alpha = self.alpha.gather(0, targets.view(-1)).view(targets.shape)
            focal_weight = focal_weight * alpha

        # 计算最终损失并处理归约方式
        loss = focal_weight * ce_loss
        if self.reduction == 'mean':
            return loss.mean()
        elif self.reduction == 'sum':
            return loss.sum()
        else:
            return loss

核心特性说明

  • 自动兼容形状差异:支持预测张量[B, C, H, W]与目标张量[B, H, W]的输入组合,内部完成one-hot转换
  • 可配置参数:gamma控制难分类样本的聚焦程度,alpha解决类别不平衡问题
  • 归约方式对齐官方:支持mean/sum/none三种损失归约模式,与PyTorch内置损失函数行为一致

使用示例

# 模拟模型输出:2个样本,5个类别,64x64图像尺寸
inputs = torch.randn(2, 5, 64, 64)
# 模拟目标标签:每个像素为0-4的类别索引
targets = torch.randint(0, 5, (2, 64, 64))

# 初始化Focal Loss,配置类别权重与聚焦系数
criterion = FocalLoss(alpha=[0.25, 0.25, 0.25, 0.125, 0.125], gamma=2.0)
loss = criterion(inputs, targets)
print(loss.item())

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 14:22:14