基于导数的损失求解网络参数梯度报错的解决方法
解决PyTorch中自动微分计算参数梯度的「One of the differentiated Tensors appears to not have been used in the graph」报错
问题场景
我需要训练一个预测y(x)的神经网络,但只有dy(x)的数据集(即已知特定x对应的y的导数,不知道y本身)。以下是极简复现代码:
import torch # 定义预测y(x)的网络 network = torch.nn.Sequential( torch.nn.Linear(1, 50), torch.nn.Tanh(), torch.nn.Linear(50, 1) ) # 数据集:dy(x) = x,对应原函数y = 0.5x² x = torch.linspace(0,1,100).reshape(-1,1) dy = x # 基于y预测值的导数计算损失 x.requires_grad=True y_pred = network(x) dy_pred = torch.autograd.grad(y_pred, x, grad_outputs=torch.ones_like(y_pred), create_graph=True)[0] loss = torch.mean((dy-dy_pred)**2) # 执行此行报错 gradients = torch.autograd.grad(loss, network.parameters())[0]
执行最后一行时抛出错误:One of the differentiated Tensors appears to not have been used in the graph,但用torch.optim.Adam配合loss.backward()却能正常运行。直接让网络预测dy不是我的可行方案,请问怎么修复?
错误原因
这报错的根源是torch.autograd.grad的默认行为:
- 计算
dy_pred时用了create_graph=True,这会构建一个嵌套的计算图(支持后续对dy_pred做微分)。 - 直接调用
torch.autograd.grad(loss, network.parameters())时,默认retain_graph=False,计算完梯度就会销毁整个计算图。但此时计算图里还有依赖x梯度的节点,这会导致部分参数的梯度路径被破坏,触发报错。
而loss.backward()能正常运行,是因为它会自动处理嵌套计算图的梯度传递,并且默认会把梯度累积到参数的.grad属性里,不会直接销毁计算图(除非显式指定retain_graph=False)。
修复方案
方案1:保留计算图
调用torch.autograd.grad时显式设置retain_graph=True,确保计算图在梯度计算后不被销毁,同时不要直接取索引[0](避免部分参数梯度为None时出错):
# 修复后的梯度计算代码 gradients = torch.autograd.grad(loss, network.parameters(), retain_graph=True)
如果你后续不需要再基于这个计算图做微分,可以用完后手动销毁:
# 销毁计算图释放内存 x.grad = None y_pred.grad = None
方案2:关闭嵌套计算图(仅适用于无需二次微分的场景)
如果你的任务不需要对梯度再做微分,计算dy_pred时可以把create_graph=False,这样后续调用torch.autograd.grad就不会有问题:
# 仅当不需要二次微分时使用 dy_pred = torch.autograd.grad(y_pred, x, grad_outputs=torch.ones_like(y_pred), create_graph=False)[0] loss = torch.mean((dy-dy_pred)**2) gradients = torch.autograd.grad(loss, network.parameters())
方案3:用backward()替代autograd.grad
如果只是要获取参数梯度来手动更新,更稳妥的方式是用loss.backward()配合参数的.grad属性,这和优化器的工作逻辑一致:
# 用backward()获取梯度 loss.backward() # 提取所有参数的梯度 gradients = [p.grad for p in network.parameters()] # 手动更新参数示例(模拟优化器步骤) with torch.no_grad(): for p in network.parameters(): p -= 0.01 * p.grad # 0.01是学习率 p.grad.zero_() # 清空梯度,避免累积
内容的提问来源于stack exchange,提问作者Thomas Wagenaar
相关产品推荐
相关产品推荐

