PyTorch仅使用基于权重的自定义损失时无梯度报错是什么原因?
问题原因分析
你报错的核心原因是代码中使用了.data获取网络权重,导致损失计算完全脱离了PyTorch的自动求导计算图,具体逻辑如下:
- PyTorch中张量的
.data属性会返回一个和原张量共享存储,但requires_grad被强制设为False的新张量,所有基于这个新张量的运算都不会被纳入计算图跟踪,自然也不会生成grad_fn。 - 你之前保留MSE损失时没有报错,是因为MSE损失基于模型前向输出
output计算,output本身是依赖网络权重、带有完整计算图的张量,两部分损失相加后,MSE部分提供了完整的grad_fn,反向传播链路不会断裂。但你之前的写法还有隐藏问题:基于.data计算的自定义权重损失项,梯度根本无法回传给网络权重,这个损失项从始至终都没有生效。 - 当你移除MSE损失后,整个损失的计算全基于脱离计算图的
.data张量,最终输出的损失张量没有grad_fn,反向传播时就会抛出你遇到的报错。
解决方案
直接移除权重获取时的.data后缀,使用网络权重张量本身参与损失计算即可,修正后的代码如下:
def custom_loss(output, target): # 去掉.data,直接使用权重张量保留计算图跟踪 weights = net.linear_layer.weight # 额外补充device参数,避免多卡训练时张量设备不匹配报错 return torch.linalg.norm(weights @ weights.T - torch.eye(weights.shape[0], device=weights.device))
提示:如果你的训练迭代逻辑不强制要求损失函数接收
output、target两个入参,也可以直接把这两个冗余参数删掉,不影响损失计算效果。
内容的提问来源于stack exchange,提问作者Kenny Smith
相关产品推荐
相关产品推荐

