PyTorch中将Sigmoid浮点输出转标签时IoULoss无梯度报错如何解决?
报错原因
你自定义的IoULoss中使用的inputs > 0.5阈值二值化、&/|位运算都属于不可导的离散操作,会直接打断PyTorch的反向传播计算图,导致最终输出的loss没有梯度传递路径,因此触发了无grad_fn的报错。
解决方案
训练阶段不能对Sigmoid的输出做硬阈值转标签的操作,你可以使用可导的软IoU损失实现,直接用Sigmoid输出的[0,1]区间连续值参与计算:
def IoULoss(inputs, targets, smooth=1e-6): # 展平张量,inputs为模型输出的sigmoid结果,值域0~1 inputs = inputs.view(inputs.size(0), -1) # 确保标签为float类型 targets = targets.view(targets.size(0), -1).float() # 可导的交集、并集计算 intersection = (inputs * targets).sum(1) union = inputs.sum(1) + targets.sum(1) - intersection IoU = (intersection + smooth) / (union + smooth) return 1 - IoU.mean()
上述实现全程使用可导运算,不会打断计算图,可以正常反向传播训练。
补充说明
如果你需要在验证/推理阶段得到最终的二值分割标签,可以在不计算梯度的场景下使用(pred > 0.5).long()做硬转换,该操作不需要反向传播,不会触发报错。
内容的提问来源于stack exchange,提问作者Wrong Wizzli
相关产品推荐
相关产品推荐

