使用Focal Loss Function平衡TON_IoT数据集时遇报错求助
问题分析与解决
从代码和报错信息来看,核心问题是数据类型不匹配,同时可能存在输入维度对齐问题:
- 报错
RuntimeError: expected scalar type Float but found Long明确显示:Focal Loss计算时,模型输出的logits是Float类型,但标签是Long类型,二者运算时类型不兼容。 - 结合你用
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
相关产品推荐
相关产品推荐

