求适用于多类别目标检测的PyTorch官方Focal Loss可靠实现
多类别Focal Loss的PyTorch实现方案
官方现状
PyTorch官方目前没有内置针对多分类场景的Focal Loss实现,你需要自定义实现来匹配nn.CrossEntropyLoss()的使用方式。
适配输入形状的可靠实现
以下是适配(b, c)预测输出和(b)目标标签的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 if alpha is not None: self.alpha = torch.tensor(alpha, dtype=torch.float32) self.gamma = gamma self.reduction = reduction def forward(self, inputs, targets): ce_loss = F.cross_entropy(inputs, targets, reduction='none') pt = torch.exp(-ce_loss) if self.alpha is not None: alpha = self.alpha[targets].to(inputs.device) focal_loss = alpha * (1 - pt) ** self.gamma * ce_loss else: focal_loss = (1 - pt) ** self.gamma * ce_loss if self.reduction == 'mean': return focal_loss.mean() elif self.reduction == 'sum': return focal_loss.sum() else: return focal_loss
使用方式
和nn.CrossEntropyLoss()完全一致:
# 初始化损失函数 focal_loss = FocalLoss(alpha=[0.25, 0.25, 0.5], gamma=2.0) # alpha为可选的类别权重参数 # 模拟输入 inputs = torch.randn(32, 5) # batch_size=32,类别数=5 targets = torch.randint(0, 5, (32,)) # 计算损失 loss = focal_loss(inputs, targets)
第三方实现参考
如果需要参考成熟的第三方实现,可关注主流目标检测框架(如MMDetection、YOLO系列)中的Focal Loss模块,它们的实现经过大量验证,且同样适配你提到的输入形状。
内容的提问来源于stack exchange,提问作者Infintyyy
相关产品推荐
相关产品推荐

