求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
相关产品推荐
相关产品推荐

