You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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}")

关键细节解释

  1. 重新构建计算图:每次迭代都重新计算y和loss,这样每次都会生成新的计算图,不需要保留之前的图结构,因此无需retain_graph=True。
  2. 清空梯度:x.grad.zero_()必须在每次计算梯度前执行,否则梯度会和上一次的结果累加,导致参数更新方向错误。
  3. 用no_grad更新参数:torch.no_grad()上下文会临时关闭autograd的追踪功能,此时对x的修改不会被记录到计算图中,避免了原地操作破坏计算图的问题。

内容的提问来源于stack exchange,提问作者MattWright

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 04:37:23