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

如何修正PyTorch梯度下降代码以正确求解目标函数最小值

代码错误原因

你的代码存在4个核心语法/逻辑错误,导致无法正常运行:

  • 可训练参数未正确注册:直接在__init__里写self.x = 1是普通Python整数,不会被PyTorch识别为可训练参数,优化器无法获取该参数进行更新,必须用nn.Parameter包装为浮点型张量,且开启梯度计算。
  • 类方法缩进错误:forward方法没有缩进在MyModel类内部,属于全局函数,不会被模型实例正确调用。
  • 变量引用错误:循环里使用的x是未定义的全局变量,实际上x是模型内部的参数,不需要作为入参传给forward方法。
  • 终止条件逻辑错误:梯度下降收敛判断应该用梯度的绝对值小于阈值,你写的x.grad < 0.001会在梯度为负(哪怕梯度绝对值很大)时直接触发终止,且引用的是不存在的全局x的梯度,不是模型内部参数的梯度。

额外说明:你给出的目标函数f(x) = 2*(x³+1)^(-1/2)在定义域x > -1上是单调递减函数,导数f’(x) = -3x²/(x³+1)^(3/2) ≤ 0恒成立,不存在有限的全局最小值(x趋近于正无穷时f(x)趋近于0),梯度下降会让x持续增大,不会自然收敛。如果你的目标是找驻点(x=0处梯度为0),可以调整终止条件或者检查目标函数是否书写正确。

修正后的可运行代码
import torch
import torch.nn as nn
from torch import optim

alpha = 0.1
grad_threshold = 1e-3  # 梯度收敛阈值
max_iter = 10000  # 增加最大迭代次数防止无限循环,适配原函数无有限最小值的特性

class MyModel(nn.Module):
    def __init__(self):
        super().__init__()
        # 正确注册可训练参数:对应初始值x(0)=1
        self.x = nn.Parameter(torch.tensor([1.0], requires_grad=True))

    # 修正缩进,forward为类内部方法,无需额外传入参数x
    def forward(self):
        return 2 * torch.pow((torch.pow(self.x, 3) + 1), -1/2)

model = MyModel()
optimizer = optim.SGD(model.parameters(), lr=alpha)
terminationCond = False
iter_count = 0

while not terminationCond:
    optimizer.zero_grad()  # 迭代开始先清零梯度,防止梯度累加
    f = model()  # 前向计算函数值
    f.backward()  # 反向传播计算梯度
    
    # 打印迭代过程方便观察
    if iter_count % 20 == 0:
        print(f"迭代次数{iter_count}, x={model.x.item():.4f}, f(x)={f.item():.4f}, 梯度={model.x.grad.item():.6f}")
    
    optimizer.step()  # 按梯度更新规则更新参数
    iter_count += 1

    # 修正终止条件:判断参数梯度的绝对值小于阈值,或达到最大迭代次数
    if torch.abs(model.x.grad) < grad_threshold or iter_count >= max_iter:
        terminationCond = True

print(f"\n迭代结束,最终x={model.x.item():.4f}, f(x)={model().item():.4f}")
运行说明
  • 若你确实要优化当前给出的函数,运行代码后会看到x持续增大、f(x)持续减小,直到触发最大迭代次数停止,符合该函数单调递减的数学特性。
  • 若你的目标函数存在笔误(比如是f(x)=2*(x²+1)^(-1/2)这类存在有限极小值的函数),只需要修改forward里的函数表达式即可,上述代码框架可以正常运行。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 21:33:19