QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2721|回复: 0
打印 上一主题 下一主题

随机梯度下降算法SGD(Stochastic gradient descent)

[复制链接]
字体大小: 正常 放大

1194

主题

4

听众

2951

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |正序浏览
|招呼Ta 关注Ta
SGD是什么3 s2 Q+ x% [, V9 Y: E/ v% O
SGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。# {) o* T, Q3 R
怎么理解梯度?9 N, m, Z8 u/ h$ K) m- Z
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。
  p' s$ `2 ]5 ~+ z: W( m6 s+ t$ o  ~
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
* X' J' H+ D1 t4 {' a4 W4 I, F
1 [  I2 E2 h. h' w每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
0 \4 I- L8 E: J( T* X4 K6 `: r, V  J: e0 T" j" u2 s  H
为什么引入SGD
3 `# H, c4 \4 d8 C/ H深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。$ S0 L! i/ t$ C6 D8 j0 x: f
* j" Z5 h3 m3 N5 a
怎么用SGD
  1. import torch* S8 M, Z' {) v- j\" n8 W, v

  2. 7 U3 ]4 Z0 m7 p
  3. from torch import nn
    6 _! n6 |0 H  [$ e) s

  4. # n& S& m2 {9 {& g8 [* b, a
  5. from torch import optim
    1 w! G+ l% b& w. {% X: q. R

  6. 1 u) \; j5 M6 a& P! x/ Z3 S

  7. * {0 R8 b. Y: o7 _; o# \
  8. 3 o: x, `& v7 M\" Z+ G/ P( Q
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    \" B! M8 I4 Z3 w6 R* i$ ]# l2 h

  10. , L' S; r/ c; D! g# n
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
    \" _\" o! ?/ S) Y' H

  12. ( S7 i( ]0 w8 \3 b8 f. o

  13. 3 F, Q( K9 T( o4 M
  14. + I& u2 s' {. |* t5 M1 ?
  15. model = nn.Linear(2, 1)4 ?$ e\" L- X* S4 R

  16. / k& X7 k# v: i  s; J/ I( Z& l
  17. + p' B# Z# U( k3 _

  18. \" x  U; W' e5 J1 g
  19. def train():% m' Y0 N0 j% b- w

  20. . S1 |2 o: _( W, y  m( @6 j
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)  u( J- m. {0 R0 ~' V

  22. , n: |7 J  j* |  G. p- z
  23.     for iter in range(20):0 R- T\" I1 u  V0 c( {2 ]7 [
  24. 9 F7 `0 U  J/ k0 z' H5 v
  25.         # 1) 消除之前的梯度(如果存在)  d4 U$ K; m3 z

  26. 1 j\" Z6 q9 q0 W8 w7 l
  27.         opt.zero_grad()# N8 m\" @/ E- w
  28. : E4 W1 d\" d& ~! p9 k, H

  29. + y% z( t- Q5 I) m# m7 `, I  ]' ^7 W

  30. 9 f: C$ \0 x. l8 K& q
  31.         # 2) 预测\" X; V) ~+ s0 |! K  E

  32. ( q: T5 r/ W8 p( ~3 X
  33.         pred = model(data)
    ) d; k9 G\" J4 H. N
  34. \" R9 d9 B6 T) [\" S3 j$ k3 R. a; s
  35. 3 v0 O- j\" F5 T: [

  36. ! u9 J3 F* d+ c
  37.         # 3) 计算损失
    & @( I2 e& c0 ~8 j

  38. # w& G  P4 p) m
  39.         loss = ((pred - target)**2).sum()* b% k) Z% X; q2 y

  40. + V; ~( p+ X* k

  41. 2 b! E3 K! v5 ?

  42. $ `2 C% J9 c\" V: e7 _
  43.         # 4) 指出那些导致损失的参数(损失回传)8 ]( E  V+ t) T5 ^. T$ \+ |

  44. $ N( v; @: D3 o) S6 |
  45.         loss.backward()
    8 C8 r1 J8 \6 o8 c

  46. \" J8 i# Z  }; ^2 e9 |; S
  47.     for name, param in model.named_parameters():
    9 d\" r1 H\" W\" n; e5 Y2 N$ f- s; j

  48. . v; x2 C7 w' i( x. Z
  49.             print(name, param.data, param.grad)
    ! y: {! j7 d* n* c3 s
  50. ' n2 i) d1 l# z4 e0 ?0 `
  51.         # 5) 更新参数' Z9 h( l' s, m! a1 }! c
  52. * L: k' b/ g( ]4 J; G/ v2 i
  53.         opt.step()
    3 T$ T0 }# S7 |5 z

  54. ' c, ^3 i5 m3 ?$ f; c! g! h

  55. $ ]# v) x) Z- q5 J3 p

  56. % q# ?+ }1 K5 Q  }, k0 G: V
  57.         # 6) 打印进程! ~+ Y, `+ g& p& z) e

  58. 7 h  n* H8 }; w2 E
  59.         print(loss.data)4 g: U, `5 g; G) V( |' D
  60. 8 s! _- B' b6 \& D
  61. , k& Q7 ^4 R0 e1 n7 V\" d
  62. # S3 [- E5 @! J
  63. if __name__ == "__main__":6 r- ?% ?2 q2 g  K4 q\" [

  64. ' k  _, W  q! X; _, N
  65.     train()8 q1 _. h, A, Q/ a
  66. ) [4 z4 S& T$ O5 H9 g( n
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。% d0 g% {2 p+ w$ g
2 o; j2 y7 r- b4 q  f
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])2 S$ }+ U. h+ r  P; o9 ?; _
  2. : x+ {5 j1 J2 V  D- \
  3. bias tensor([-0.2108]) tensor([-2.6971])
    8 x, J\" t4 D+ F6 y& m

  4. : B& f4 x2 D6 i
  5. tensor(0.8531)& w8 `7 Z0 [1 E) Q
  6. \" `6 N- S: B2 c5 o( C+ a$ j0 P
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])
    ! J+ N0 w7 N3 I  L\" `) O

  8. ( C( ?/ B. J5 @( t4 \% f
  9. bias tensor([0.0589]) tensor([0.7416])
    , i# ~% p. s0 `& G( B- i: D
  10. + T2 \, |% Y3 B, P% M) y
  11. tensor(0.2712)9 T& E. u! j5 f% c, W8 X5 V- X' `! ]( w
  12. 7 g4 D1 e, p3 `+ T% g# z
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    : h, `. \1 `9 U3 r( d/ I

  14. - f# y* @0 C$ l7 B. M0 y( K
  15. bias tensor([-0.0152]) tensor([-0.2023])
    # [( `5 g: ?, c6 k+ h$ @' `

  16. 4 |4 B8 a/ p' T2 R# }3 j3 p% p
  17. tensor(0.1529)& N, A& ]$ G/ V  f( t

  18. ; o( U  H) d. ]- p2 v2 G/ Y+ u
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])
    7 F# O6 V$ o% g/ ^$ F6 X. b! u
  20. ( R4 r! C6 r0 C  `9 ?\" J* r
  21. bias tensor([0.0050]) tensor([0.0566])
    & |+ u/ S5 a( i2 X, N

  22. 7 R\" w2 q: o6 z, b7 T  Z- i
  23. tensor(0.0963)# t8 p1 x& B9 i+ ^5 ^9 C( Y4 D
  24. ( x( g% ]\" V9 x# V) e( n' d3 P
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])
    \" G6 L7 x7 z2 H  N7 Q  j' g3 `8 g( j

  26. 1 Q% o6 X8 j& ~' i9 v& h$ j
  27. bias tensor([-0.0006]) tensor([-0.0146])# O# E+ ^* m6 o2 T* u* t) F
  28. $ Y6 \0 l4 V# {  S/ t4 B( F3 T
  29. tensor(0.0615). U  x& a7 t! G( z& x  _5 \

  30. 9 L7 ?; a; V3 \4 c3 N) X
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])- q% _& C! H5 v
  32. * e! l3 K/ S* E9 Y
  33. bias tensor([0.0008]) tensor([0.0048])6 C8 @; ?  `  m, a1 Q( q8 D+ L

  34. $ F8 c\" P7 z2 k! f6 m: \
  35. tensor(0.0394)' Z! h1 @2 F+ q  m1 N
  36. 5 T) P2 f: ^! i! p% }' D
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])9 _( C, _7 d; i( `. _2 k: w5 O

  38. * l- t+ n  E+ g! w7 G0 @\" g
  39. bias tensor([0.0003]) tensor([-0.0006]), u% y\" s$ U7 F7 K* f0 B  b! B7 A. g/ ~
  40. . \5 I8 j: m8 [* X
  41. tensor(0.0252)8 U1 M- B  G+ r( a

  42. , f0 a6 C$ h& l+ ]$ a% G\" z* p\" [
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])- l+ D) G3 x1 A3 G6 \4 j: w
  44.   u- z9 X7 n' j  T# d: Y
  45. bias tensor([0.0004]) tensor([0.0008])
    . k0 F/ x0 [' q\" p+ Q! `; q) D

  46. 7 S$ D, n4 {6 f' n7 v& g
  47. tensor(0.0161)
    + y  D+ B1 j7 M) {  Z

  48. 2 T+ N% D) B6 e# ?
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])5 `7 m% w- ?/ s  T' E$ C

  50. $ _2 |( B  z  r- @5 y+ k% b
  51. bias tensor([0.0003]) tensor([0.0003])
    * }- r  [\" q9 g  `
  52. ; f$ s% Q3 q! y' [  l
  53. tensor(0.0103)% O! h4 L0 S4 c5 S# }3 m

  54. * |. \  d$ _- _/ [
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])$ t& }. s, C# O\" n$ F/ r
  56. 5 B9 `* v4 T- g9 ?, p- c
  57. bias tensor([0.0003]) tensor([0.0004])% U- V. p8 i3 ]& }

  58. 0 f  X3 D7 u1 c! ]
  59. tensor(0.0066)
    % V% u# H% a7 m  h, l

  60. , @* }3 A$ `9 C. |. b2 m0 m2 y! G
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])
    5 w, p1 Z, K5 _% f/ f

  62. . f0 D4 ?0 h& K3 Y
  63. bias tensor([0.0003]) tensor([0.0003])8 v; h* A/ Y% J
  64. + n8 f3 C$ `4 i
  65. tensor(0.0042)
    7 K7 E, c+ q: i2 `$ j! x

  66. 6 G- [$ C- ~! c\" m9 f' ~
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])% ?7 N/ _* `. s6 }

  68. 6 a( _1 Z: s+ e& z% z$ N  {
  69. bias tensor([0.0002]) tensor([0.0003])
    6 h8 E  [, z! [% B
  70. . O* K, \9 O0 O- v7 i8 u
  71. tensor(0.0027)
    + k\" X1 K% a) Q

  72. 5 o' @! t8 x& g# e+ a
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])
    $ ]( {\" d6 h0 w

  74. % J9 i  ?! k7 h\" D) b5 }+ L$ c
  75. bias tensor([0.0002]) tensor([0.0002])& J$ @  v' M, W* ^8 f

  76. - \. W% [% q& T4 ?& E: ?& A% v# z1 G
  77. tensor(0.0017)' Y7 ^0 |$ M7 ?6 B* R2 _

  78. - k! n5 k9 j3 f0 b% b/ c) p
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])
    0 o* O\" ?- Y# R9 V
  80. + @4 U' W8 A/ |  Z5 ~
  81. bias tensor([0.0002]) tensor([0.0002])) x3 @' N! t2 N1 [' v3 j# G
  82. & a3 Z% Y\" t8 O0 r- E* t; F
  83. tensor(0.0011)! K7 b5 B4 E' P+ F6 z/ {
  84. % B8 C4 O* |$ a\" g3 X( S, W
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]]), j\" V- x1 m! ?. ]7 p  X- j

  86.   w+ O% Y# ~- T2 E1 @% Q& n
  87. bias tensor([0.0001]) tensor([0.0002]), R: T+ `9 ^% n8 I- O$ K
  88. \" U\" m1 _. ]: m! W2 N3 x
  89. tensor(0.0007)
    0 M; T. g# S! I2 t% b+ b* G\" U
  90. 3 D5 h( g: J! L
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])- U  Z; f4 ^, ]: D, L5 ~+ C
  92. 0 ?$ {* v% R* ~; q. x
  93. bias tensor([0.0001]) tensor([0.0002])
    \" y! G8 N& Q4 ~/ P\" Q* P6 ?' w
  94. : n9 W( |8 I( p6 p7 G# O
  95. tensor(0.0005)
    : n7 H: K3 W; K+ M

  96. ; c4 g7 H2 T; T9 i1 M& n$ N
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])5 g, t. n* H1 [
  98. 5 c2 Z7 t9 k  g+ J1 R\" X
  99. bias tensor([0.0001]) tensor([0.0001])
    + O' G( M; N5 n- n3 z\" q
  100. 2 p8 S, y8 n1 I( o6 ?5 }
  101. tensor(0.0003)
    4 J- |; D7 \& [$ @) W

  102. 2 o% c; S8 H9 }
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]]). y# Y/ w3 R$ n3 x  M  M: [
  104. % z/ k( t. d8 I2 R% n
  105. bias tensor([9.7973e-05]) tensor([0.0001])\" [' @  ^- U7 }' v3 A

  106. ! f+ ?0 }2 X* w2 a
  107. tensor(0.0002)9 O0 e# x# Z' C* d6 F% n
  108. ' \! m4 D8 r% u3 W  x7 R3 ~! J4 w
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])! ?, y% t2 N0 u

  110. ' [1 C) U9 e* H+ T: T+ v: a
  111. bias tensor([8.5674e-05]) tensor([0.0001])+ {7 O# z. B) L& X4 y) `4 n

  112. 4 u) f% b8 ^0 Z, s# e9 Q6 A
  113. tensor(0.0001)
    $ K8 q- A6 y% [$ h- j$ ^4 {1 L

  114. / U3 g4 {9 a0 u8 G9 |& T4 L8 h
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])
    7 T: m\" G; ?: h( D: ~% E\" W! w+ W
  116. 9 @% }. e; E$ [& L. ^2 o, c( `/ D- m
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])
    2 |9 k- B+ E. K# T: v- E8 j2 P$ h

  118. ( Z6 N8 H& B! x2 u  `$ o; c
  119. tensor(7.6120e-05)
复制代码

5 c4 A* [4 Z1 N8 ]0 c( r
zan
转播转播0 分享淘帖0 分享分享0 收藏收藏0 支持支持0 反对反对0 微信微信
您需要登录后才可以回帖 登录 | 注册地址

qq
收缩
  • 电话咨询

  • 04714969085
fastpost

关于我们| 联系我们| 诚征英才| 对外合作| 产品服务| QQ

手机版|Archiver| |繁體中文 手机客户端  

蒙公网安备 15010502000194号

Powered by Discuz! X2.5   © 2001-2013 数学建模网-数学中国 ( 蒙ICP备14002410号-3 蒙BBS备-0002号 )     论坛法律顾问:王兆丰

GMT+8, 2026-9-7 15:17 , Processed in 0.625273 second(s), 52 queries .

回顶部