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

U²Net启用AMP混合精度训练报BCELoss autocast错误如何解决

问题背景

在实现面向显著性目标检测任务的U²Net时,原代码未针对训练流程做优化,参照PyTorch官方AMP(自动混合精度)训练文档,在个人代码分支中对原代码做了适配修改,用于验证混合精度训练的效果。
使用修改后的代码在Colab平台运行训练脚本,执行命令如下:

! git clone U-2-Net项目代码仓库
%cd ./U-2-Net/
!python u2net_train.py

运行后触发报错,经排查问题源于自定义损失函数muti_bce_loss_fusion,该损失函数实现代码如下:

bce_loss = nn.BCELoss(size_average=True)

def muti_bce_loss_fusion(d0, d1, d2, d3, d4, d5, d6, labels_v):

    loss0 = bce_loss(d0,labels_v)
    loss1 = bce_loss(d1,labels_v)
    loss2 = bce_loss(d2,labels_v)
    loss3 = bce_loss(d3,labels_v)
    loss4 = bce_loss(d4,labels_v)
    loss5 = bce_loss(d5,labels_v)
    loss6 = bce_loss(d6,labels_v)

    loss = loss0 + loss1 + loss2 + loss3 + loss4 + loss5 + loss6
    return loss0, loss

模型定义文件第526行(即模型输出层)会先返回7个经过sigmoid激活的预测值,再传入上述损失函数,对应代码如下:

F.sigmoid(d0), F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)
报错信息

运行时抛出的完整错误栈如下:

/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py:780: UserWarning: Note that order of the arguments: ceil_mode and return_indices will changeto match the args list in nn.MaxPool2d in a future release.
  warnings.warn("Note that order of the arguments: ceil_mode and return_indices will change"
/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py:3704: UserWarning: nn.functional.upsample is deprecated. Use nn.functional.interpolate instead.
  warnings.warn("nn.functional.upsample is deprecated. Use nn.functional.interpolate instead.")
/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py:1944: UserWarning: nn.functional.sigmoid is deprecated. Use torch.sigmoid instead.
  warnings.warn("nn.functional.sigmoid is deprecated. Use torch.sigmoid instead.")
Traceback (most recent call last):
  File "u2net_train.py", line 148, in <module>
    loss2, loss = muti_bce_loss_fusion(d0, d1, d2, d3, d4, d5, d6, labels_v)
  File "u2net_train.py", line 33, in muti_bce_loss_fusion
    loss0 = bce_loss(d0,labels_v)
  File "/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py", line 1110, in _call_impl
    return forward_call(*input, **kwargs)
  File "/usr/local/lib/python3.7/dist-packages/torch/nn/modules/loss.py", line 612, in forward
    return F.binary_cross_entropy(input, target, weight=self.weight, reduction=self.reduction)
  File "/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py", line 3065, in binary_cross_entropy
    return torch._C._nn.binary_cross_entropy(input, target, weight, reduction_enum)
RuntimeError: torch.nn.functional.binary_cross_entropy and torch.nn.BCELoss are unsafe to autocast.
Many models use a sigmoid layer right before the binary cross entropy layer.
In this case, combine the two layers using torch.nn.functional.binary_cross_entropy_with_logits
or torch.nn.BCEWithLogitsLoss.  binary_cross_entropy_with_logits and BCEWithLogits are
safe to autocast.
可行解决方案

报错核心原因是AMP自动混合精度场景下,nn.BCELoss本身不支持autocast运算,PyTorch官方明确要求sigmoid激活+BCE损失的组合要替换为内置sigmoid实现的BCEWithLogitsLoss,从根源规避精度异常问题。

推荐方案(无精度损失、符合AMP规范、数值稳定性更好)

分两步修改:

  • 修改模型输出层,去掉输出前的sigmoid激活,直接返回原始logits:
    # 替换原有的F.sigmoid包裹的输出逻辑,直接返回原始输出
    d0, d1, d2, d3, d4, d5, d6
    
    推理阶段如果需要0-1区间的概率输出,单独在推理代码中对输出做sigmoid运算即可,不影响训练逻辑。
  • 修改损失函数定义,将nn.BCELoss替换为nn.BCEWithLogitsLoss,其余多损失融合逻辑保持不变:
    # 替换原有的bce_loss定义,注意原size_average参数已弃用,用reduction='mean'等价替代
    bce_loss = nn.BCEWithLogitsLoss(reduction='mean')
    
    def muti_bce_loss_fusion(d0, d1, d2, d3, d4, d5, d6, labels_v):
        loss0 = bce_loss(d0,labels_v)
        loss1 = bce_loss(d1,labels_v)
        loss2 = bce_loss(d2,labels_v)
        loss3 = bce_loss(d3,labels_v)
        loss4 = bce_loss(d4,labels_v)
        loss5 = bce_loss(d5,labels_v)
        loss6 = bce_loss(d6,labels_v)
    
        loss = loss0 + loss1 + loss2 + loss3 + loss4 + loss5 + loss6
        return loss0, loss
    

临时兼容方案(不修改模型结构,损失少量加速收益)

如果不想调整模型输出逻辑,可以在损失计算前,将输入的预测值和标签手动转换为float32精度,再送入BCELoss计算,跳过autocast对该段逻辑的精度转换:

bce_loss = nn.BCELoss(size_average=True)

def muti_bce_loss_fusion(d0, d1, d2, d3, d4, d5, d6, labels_v):
    loss0 = bce_loss(d0.float(),labels_v.float())
    loss1 = bce_loss(d1.float(),labels_v.float())
    loss2 = bce_loss(d2.float(),labels_v.float())
    loss3 = bce_loss(d3.float(),labels_v.float())
    loss4 = bce_loss(d4.float(),labels_v.float())
    loss5 = bce_loss(d5.float(),labels_v.float())
    loss6 = bce_loss(d6.float(),labels_v.float())

    loss = loss0 + loss1 + loss2 + loss3 + loss4 + loss5 + loss6
    return loss0, loss

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 20:09:34