数学建模社区-数学中国

标题: 随机梯度下降算法SGD(Stochastic gradient descent) [打印本页]

作者: 2744557306    时间: 2023-11-28 14:57
标题: 随机梯度下降算法SGD(Stochastic gradient descent)
SGD是什么
# m- n- a/ j) d; X* M% [- `/ TSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。
' ]: F7 x# S9 Y+ M, T  P2 J- S; l9 V怎么理解梯度?  ]# g2 a! S9 g% x$ M7 i
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。
, `. F6 n- \6 H! b0 i9 T8 L+ P, k9 f6 R+ d& A; x$ {1 ^- ?
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
3 o& o& I& T: n" a# M$ U' ^* s1 F' y
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
/ e; F0 C1 `- {  z: R, e/ K( o; n( Z( ^- w: l
为什么引入SGD! q+ j! ~' b5 d, ^( |: z
深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。7 H& J9 T- c  w' q. k

4 b; E  {9 H# C, o0 x1 n1 X  a怎么用SGD
  1. import torch+ ~  u8 x- c* J! X+ l, I9 t
  2. / s# b4 T3 x/ p
  3. from torch import nn7 n! Z+ t3 B) L: }' i# T9 u& W6 X

  4. * f, V* f) `5 U& |- M& g
  5. from torch import optim
    & X& S. o' o7 j$ p9 v- ]) L' p; b
  6. ' S0 C( E3 F  h: s7 r
  7. 5 H8 B, K2 T& C

  8. . B" O! v# O2 H  k
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    * k8 Z. l* R5 [( f% {1 N# b
  10. 4 J1 [9 I5 g0 D2 Q" c- X$ l
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)5 X* R# _2 _1 O/ G
  12. : A# j2 Y3 Q: t- u; |

  13. : r* q+ n! F+ E: u% f1 ?9 B

  14. ' ]+ W' A+ n3 i- S- q- L( S* }
  15. model = nn.Linear(2, 1)8 m: C0 ^* }6 }# |/ d- X' {
  16. ! G* ~: k- }3 V
  17. $ E8 X4 J1 P$ J/ M* i7 ~
  18. & @) f2 ]+ `$ P1 j8 V+ F" J
  19. def train():+ ]& u& _$ K% ]

  20. , k* A; ~" l0 k0 V0 J6 g, `
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)6 m) V) p  h& n8 \: B

  22. $ C  w# O2 f; B
  23.     for iter in range(20):
      E. e* V! s; G" [, c) Q2 N" K

  24. % V* {% K3 o/ J' N
  25.         # 1) 消除之前的梯度(如果存在)
    6 j! _4 n/ d. g3 \

  26. 0 F: {' a, K9 l
  27.         opt.zero_grad()
    ! E- i! x! p0 ~: s# e) J

  28. ) F5 e; t) ]% e. B5 V, [

  29. 2 G6 T; W4 ^" m8 D4 k9 L
  30. 3 K- o+ F/ n' p$ Y+ Y* v
  31.         # 2) 预测8 _! T. |+ r2 {6 T
  32. - }& H* |; x9 E) k% }
  33.         pred = model(data)
      {4 ^5 D6 c' k: V/ L) p4 {3 @
  34. + f1 F; P+ Y3 ^# ~7 n4 l; `
  35. . M5 b+ M4 R, ?: C, t, u
  36. / R* D7 S4 [) e
  37.         # 3) 计算损失
    . k- N- y6 O7 M# U( i* A
  38. " W. S2 n$ h7 C2 Y" \; M
  39.         loss = ((pred - target)**2).sum()
    ' O( p* {' S- y

  40. 6 U1 ]3 e7 O7 ]7 `, B! b
  41. # ]7 K7 z3 Y  m0 h9 _9 e
  42. 6 d  Y" y  @( Z! ?
  43.         # 4) 指出那些导致损失的参数(损失回传)$ i* w: U6 i8 u& k: z5 F! u

  44. ( K# n9 f# x0 N# y  R" \/ S
  45.         loss.backward()
    ) x$ }7 k+ L, f& ]) Y
  46. - y) Z" ~& v6 {" F9 F0 S9 C1 G3 d3 I6 J
  47.     for name, param in model.named_parameters():3 `* y. O, i- ~# R) W
  48. 2 Y: J( H' e. Q7 C, M1 `
  49.             print(name, param.data, param.grad)
    5 x6 e/ P: w. E1 x" C7 _6 E
  50. ( k( r9 K! [2 V$ J# j4 F- o
  51.         # 5) 更新参数
    6 p$ Z$ T5 c4 i- t; ~, F% `. i1 Y) H

  52. ) a" k% v- k' B1 ]
  53.         opt.step()% W  [! l0 k- l! A
  54.   J* b* H; K" c

  55. : ^5 p1 d, X/ ]3 q. `
  56. ' [7 F' k; M# R' i9 |+ m
  57.         # 6) 打印进程
    ( E* T" n9 k) x5 V5 T/ ]. t
  58. & b3 B% `, T& z) y
  59.         print(loss.data)  z4 q' x% }. G* L! M% E7 ~
  60. & o, V2 r& e- K: \! E

  61. , Z. [3 z: ^$ [6 z, L

  62. 7 c" V- [% H2 _! r- f+ t) k
  63. if __name__ == "__main__":+ y; R& E7 N% p1 `8 ^

  64. * e8 U; E( I' ^$ `$ t, B
  65.     train()
    8 c; _+ i2 {; }4 ~8 b+ D7 L* W

  66. % v% |' y2 k( A9 f4 p. N6 j
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。1 e  E" v1 d! P4 }4 N% o: y
& q& o+ V. v5 M- B% m
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]]). j" D; Z: j- o+ A5 p0 q, T, A
  2. / j+ ^* L. Q) k: M
  3. bias tensor([-0.2108]) tensor([-2.6971])
    4 x7 r, n! }# X- D, z

  4. . n& l( W* d4 p! \2 B
  5. tensor(0.8531)3 [  r( c! c% q, X2 ~

  6. , V  E+ d$ j6 q7 g- m* i
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])% r$ r* c, ?6 N& j$ u% Z

  8. * |( j( p& e) |- [; P/ C
  9. bias tensor([0.0589]) tensor([0.7416]): m7 m% Y! J5 U8 N8 s( h
  10.   L' I  M1 R) ~9 A4 e# |
  11. tensor(0.2712)4 Y( I3 ~  Q: k, Y/ Q4 T
  12. 2 ?4 x' c( I5 w+ t5 O5 Y' T
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    " t: u! e( @) e+ l

  14.   U6 G% P% k" p/ ?( C
  15. bias tensor([-0.0152]) tensor([-0.2023])
    / D2 V, |; w% q- j7 t
  16. 1 ]% c) H, r! P* ?9 M: m  M
  17. tensor(0.1529)# Q1 K% T! N4 X  D4 ]6 [

  18. / E: D. S6 D8 l& e) n& J  |1 H! u
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])
    : P, h, {8 R7 a- d

  20. . Y  J4 `6 f% a8 Z2 a1 z/ _
  21. bias tensor([0.0050]) tensor([0.0566])
    : c* Z% [- d8 B# l, c

  22. / Y8 ]8 p6 z, ]( A( k
  23. tensor(0.0963)7 I- l/ e2 D3 v- E8 Z& Z/ G' p

  24. 6 s( O5 t9 }+ K1 i) s& j  x. b2 K+ J
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])9 j  I$ T0 b- D$ g4 b/ O

  26. % Q- ^4 o  O6 h. R$ I
  27. bias tensor([-0.0006]) tensor([-0.0146])  n: c( G6 N4 A% s5 _) r; R

  28. 8 M$ ~2 @; {. {$ u  Q' z( B0 r
  29. tensor(0.0615)+ v! O, l' s, |1 v2 T0 W, x5 u8 V" o
  30. , n, M( ?+ |; ]
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])5 S6 n4 H$ `& K3 n7 T3 M' F) H

  32. " j. z# c+ V0 F8 M# S- z! ?7 }
  33. bias tensor([0.0008]) tensor([0.0048])/ I4 X' Z6 s. `/ i' W8 O

  34. , P; F% L2 U& @. L
  35. tensor(0.0394)) W" _+ v9 ]% U2 L( x

  36. , x2 k0 E( v/ `( P5 i# g- A
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])5 A* V$ D: y4 i* {% c7 o  S+ e

  38. + L; w0 W1 g! I& ?, ~5 s6 I
  39. bias tensor([0.0003]) tensor([-0.0006])" b# w6 W; C8 B. n4 ?5 S4 m

  40. 5 g% Q# V3 h; x5 f( t( H( E
  41. tensor(0.0252)
    " l' E: i! X9 q

  42. ) K7 O. S$ G% ?& p4 r
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])8 g0 l/ O" ]% `' h+ p- N/ h

  44. 0 X# [+ W4 Y# l% k
  45. bias tensor([0.0004]) tensor([0.0008])
    ; q0 u7 m% {& h$ P

  46. ; v3 B5 K2 a! I2 {
  47. tensor(0.0161)& S+ M, [  D/ e8 Q" p6 L! Y, w
  48. # |/ |5 T' V% m  B: {: |* C
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])7 t/ k9 A4 x  B+ U
  50. 3 M& M' G4 Y* u! @
  51. bias tensor([0.0003]) tensor([0.0003]): R" W+ [0 i+ Q1 f
  52. ' L" N9 }8 \6 B+ @1 K* @
  53. tensor(0.0103)
    - V* @+ p' m2 y5 Q6 m( W' T

  54. ' l& i0 t0 F9 q: U7 v; c( t
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])7 s7 Q0 j% H7 t1 o& g

  56. $ d7 k. i, ]; [: c, w. w% z
  57. bias tensor([0.0003]) tensor([0.0004])& }% z6 G, Y) i$ Y: P! d( s+ z

  58. 8 U' u! G0 ~! p3 n" I  u, S
  59. tensor(0.0066)& T) Q, B) E( K: ^* V
  60. 7 D* N8 X8 a5 q+ ]
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])
    , h* t( V$ D! L. g& _
  62. 2 |. j: i! k: V
  63. bias tensor([0.0003]) tensor([0.0003])
    * |+ |! Q6 P0 o/ g6 _

  64. 5 P1 _/ U. Q9 L3 z
  65. tensor(0.0042)) Q. x5 Q" J  `4 u$ P& y+ n% y

  66. $ U+ ~5 J+ c' w0 E1 _
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])9 `5 P0 F( N* C0 E

  68. 8 E6 G- F( I8 K- H
  69. bias tensor([0.0002]) tensor([0.0003])
    % N) O8 v  H6 o* e
  70. / `! ^2 p' r) ~9 Y
  71. tensor(0.0027)$ L; N# N/ L  q/ ?3 ^0 Q

  72. # m5 Y" M3 H  f$ k& x( p
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])
    - Y8 D5 l5 d8 b

  74. 3 \& h9 I# E9 C" |# m) U, s
  75. bias tensor([0.0002]) tensor([0.0002])
    1 i2 l+ o+ L- m7 u
  76. 4 N: x7 g+ r' X4 V) ^, |
  77. tensor(0.0017)/ ^9 W* D/ ~* O9 [8 y% v) @
  78. 5 \: N1 S5 a' d/ y
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]]); _4 }/ x  @8 S/ K" l7 q# e( \2 S5 _

  80. , H+ _  X* j. D) e; H) ?' n, c
  81. bias tensor([0.0002]) tensor([0.0002])
    1 E9 V- L; t/ Z/ w+ N

  82. / B9 R- L% m' c: t5 ~. }; h
  83. tensor(0.0011)3 p% ~4 u2 @5 A8 x  |) G+ f- U
  84. - B3 M% Q- K( t+ {/ ~
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])
    / C6 r0 f9 J2 `$ j
  86. 6 T2 a: P# b+ `' s7 _! I. d
  87. bias tensor([0.0001]) tensor([0.0002])' Z0 @, |/ G& ]8 S
  88. / G' B! k. Q  W: v
  89. tensor(0.0007)4 I4 Q9 g& R% k9 o3 y$ Z' a; P
  90. . j* Y! P9 Z, Y6 g) l2 E$ [( K( p
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])
    , O$ h0 M, L) n* [5 [

  92. 2 Q: o+ V8 C0 y0 v; o! h& {
  93. bias tensor([0.0001]) tensor([0.0002])7 H* o- j0 H3 p  g+ n3 G7 Z- ^* Y
  94. 1 e- G- r* o( s8 K2 c
  95. tensor(0.0005)
    % D( h3 U. w2 m% N6 M
  96. # H, q2 N9 R8 _5 R$ _2 \
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])/ |! t6 q7 H( _2 j0 n, m, d" v

  98. 1 _( ?+ l' v" a: ]) Z
  99. bias tensor([0.0001]) tensor([0.0001])
    ; ~4 K6 G& L3 O9 h

  100. " K9 s4 s$ O3 s( T0 \
  101. tensor(0.0003)
    4 R; v; h' C$ g' M4 p1 `% m0 u

  102. 8 ^; C& }- m+ M' h9 n4 e# ^* b
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])9 |" B% d& n: e, S3 n( M
  104. / T2 j7 T2 B$ n8 `# D# l- J9 h- F
  105. bias tensor([9.7973e-05]) tensor([0.0001])
    5 l5 x# c+ q) E& @5 U3 _
  106. 6 M! J( Z/ f" Y: w, j
  107. tensor(0.0002)4 ?% l  |' O1 ~+ E4 G9 t) r* f
  108. 7 Z/ f9 E3 e& j/ B: r. U" g
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])1 X2 b$ l4 m0 u3 H" U9 {
  110. 5 C+ q* N8 I% v8 a- G
  111. bias tensor([8.5674e-05]) tensor([0.0001])* x; d# F7 V$ t5 s2 V' P- ^3 [; D  v& \

  112. 8 P1 q4 j& Q/ w9 w6 ?/ V
  113. tensor(0.0001)" \1 Q4 K: u. q; `9 ^
  114. , X$ ^$ J1 i" s1 K. _  U! ~
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])
    & d* r# A+ b* K
  116. $ V& y4 {1 t  j4 ^; O0 S3 Z% L: l" Q
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])
    7 e3 q7 E& A" Z
  118. % d8 M3 Q6 v+ X/ L
  119. tensor(7.6120e-05)
复制代码

0 C' `: G) R3 @  u9 i+ Y2 c




欢迎光临 数学建模社区-数学中国 (http://www.madio.net/) Powered by Discuz! X2.5