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

使用Focal Loss Function平衡TON_IoT数据集时遇报错求助

问题分析与解决

从代码和报错信息来看,核心问题是数据类型不匹配,同时可能存在输入维度对齐问题:

  1. 报错RuntimeError: expected scalar type Float but found Long明确显示:Focal Loss计算时,模型输出的logits是Float类型,但标签是Long类型,二者运算时类型不兼容。
  2. 结合你用BCELoss实现二分类Focal Loss的场景,还需确保模型输出维度和标签维度一致,避免广播错误。

具体修复步骤

  • 统一数据类型:
    计算损失前将标签转换为浮点型,和logits类型对齐:

    loss = focal_loss(output, target.float())
    

    或者在Focal Loss函数内部直接处理类型转换:

    def focal_loss(inputs, targets, alpha=0.8, gamma=2):
        targets = targets.float()  # 新增类型转换
        BCE_loss = nn.BCELoss(reduction='none')(inputs, targets)
        pt = torch.exp(-BCE_loss)
        F_loss = alpha * (1-pt)**gamma * BCE_loss
        return torch.mean(F_loss)
    
  • 对齐输入维度:
    二分类场景下,若模型输出形状为(batch_size,1),而标签是(batch_size,),需要给标签增加维度:

    target = target.unsqueeze(1)
    

    或者将模型输出压缩维度:

    output = output.squeeze()
    
  • 优化损失函数实现:
    推荐用BCEWithLogitsLoss替代BCELoss,它自带sigmoid激活,数值稳定性更好,适配二分类Focal Loss:

    import torch
    import torch.nn as nn
    
    class FocalLoss(nn.Module):
        def __init__(self, alpha=0.8, gamma=2, reduction='mean'):
            super().__init__()
            self.alpha = alpha
            self.gamma = gamma
            self.reduction = reduction
            self.bce_loss = nn.BCEWithLogitsLoss(reduction='none')
    
        def forward(self, inputs, targets):
            targets = targets.float()
            bce_loss = self.bce_loss(inputs, targets)
            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)
            return focal_loss
    

    使用示例:

    loss_fn = FocalLoss(alpha=0.8, gamma=2)
    loss = loss_fn(output, target.float())
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 12:12:15