PyTorch自定义损失函数梯度计算原理及不可导点无报错原因
PyTorch自定义损失梯度计算相关问题解答
你提到的自定义铰链损失实现代码如下:
class MarginRankingLossExp(nn.Module): def __init__(self) -> None: super(MarginRankingLossExp, self).__init__( ) def forward(self,input1,input2,target): # loss_without_reduction = max(0, −target * (input1 − input2) + margin) neg_target = -target input_diff = input2-input1 mul_target_input = neg_target*input_diff add_margin = mul_target_input zeros=torch.zeros_like(add_margin) loss = torch.max(add_margin, zeros) return loss.mean()
自定义损失的梯度计算逻辑
这类只实现了构造函数和forward方法的自定义损失,完全通过PyTorch内置的Autograd自动微分引擎完成梯度计算,不需要手动编写反向传播逻辑,核心原理如下:
- 你在
forward中使用的所有运算(张量取负、逐元素减法、逐元素乘法、构造全零张量、逐元素取最大值、求均值)全都是PyTorch原生实现的算子,每个算子都预先注册了对应的反向求导规则。 - 前向传播执行时,Autograd会自动跟踪所有开启梯度记录的张量运算,按执行顺序构建动态计算图,记录每个输出和输入之间的依赖关系。
- 反向传播触发时,Autograd会从最终的损失张量出发,沿着动态计算图反向回溯,通过链式法则依次调用每个算子内置的求导逻辑,自动算出所有可学习参数对应的梯度值。
- 你自定义的
nn.Module子类本质只是对原生算子的逻辑封装,没有引入Autograd无法识别的外部运算,自然不需要额外实现反向方法就能正常完成梯度计算。
不可导点运行不报错的原因
数学上铰链损失在分段临界点(也就是你提到的y=margin位置,对应当前实现里add_margin和0值相等的点)确实不存在严格的双侧导数,运行时不报错是框架设计层面的合理选择:
- 深度学习框架的求导逻辑不要求严格满足数学上的可导定义,只需要存在可用于参数更新的次梯度即可。对于
torch.max这类分段算子,PyTorch在两个输入值相等的不可导点,会直接预设一个次梯度返回(比如max(x,0)在x=0点默认返回0梯度),不会因为数学定义上的不可导中断计算。 - 这类临界点在实际训练中属于零测集:由于深度学习模型的参数、输入都是浮点数值,前向传播的结果刚好精确落在不可导点上的概率极低,就算真的命中,单次更新带来的偏差也完全可以忽略,几乎不会对训练收敛造成影响。
- 框架本身没有必要为了这种极低概率、不影响实际训练效果的边界点抛出错误打断训练流程。
内容的提问来源于stack exchange,提问作者Anirudha Kulkarni
相关产品推荐
相关产品推荐

