QQ登录

只需要一步,快速开始

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

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

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

1194

主题

4

听众

2951

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么
! p" q( w3 I- O+ TSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。8 k' E) C7 Q4 Q
怎么理解梯度?  t/ D1 b( I) Z3 K
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。
( N# ?$ k$ i$ L% O* f6 E) {# b6 G$ m; j& H* Z
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
. G$ h  i8 I( |" R% [" B9 p$ b/ a  z' n
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
) g6 S0 h* k) p( n$ d
$ P1 ~! i8 j3 r$ G  s为什么引入SGD& k4 E0 f0 c. r) p% |
深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。
2 `1 J. R5 }& @7 s- Y& J  @" N! S+ Q0 A
怎么用SGD
  1. import torch, L& Z. M/ K; X' d# _0 F3 \' s

  2. 8 H& W9 p2 Z& j
  3. from torch import nn
    # |5 v5 w/ B+ ^
  4. ; G' q5 p7 ^: W6 `4 X
  5. from torch import optim$ f4 a8 o6 Y8 ]! w3 a% \2 ]( r
  6. & A6 b\" b- C4 [4 y: a/ n5 Z9 _% c0 V4 r
  7. + i6 c' a. g) m; {- s5 y
  8. & E& E% {/ }- h3 o; ?. Z: E  t+ n! r4 e
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)! z  y% v* \% ?6 I\" t7 ^% o
  10. 7 |+ f# V- h# ~1 i7 n- G
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
    * @7 X& ?8 S9 O3 i  ]7 m0 U. \) A

  12. 3 `1 O2 c+ A, P

  13. * V' ~) h\" m5 E  a/ y  q& c

  14. 6 F; }7 c. F5 S$ F
  15. model = nn.Linear(2, 1)
    7 x9 t: a. }8 Q0 H* V

  16. 8 q5 a% s# K0 `( j' j# O- u% a
  17.   j, Y7 k  x2 m4 I2 _: g7 X' ~
  18. ) U. A8 t) M6 c, H3 V0 v, V' N
  19. def train():
    4 ?& Z2 J( g- U# a# G) l5 A
  20. 9 ?4 c% v5 f; `: t
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)7 K\" m! O7 v7 g, x+ q; r

  22. ; T% }5 m' k& F! B/ S1 j1 ]+ f
  23.     for iter in range(20):
    . @& S* o) C7 U# t# B
  24. \" D9 g/ t; Y! c7 I* K
  25.         # 1) 消除之前的梯度(如果存在)2 f3 ]2 a5 F2 f3 s1 V4 u

  26. * q4 d2 {: p6 @' S2 d
  27.         opt.zero_grad(), P. `5 u! z, ]1 _2 J

  28.   S9 c9 E: e3 A4 a! Y% R& ]+ L: k
  29. 4 F6 G( n, h- Y3 y. i1 {% W8 ]

  30. 5 p- _9 \. B; N, h8 q\" p
  31.         # 2) 预测6 A4 [; Z1 M/ y
  32. . p' g3 `+ @& [# C
  33.         pred = model(data)1 F5 p7 i0 d$ d( X/ _% ?9 F( j

  34. 0 d5 ^# ~$ X1 z( P7 k

  35. ' @& y0 d3 k1 `8 k
  36. / V( C* n' G6 z, T' `  K5 D' F) u
  37.         # 3) 计算损失
      h9 ]% {+ g3 f, M* u
  38. - {% Y  l  m: k6 J, e# d
  39.         loss = ((pred - target)**2).sum()! D8 n) Q( a1 k8 k; M  {' X
  40. 4 _. X3 r1 Y/ Q+ \5 K
  41. & E# I; U) p. |5 N* ~
  42. $ n$ }: K5 j, o3 X
  43.         # 4) 指出那些导致损失的参数(损失回传)$ h  `' v8 F' M2 B; X4 H
  44. . D) S\" h0 w7 l
  45.         loss.backward()+ ^! A: R1 U0 I, I! k7 p7 [9 |
  46. 6 L. y6 c4 I0 `3 {' O\" K  e+ h  H+ y
  47.     for name, param in model.named_parameters():+ |4 y! m$ _8 D7 e5 g
  48. , r* F7 _( e! e, `
  49.             print(name, param.data, param.grad)
    4 c% [; n0 S4 s: ^6 w- c
  50. $ b- ~/ f, O2 R8 R, b. q+ r) _
  51.         # 5) 更新参数
    8 V: u$ M& }& A# _; ]: x
  52. ! [, E. [5 P+ g\" \, M! L* Q; n
  53.         opt.step()
    ) I! b# {, C; ?6 H: t5 t: o

  54. 1 R0 d& Z% x! b3 ]9 v  U' m
  55. , W+ u( u8 }: Q& h, `
  56. 9 b8 c1 U! }) I- e, Q) q9 p
  57.         # 6) 打印进程
    4 s- c# S# z9 H$ Y\" B
  58. # Y/ |1 G. ~* ~; ~8 A
  59.         print(loss.data)) I* ^4 ]# x: O$ W4 |2 s7 ^

  60. 7 c1 G$ P) }\" ?# h9 Y( I7 O8 o0 s
  61. - K4 g4 M1 i2 W% }; z

  62. \" {+ S3 Q/ T1 c( K
  63. if __name__ == "__main__":
    9 e5 @5 a8 P9 y' ~+ k& {0 j1 F
  64. 7 J- u0 t0 Z, Y. ^2 g\" ?2 P& y
  65.     train()5 j& Q( F6 M( @  V
  66. 5 K. d1 U8 a  V
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
' B  B) i" J6 l
; G; L0 {* T- Q: N# m* n计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])
    * h8 v8 I4 V; |3 Q5 T# Z
  2. ! {2 c1 P! o; ~' F, m& y/ X7 i
  3. bias tensor([-0.2108]) tensor([-2.6971])9 Q5 \7 l3 w* _1 {

  4. ( a2 k! b( F5 U( p8 X
  5. tensor(0.8531)& j0 C+ p1 s  n  l/ z* t9 B
  6.   q  W4 L% j+ n; _' @
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])! a. n2 `7 x4 j3 |. f

  8. - ]$ n/ d, w1 M4 m! o
  9. bias tensor([0.0589]) tensor([0.7416])
    3 `1 U4 P7 z% c- R

  10. # e9 l! ?( z5 j
  11. tensor(0.2712)
    2 U* o' q) Q! {8 M% ^0 F( N

  12. - N/ b; s$ r0 V\" F: [, _
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])4 L) @# W- H  t5 V8 E+ ^4 \

  14. ; |  S! i* k/ C- g& O* l, Y) N* V
  15. bias tensor([-0.0152]) tensor([-0.2023])  F9 O/ o: C+ S8 @/ ~

  16. - u. b; Y3 }; X1 j
  17. tensor(0.1529)* u: i+ N( @+ }/ k! h, E6 h5 V, Y% F
  18. : J  E+ P1 S2 b; n6 l$ j
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])6 m* ^9 i# F\" y) d0 A

  20. 2 [. s) i; g$ F$ @/ B
  21. bias tensor([0.0050]) tensor([0.0566])
    4 U, v! O& [' O' Y6 P0 y

  22. 3 U( v6 B' |  S0 x& _. P
  23. tensor(0.0963)
    7 h, j5 U( @4 }

  24. ! r! m2 K3 m& @% d% \3 f# {- L6 z& [
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])& f8 N  D4 w4 l; ^' b
  26. 1 Q* V, n7 @3 Z# k\" y$ {8 H
  27. bias tensor([-0.0006]) tensor([-0.0146])
    ! v7 q3 T0 V$ O& z

  28. ; U8 D1 M: j4 G\" y5 y0 U# f
  29. tensor(0.0615)
    $ ?* x# c; X# [- Q- }9 A: T

  30. 8 ?. H- V! Q& E' E& i\" @
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])6 o6 a; j' L  m- }  y- |
  32. ; @3 R7 N2 o  J# Z% A) f. b
  33. bias tensor([0.0008]) tensor([0.0048])! t/ b* Y: l% Y
  34. 8 E) _/ r1 ]! K/ g) B
  35. tensor(0.0394): {, Z0 Y, P' t: s, }. r* o

  36. \" j3 b- {\" o8 m3 F$ H5 s. {
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])
    - X6 x3 Y) f% _
  38. : F& T( G6 ]8 o$ S( d+ n4 I
  39. bias tensor([0.0003]) tensor([-0.0006])
    2 U/ B6 C5 Y, |- D& I+ \

  40. , @( A! L- c+ H  _9 w* ]& R
  41. tensor(0.0252)
    2 e& ~+ U  L% }# F
  42. $ o$ h5 _# d- `! `4 T6 Z2 G0 e# t- j
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])
    * [: Y/ p5 f& C* t/ L

  44. : W6 `- ~- I/ {% r4 B
  45. bias tensor([0.0004]) tensor([0.0008])
    . q- W1 P; {/ T$ b  ]( k9 O
  46. , G/ w8 L( O) K! G/ Q
  47. tensor(0.0161)
    7 E# w  s  |* c0 s6 Y) R% v- }' z

  48. * V/ d$ N8 \- j( ]0 T
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])' S' p\" v$ R+ B! o0 y
  50. / Y9 i3 j* |9 `; E/ X! `7 D7 q1 U+ ~
  51. bias tensor([0.0003]) tensor([0.0003])- V& R2 M0 S8 W\" o

  52. + Q8 J; U. V+ r' M6 e
  53. tensor(0.0103)0 b# l) ]9 T& e$ [+ _8 h\" e
  54. \" v9 O! v4 X3 v/ P9 f0 k- U
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])0 R! a$ X: D- b! \2 |* @
  56. , _! t1 ?\" ?9 f+ B) A6 i
  57. bias tensor([0.0003]) tensor([0.0004])5 ]2 Y6 ~9 Z- Q\" x; v' T2 l0 ^% g

  58. 8 U* G\" N6 s% n/ {9 p
  59. tensor(0.0066)3 o' w  ]6 [/ u4 p: _( F) t* n- w
  60. % @& M) Z3 l4 p6 l  F5 I- b( d$ B* W
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]]). y5 N/ [4 S8 V0 ^4 R2 Q% u

  62. % y7 I; p. h6 c1 t# N; d: _' {& N
  63. bias tensor([0.0003]) tensor([0.0003])
    & R% ?0 |8 K% |7 V6 w
  64. # K  ?+ ~. p' j3 D( D7 P
  65. tensor(0.0042)
    5 O) B9 O\" l* x+ Y$ Z% ?

  66. & O5 Y: Q& J9 z! {$ J
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])
    $ K. t- P3 `8 W# F# G

  68. / B+ Y# s$ J, |\" ~
  69. bias tensor([0.0002]) tensor([0.0003])1 f0 O9 w% }1 o9 U: }! z

  70. ; X! x9 i+ e5 a2 [4 I* F/ _6 s0 m
  71. tensor(0.0027)
    ' a\" \' B0 w3 W6 S* Q- c. i$ t
  72. ' z2 f0 ]1 Y4 O) e0 X. c& d) _
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])
    ) L/ X& w: H( p$ F. L\" p( O

  74. 0 Z+ c' X* m: o$ l* p
  75. bias tensor([0.0002]) tensor([0.0002])) N* ^/ B: |7 Q! F$ \6 a/ F( `5 b
  76. ( n& G1 ]/ w! M+ e1 `' q
  77. tensor(0.0017)& z8 D! v7 b+ Q2 d
  78. / }6 v9 i% x) B\" G& M
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]]). a; D( f/ O3 O1 e5 F

  80. * K/ P, l7 V7 k% m2 y
  81. bias tensor([0.0002]) tensor([0.0002])
      s3 d/ D, h3 ]( Q# m3 K

  82. : L% }8 H: J! G  r' \# _) {' K
  83. tensor(0.0011)\" H) i7 v9 i3 F9 ~\" w7 m5 P  T

  84. ; D2 m0 ^- u9 X7 {1 n( A
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])- @$ ~$ N\" o! \6 }% ]

  86. 8 J+ G9 c- P; t2 ?. d5 `
  87. bias tensor([0.0001]) tensor([0.0002])/ Q$ I; }& _) S+ o, P% u\" l& g

  88. 9 D( i: |3 Y6 l9 F6 B% X4 f. c7 T
  89. tensor(0.0007)
    4 O. e' L\" j; s. F/ V

  90. : _+ w% ^* l- ?3 y
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]]), b0 w\" J8 N' c

  92. ; h9 N: t2 E9 g2 T# T0 d5 C\" A
  93. bias tensor([0.0001]) tensor([0.0002])& g, n\" A% W2 Y- w9 V% `
  94. 1 r. Z4 U5 u; |  S) R
  95. tensor(0.0005)
    # U2 k4 d9 U8 [0 W

  96. + W1 o7 i\" b# k$ ~0 w  F
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])
    9 r- \$ f3 d2 h7 C) [' P- f# l

  98. & D5 K, S8 N  w5 H
  99. bias tensor([0.0001]) tensor([0.0001])
    8 y7 Y6 {, C9 x$ Q+ |

  100. # {1 a7 q5 X; ^6 O$ h+ ?) R
  101. tensor(0.0003). n, J3 }5 [# C4 {1 h2 p
  102. - t! Q0 p( d. g5 n2 Y9 {& g  I
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])3 T1 V* m$ w; o2 s
  104. \" C; E7 H- S6 k
  105. bias tensor([9.7973e-05]) tensor([0.0001])# P: G4 A* n5 \$ v5 p7 [# l6 Z4 p
  106. \" n0 f; i9 ^0 |, c7 t4 Z, ?
  107. tensor(0.0002)
    $ B( O' L3 |0 P8 B7 m

  108. 4 [+ s9 i4 R9 W2 d7 C% V
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])6 o2 X7 }1 S# Y% R) w/ @

  110. * ^! ~1 c0 o% L: [
  111. bias tensor([8.5674e-05]) tensor([0.0001])8 i& u0 W! K( P# [2 `

  112. 7 |5 p) g1 }  B7 {5 J
  113. tensor(0.0001)
    ' I4 W# Q, i, {' e4 i
  114. ) g. `, ?: C( q
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])8 i& {$ e; {+ y+ h1 _5 ~

  116. $ l* ]/ o: ~  v; W. G, e
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])1 S3 W/ g6 t/ Q- b& F1 X5 x

  118. ) b( L* K$ Z# u
  119. tensor(7.6120e-05)
复制代码

5 R8 p" D0 r. v
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 14:24 , Processed in 0.355584 second(s), 51 queries .

回顶部