梯度下降求解函数f(x)=x−5+√(4−x²)最大值的代码问题排查
梯度下降求函数最大值的代码问题排查与修复
核心问题分析
- 输入参数被硬编码覆盖:函数
f_gradient_descent内部直接赋值x0 = 1.23和eta = 0.001,完全忽略了调用时传入的x0=0、eta=0.01参数,导致用户指定的初始点和学习率根本不生效。 - 迭代方向错误(梯度下降≠梯度上升):我们要找的是函数最大值,梯度下降是用来找最小值的,求最大值应该用梯度上升:
- 原迭代公式(梯度下降,求最小值):
x_{t+1}=x_t−ηf’(x_t) - 正确迭代公式(梯度上升,求最大值):
x_{t+1}=x_t+ηf’(x_t)
原代码用了减号,相当于往函数减小的方向走,自然找不到最大值点。
- 原迭代公式(梯度下降,求最小值):
- 导数计算可简化(非功能性问题):原导数计算代码可以简化,提高可读性:
两者数学上等价,但简化后更直观。# 原代码 return 1 + 1/(2*math.sqrt(4-x**2))*(-2*x) # 简化后 return 1 - x / math.sqrt(4 - x**2)
修复后的完整代码
import math def f(x): return x - 5 + math.sqrt(4 - x**2) def f_derivative(x): # 简化后的导数计算 return 1 - x / math.sqrt(4 - x**2) def f_gradient_descent(x0, eta, n_step): """ Parameter: x0: Start point eta: Learning rate n_step: algorithm will stop after `n_step` cycle """ # 移除硬编码的参数覆盖 x_current = x0 for _ in range(n_step): grad = f_derivative(x_current) # 梯度上升:加梯度乘以学习率 x_current = x_current + eta * grad # 判断梯度是否足够小(接近极值点) if abs(grad) < 1e-5: break return f(x_current) # 用更精确的预期值(2√2 -5 ≈ -2.171572875) assert abs(f_gradient_descent(0, 0.01, 1000) - (2*math.sqrt(2)-5)) < 1e-4
验证结果
修复后调用f_gradient_descent(0, 0.01, 1000)会返回约-2.1716,符合预期的最大值。
内容的提问来源于stack exchange,提问作者Chris
相关产品推荐
相关产品推荐

