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

求适用于多类别目标检测的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 02:20:25