【问题标题】:gradient descent MATLAB script梯度下降 MATLAB 脚本
【发布时间】:2015-11-20 19:30:13
【问题描述】:

所以我编写了以下 MATLAB 代码作为 梯度下降 的练习。我显然选择了一个最小值为 (0,0) 的函数,但算法将我抛出到 (-3,3)。

我确实发现在xGradyGrad 之间在线切换:[xGrad,yGrad] = gradient(f); 可以实现正确的收敛,尽管xGradyGrad 与预期的差不多2*X2*Y .我想我在这里颠倒了一些东西,但我一直在试图弄清楚它是什么,但我不明白,所以我希望有人能注意到我的错误......

dx=.01;
dy=.01;
x=-3:dx:3;
y=-3:dy:3;
[X,Y]=meshgrid(x,y);

f=X.^2+Y.^2;

lr = .1; %learning rate
eps = 1e-10; %epsilon threshold
tooMuch = 1e5; %limit iterations
p = [.1 1]; %starting point
[~, idx] = min( abs(x-p(1)) ); %index of closest value
[~, idy] = min( abs(y-p(2)) ); %index of closest value
p = [x(idx) y(idy)]; %closest point to start
[xGrad,yGrad] = gradient(f); %partial derivatives of f
xGrad = xGrad/dx; %scale correction
yGrad = yGrad/dy; %scale correction

for i=1:tooMuch %prevents too many iterations    
    fGrad = [ xGrad(idx,idy) , yGrad(idx,idy) ]; %gradient's definition
    pTMP = p(end,:) - lr*fGrad; %gradient descent's core
    [~, idx] = min( abs(x-pTMP(1)) ); %index of closest value
    [~, idy] = min( abs(y-pTMP(2)) ); %index of closest value
    p = [p;x(idx) y(idy)]; %add the new point
    if sqrt( sum( (p(end,:)-p(end-1,:)).^2 ) ) < eps %check conversion
        break
    end
end

感谢所有帮助的人

编辑:纠正错别字并使代码更清晰。它仍然做同样的事情并且有同样的问题

【问题讨论】:

  • 切向评论:在这样的网格上预先计算梯度是不寻常的。我认为你不应该在网格上工作。它将大大简化代码,并且更正确并提供更好的结果。
  • 我最终将它显示为箭头图(使用quiver),但如果你能告诉我一种更准确地计算梯度的方法,我会很高兴(也许通过链接到它的一些文档)
  • 我认为这主要是出于学习目的。对于这样的问题,没有必要将所有内容都网格化。您可以在没有索引、网格等的情况下跟踪 x_val 和 y_val……并通过自己获取偏导数来计算梯度:fGrad = [2*x_val, 2*y_val]。对于许多功能,您实际上可以使用自动微分自动计算梯度(有一些包)。
  • 我需要它运行更多的功能,例如我稍后检查-20*(X/2-X.^2-Y.^5)*exp(-X.^2-Y.^2) 的收敛性。所以在我看来,一般功能会更好
  • 如果你有符号数学工具箱,那么我建议使用funstr='-20*(X/2-X^2-Y^5)*exp(-X^2-Y^2)'; syms X Y; mygrad=matlabFunction([diff(funstr,X) diff(funstr,Y)]); 之类的东西来计算符号函数的梯度。请注意公式中缺少句点。

标签: matlab gradient


【解决方案1】:

meshgrid 返回的 X 矩阵的 X 值在列中增加,而不是在行中!例如[X, Y] = meshgrid(-1:1, 1:3) 返回

     [-1  0  1;           [1  1  1;
X  =  -1  0  1;       Y =  2  2  2;
   =  -1  0  1];           3  3  3];

注意x-index应该如何放在X或Y的列中,y-index应该放在行中。具体来说,您的线路:

fGrad = [ xGrad(idx,idy) , yGrad(idx,idy) ]; %gradient's definition

应该是:

fGrad = [ xGrad(idy,idx) , yGrad(idy,idx) ]; %gradient's definition

idy 变量应该索引 idx 变量应该索引

【讨论】:

  • 我知道为什么我写错了,但是当我尝试你的解决方案时它仍然不起作用(它在X.^2+Y.^2 上有效,但在-20*(X/2-X.^2-Y.^5)*exp(-X.^2-Y.^2) 上无效)
  • 我认为问题不是直接的编码问题,而是更复杂的数值优化问题。这个问题是deeply not-convex,除非你对其施加一些限制。
  • 如果你看图,到处都是局部最小值。
  • @OdedSayar 如果 Matthew 的回答解决了您的问题,请考虑将他的回答标记为已接受。如果您自己解决了问题,请考虑添加您自己的答案并接受(48 小时后),以表明问题已解决。
  • 谢谢,我没有直接得到答案,而是找到了另一种写算法的方法。我会复制到这里并按照你的建议将其标记为答案。
【解决方案2】:

最终我没有弄清楚以前的方法有什么问题,但是这里有一个梯度体面的替代脚本,我用它来解决同样的问题:

syms x y
f = -20*(x/2-x^2-y^5)*exp(-x^2-y^2); %cost function
% f = x^2+y^2; %simple test function

g = gradient(f, [x, y]);
lr = .01; %learning rate
eps = 1e-10; %convergence threshold
tooMuch = 1e3; %iterations' limit
p = [1.5 -1]; %starting point
for i=1:tooMuch %prevents too many iterations
    pGrad = [subs(g(1),[x y],p(end,:)) subs(g(2),[x y],p(end,:))]; %computes gradient
    pTMP = p(end,:) - lr*pGrad; %gradient descent's core
    p = [p;double(pTMP)]; %adds the new point
    if sum( (p(end,:)-p(end-1,:)).^2 ) < eps %checks convergence
        break
    end
end
v = -3:.1:3; %desired axes
[X, Y] = meshgrid(v,v);
contour(v,v,subs(f,[x y],{X,Y})) %draws the contour lines 
hold on
quiver(v,v,subs(g(1), [x y], {X,Y}),subs(g(2), [x y], {X,Y})) %draws the gradient directions 
plot(p(:,1),p(:,2)) %draws the route
hold off
suptitle(['gradient descent route from ',mat2str(round(p(1,:),3)),' with \eta=',num2str(lr)])
if i<tooMuch
    title(['converged to ',mat2str(round(p(end,:),3)),' after ',mat2str(i),' steps'])
else
    title(['stopped at ',mat2str(round(p(end,:),3)),' without converging'])
end

只是部分结果

在后一种情况下,您可以看到它没有收敛,但 梯度下降 没有问题,只是学习率设置得太高(因此它反复错过了最小值)。

欢迎使用它。

【讨论】:

  • 虽然我们正在等待 48 小时的宽限期,在此之后您可以接受自己的答案,请考虑将图像直接上传到 Stack Overflow。这样,只要 SO 存在,您的答案就会保持完整。
  • 哦,好吧,我在考虑空间效率
  • 这是一个很好的观点,但链接衰减是一个更大的问题:) 这也是不接受仅链接答案的原因。
猜你喜欢
  • 1970-01-01
  • 2014-07-22
  • 2014-03-14
  • 2011-07-04
  • 2013-10-29
  • 1970-01-01
  • 2016-06-13
  • 1970-01-01
  • 2017-02-18
相关资源
最近更新 更多