PyTorch更新叶子变量遇RuntimeError,如何正确实现梯度下降求二次函数根
PyTorch梯度下降求解二次方程(理解autograd)
问题根源
你遇到的两个错误本质都是对PyTorch计算图和autograd机制的误解:
- 多次调用
.backward()报错:默认情况下,PyTorch执行.backward()后会自动销毁计算图,第二次调用时没有可用的图结构,因此抛出错误。 - 添加
retain_graph=True后出现原地操作错误:直接对x做原地修改(比如x -= lr * x.grad)会破坏计算图中依赖的原始张量,而autograd需要这些原始值来完成梯度回溯,因此触发RuntimeError。
另外补充:你要解的二次方程y=3x²+4x+9判别式为负,没有实根,我们实际是通过梯度下降找函数的最小值点(梯度为0的位置),这同样能很好地演示autograd的工作逻辑。
正确实现代码
以下是完整的可运行代码,附带关键步骤说明:
import torch # 初始化可训练变量,requires_grad=True开启梯度追踪 x = torch.tensor(0.0, requires_grad=True) learning_rate = 0.1 iterations = 100 for i in range(iterations): # 每次迭代重新计算目标函数(构建新的计算图) y = 3 * x ** 2 + 4 * x + 9 # 最小化y²等价于逼近y=0的点,也可以直接用y,但y²的梯度更适合梯度下降 loss = y ** 2 # 清空上一次迭代的梯度,避免累加 if x.grad is not None: x.grad.zero_() # 计算梯度,无需retain_graph=True——因为每次迭代都重新构建计算图 loss.backward() # 在no_grad上下文更新x,禁用梯度追踪,避免原地修改破坏计算图 with torch.no_grad(): x -= learning_rate * x.grad # 每10次迭代打印状态 if (i + 1) % 10 == 0: print(f"Iteration {i+1}: x = {x.item():.4f}, y = {y.item():.4f}") # 最终结果:逼近函数最小值点x=-2/3≈-0.6667 print(f"\nFinal x: {x.item():.4f}, Corresponding y: {3*x.item()**2 +4*x.item() +9:.4f}")
关键细节解释
- 重新构建计算图:每次迭代都重新计算
y和loss,这样每次都会生成新的计算图,不需要保留之前的图结构,因此无需retain_graph=True。 - 清空梯度:
x.grad.zero_()必须在每次计算梯度前执行,否则梯度会和上一次的结果累加,导致参数更新方向错误。 - 用no_grad更新参数:
torch.no_grad()上下文会临时关闭autograd的追踪功能,此时对x的修改不会被记录到计算图中,避免了原地操作破坏计算图的问题。
内容的提问来源于stack exchange,提问作者MattWright
相关产品推荐
相关产品推荐

