QQ登录

只需要一步,快速开始

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

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

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

1189

主题

4

听众

2934

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |正序浏览
|招呼Ta 关注Ta
SGD是什么
- c, l1 t7 |' Z. r% g; aSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。$ V9 X, g8 {" s3 K% F7 `! z. [; x
怎么理解梯度?
; M* x+ A0 r; W& H$ D/ u) L; ~假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。  }+ X! ~8 |, g* `* C
- g9 t% o* `0 {( {+ B8 R
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。% m2 L9 }  A- x' J0 Q

! F& M: d0 T' C' @: L7 v9 B% E每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。3 d3 L+ Z$ D  O' @! W

7 F4 Y/ F) L( n) x1 B# h* L" \为什么引入SGD
  h  i3 F& U. i) z深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。5 `1 O" W- _2 x/ n

3 a; p9 f% z+ P' u0 b9 D- E怎么用SGD
  1. import torch# X3 I7 C) [+ f- Y- Z# o

  2. + s- o; e  T- o) a( g6 X\" C
  3. from torch import nn
    * h1 {6 @. X* Q( j; z

  4. % \8 Q! ?6 F& u6 g  g
  5. from torch import optim
    7 S9 n8 D\" |. I8 J) R

  6. 4 ]& c. M/ I0 I; l: ~' ]/ E* P
  7. ; y: D8 O& S/ D5 Y  B4 b( h
  8. ; A1 }3 w! M3 k. h
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    , k  H) s* B) S& Z! ^# k
  10. ' D1 ^: Z* d5 p9 p\" ?
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
      b3 X0 \; {4 N) X) V
  12. 5 c: u1 ?  l! {: L2 G7 ?6 @

  13. . W5 g$ |. c. S( q, k

  14. # R( q/ A6 M2 Y5 a6 {
  15. model = nn.Linear(2, 1)
    7 o9 `8 Q% K5 t3 j

  16. & B, t( ]) o+ |8 |1 k9 A% ?
  17. 5 P' l3 j1 `\" ^- L  ?

  18. 6 m, G/ D& B) j+ r6 w
  19. def train():: P' P7 L7 C: c& q* z

  20. ' I$ |7 b; z- }9 E1 j5 W4 M
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)
    * J/ L5 f: p6 a\" J! [8 L4 B; q

  22. ; m& T- P  l7 q3 M. U7 r* M. H' W
  23.     for iter in range(20):
    & e: l5 @# v! l: h% G4 y- e
  24. $ R. p: R0 F& ]  ?7 B
  25.         # 1) 消除之前的梯度(如果存在)2 x& s' h) S2 I\" H* ~8 z

  26. : U0 S$ @& R: ~& t6 |0 n1 d/ V
  27.         opt.zero_grad()7 K* X5 K% B& }# x+ \6 c

  28. 8 A: K  t$ p9 z5 Y% d& R  @

  29. - ?9 Q* X3 q4 o3 ^6 ]7 U\" k- V- R

  30. : ~) Q5 Q- J! l8 h4 _  K& e
  31.         # 2) 预测
    0 i0 e- @8 S! _0 M8 g, @3 G

  32. 8 t' N) p! `2 l7 F2 j9 ~
  33.         pred = model(data)
    2 f6 o$ x' X3 V5 k
  34.   n6 W) D9 I0 S* t0 C' }; \2 Z& C
  35. 1 a* q  \  `# s6 N# |

  36. 2 X( P( W5 \; y% t' V
  37.         # 3) 计算损失* U\" P* A) N5 Y! q+ R& C
  38. / g0 j% R! [; _1 W
  39.         loss = ((pred - target)**2).sum()& y) O/ K0 }1 g6 {9 `4 R5 c6 K

  40. - J5 @5 F$ R' k2 p

  41. 7 f. _- t1 c' I5 I2 E

  42. / B- _5 I3 j. O, |
  43.         # 4) 指出那些导致损失的参数(损失回传)2 ^& _\" g+ s3 Z! C1 b! g

  44. 4 y. `4 @( h( Z, \; f3 V0 q3 O7 E
  45.         loss.backward()
    & }% q. L: \' R# K2 I' W* V- @, r

  46. ! ^) R4 `1 l\" F6 c- _
  47.     for name, param in model.named_parameters():
    6 B5 q\" m/ Y\" P, d) X
  48. \" [# \7 ?/ Q& \# R
  49.             print(name, param.data, param.grad)
    1 a. a8 _) |5 S4 D

  50. 6 [; Q! \6 k5 U$ m: C# i
  51.         # 5) 更新参数' q3 W9 B5 `7 {7 W3 l
  52. 5 P( E4 v, j: L1 O9 k\" O& a9 v  |
  53.         opt.step()% H  S7 K1 ^; ]# c

  54. 1 e, B0 C) }' e. J\" w: e

  55. 0 V4 W5 {% R- l

  56. , b) P  o; d* v) f2 A* [
  57.         # 6) 打印进程* k5 ?) b: J1 x5 a3 r4 a9 P
  58. / E$ L; q% ~+ @+ D2 H6 [- i
  59.         print(loss.data)8 }( L5 K& R: s, A# v

  60. ' q; s  }+ `3 {8 @1 L
  61. ! @4 y# E( v- O. q' A& x\" ?

  62. + G4 x: K  d\" j7 L\" i+ Q- n
  63. if __name__ == "__main__":
    9 Z# c- e7 p9 o
  64. 7 `* _) |2 E% j1 V) r5 X  G/ F3 f
  65.     train()8 z5 N, j6 i% K: r

  66. 7 ^# a1 M3 M0 Z+ U8 B& L
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
- v5 M( }9 {) \  I2 _) h
$ f7 d; l6 I7 w计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])\" [\" G, r! U5 v3 _2 N+ o
  2. # b! U1 U) Z# R( t' C& ^$ k
  3. bias tensor([-0.2108]) tensor([-2.6971])2 G0 w3 p, M6 W5 ]
  4. + y\" q8 m- j. I: v
  5. tensor(0.8531)
    ) B: B' G* E6 W( O1 {- _

  6. ) q1 [/ j; b. b3 o( }! O! A% e
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])
    ) c8 F9 Z1 r: z1 F( r* {+ X9 j: ]
  8. 6 g  l* Y8 c& d1 H' D& m/ ?
  9. bias tensor([0.0589]) tensor([0.7416])
    2 T! X  Z# G% w1 t
  10. ; p$ b  M* z3 J6 j. q: L
  11. tensor(0.2712)3 U5 o1 m7 N5 T) p7 I& B, w- b
  12. 9 O5 ]4 g) d$ E# g( D6 Y
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]]): J3 L9 {, v; e! ?) b
  14. 0 s! g+ z2 Q0 W
  15. bias tensor([-0.0152]) tensor([-0.2023])
    . q4 N' |. r* n/ _! i( ^* E: g
  16. , u9 a1 H$ A7 B3 F, L% A
  17. tensor(0.1529)& {# o' ~( v8 d, V# h6 v
  18. 8 W) v, }) |0 I9 m0 U
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])4 A1 x) Y- z  G7 ?% ?3 v
  20. \" F* G. t& J. x: g! M5 x4 K
  21. bias tensor([0.0050]) tensor([0.0566])
    9 e4 ?7 L\" D- k0 v

  22. $ F; Q# y! {4 S( Z
  23. tensor(0.0963)
    ' T: K# S4 b6 R

  24. % K4 x% ?4 J8 \$ X/ F: c
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])0 t7 J# ^7 r2 Y* n3 Y
  26. 0 T$ L$ M: @8 F7 _  F
  27. bias tensor([-0.0006]) tensor([-0.0146])
    8 F2 e! {0 @; g% V- V& l# s

  28. ' F\" F4 k8 F7 g
  29. tensor(0.0615)/ N0 [4 {; \# w; ]2 v

  30. $ A2 I. e+ J' _4 k# S* w- @
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])
    $ j8 [) X9 s  W+ ]6 [8 t
  32. ( j2 c\" {0 ?9 F$ L7 v# I% w8 B
  33. bias tensor([0.0008]) tensor([0.0048])
    . o/ }$ m0 [5 H7 X& N' L, i
  34. ( a. }( g+ U  e  S
  35. tensor(0.0394)
    . Q5 M4 q* N& I1 S\" n% Y7 u
  36. & {  p5 p# O- @$ v0 t$ q
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])
    + K5 x2 y4 @1 U7 m9 Y, X  D- j
  38. 6 `: c- {# {( E' R, Q' \0 C
  39. bias tensor([0.0003]) tensor([-0.0006])
    . V3 s. Q2 N$ C* \  d( e4 P

  40. 5 a! W! T7 P) W\" c# k+ {
  41. tensor(0.0252)4 ~/ ?: T$ B' s/ Z* J( ?
  42. ( D, \% U* C; Q& Z
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])
    \" w% B7 A2 B. l

  44. 0 g' x\" z$ \: O) k- k  l; g: \
  45. bias tensor([0.0004]) tensor([0.0008])
    ! ^8 J  h+ ~$ W

  46.   Z; {( D& h2 u1 j% M9 U
  47. tensor(0.0161)+ j2 F$ d\" ?# r: L, X& l9 W9 F\" ]# F
  48. / p5 _8 k0 U0 Y8 ~- \# }0 R\" s
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]]). z8 p2 E0 \( ]\" {  U
  50. & k- ^6 R1 Q  q8 @: }0 ]
  51. bias tensor([0.0003]) tensor([0.0003])  D! v8 Y& Q2 W

  52. 7 p\" Z0 W/ A  R
  53. tensor(0.0103)
    , e9 E9 q5 z  A6 U9 E+ \& k
  54. . S. O, H# S) ]
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]]); _, j) S\" k5 [8 _( ]2 z\" W
  56.   j+ G2 g( R$ ^) o
  57. bias tensor([0.0003]) tensor([0.0004]); _0 y7 o  {; B4 k3 x
  58. 8 |2 M& C+ ~$ [# X* y) V\" D
  59. tensor(0.0066)
    / X# Y8 U! A% i7 w: W; Y  m. M- P

  60. + x& b8 J* ?! s$ V
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])* E, O) a+ d/ k7 f- C\" }
  62. 3 p4 [2 Z5 c& G
  63. bias tensor([0.0003]) tensor([0.0003])
    ' @- \6 u0 T2 |: J( n& C

  64. 4 [* V# Y' n% ]: S8 s
  65. tensor(0.0042)6 V$ Y0 a7 f: ?- o\" G5 F! c
  66. 3 W* `. Z) P, a\" y0 X
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])+ J  c& v+ T0 X( C* z0 G4 h& [

  68. ( H0 \/ l! o: I\" F( [+ J' J4 B
  69. bias tensor([0.0002]) tensor([0.0003])\" p\" R4 U' M) a4 \, h6 J5 _( U
  70. 5 q/ n; V7 q$ Y3 \0 j# p* F
  71. tensor(0.0027)% g7 T! p4 _( k- z

  72. : E3 R+ O9 r# o7 Q& u; @, [
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])+ F$ t5 _  U7 e/ G$ g

  74. ( {+ U% ]4 i/ H5 {
  75. bias tensor([0.0002]) tensor([0.0002])
    - s9 y4 T0 N' i1 X$ t4 o1 F% H- P

  76. \" C1 Z8 X6 @! t# X6 Q3 O
  77. tensor(0.0017)
    5 n1 u2 M* Z. f
  78. 4 `7 c5 l5 a2 h5 Z, L; g
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])  i1 z0 {3 k/ ~1 e; \0 A. f3 L5 E
  80. + P# l) C7 ?\" T
  81. bias tensor([0.0002]) tensor([0.0002])2 U; M0 Q: X7 V( @; ]* r# b: v5 o; _

  82. + U\" |\" W$ p4 H9 a! z7 `5 r! R
  83. tensor(0.0011)% ^- ~: `# v6 Z* F

  84. 1 k8 C0 F1 {* e% l' `7 i
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])4 q1 m: {3 {3 ^! _
  86. \" I; F4 S7 u4 c8 x, t2 i2 g
  87. bias tensor([0.0001]) tensor([0.0002])$ g8 U/ o7 [: X! \+ E

  88. 0 p; f& C; ^0 m5 E! X
  89. tensor(0.0007); b- ]& l% }7 w9 m\" g+ m, a' O

  90. $ T  R/ }8 M; g0 l( @; t
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])
    1 m  F' Y% F9 [4 F( N

  92. ! z7 F+ H3 U9 v8 ~& Q, x! o2 c% I- w
  93. bias tensor([0.0001]) tensor([0.0002])\" e7 Y& X! F  g
  94. ; z# Z& D, g8 f* \2 a( Z0 ?
  95. tensor(0.0005)
    % V\" x+ m2 n\" |4 w* s
  96. 2 b' R% a. D; D1 r
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])- S/ K\" b$ a; _: e\" Y- W' o
  98. * R2 O. \% P& F( }
  99. bias tensor([0.0001]) tensor([0.0001])
    1 m) O1 R\" G$ v
  100. ) \$ J; k8 z7 e8 c1 O
  101. tensor(0.0003)
    9 {0 S% P! v; X- m8 B+ W
  102. 0 _* m) d( h, E2 n/ J6 n
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])1 ]5 g. c. ]4 [# K, N/ ?; u

  104. 7 P! s: o% Q) B5 G4 _
  105. bias tensor([9.7973e-05]) tensor([0.0001])0 w, p2 I0 }: S9 m

  106. ' ]. J: l: e( q  \4 g
  107. tensor(0.0002)
    ' u# `' h1 P9 T- @
  108.   I$ I7 |2 G6 f) v4 L2 z
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])
    2 B0 f7 y: v6 C/ U8 s5 e# C3 ]

  110. 9 _, B. G' q4 `& ?; K4 s/ P6 p
  111. bias tensor([8.5674e-05]) tensor([0.0001])9 B. a# b3 P: S+ O
  112. & K2 i0 e+ O9 _
  113. tensor(0.0001)  V9 m0 |5 t3 s- @/ M: s4 e

  114. - p+ a2 a8 z5 Z* S+ g8 V
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])
    3 A4 M+ p# _8 {0 A& Q1 a2 H: C
  116. ( Z9 ^7 U8 B0 W* i- K5 Z* P! p
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])$ O- ]! n! ^9 K% _0 ]1 u, R
  118. + M0 j- F: W0 M
  119. tensor(7.6120e-05)
复制代码
# O8 u2 P% o* h! {/ x% ]* L
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-7-24 04:07 , Processed in 0.323635 second(s), 51 queries .

回顶部