QQ登录

只需要一步,快速开始

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

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

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

1189

主题

4

听众

2934

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么
) y7 Z% b7 Q' sSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。
/ T# e2 w) z- `0 ^; r9 u+ T8 k怎么理解梯度?0 O4 k  R# w( L  ?! w2 _: n' b9 Q
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。- u! z( |" t) D
( N" w: I$ w9 x$ V. u( m# ~
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。& m' p- X  m6 H9 I2 q/ m8 p
9 T3 \' G) W- j% e! q3 T7 l
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。5 |! a3 p9 b9 r/ r3 E8 N$ U
0 x. r, E& ~0 |' F+ Z
为什么引入SGD( q6 a8 E  b: t( M' d6 A4 I
深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。% q& a" l: T' z  j, z

3 N) I) _2 Y# p* m3 K4 G8 s怎么用SGD
  1. import torch
    . a( ^  U, W8 }1 Z
  2. 7 _4 V9 ^7 q; G* p2 L: X
  3. from torch import nn7 T9 \3 G' q. s  Z' P

  4. : S6 ]8 T6 }: H2 @( K- A
  5. from torch import optim6 N3 Y. O8 \! L, r

  6. - n\" [# v) s0 Q- O1 w\" o
  7. * ~\" `1 _8 {! e1 a

  8. % g2 A3 q% s+ l
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)1 y& y% ^; r7 B4 w7 z
  10. % r- n0 n& N& W* @# e+ p
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)& _$ M$ K; j8 k. w; [1 }
  12. # A5 H/ b9 t\" O  D
  13. ! N8 W( i8 a* _+ a\" [
  14. ( A  x: ^5 g2 G3 d1 G! q
  15. model = nn.Linear(2, 1): H' L2 b1 ]& v' @+ z

  16. % H6 H4 q# j  w3 e\" Q
  17. , I1 x6 C: T2 W3 P/ z$ @- ]- ?

  18. & T$ O* g4 l5 H- @. s. b
  19. def train():0 ^5 i, u+ F1 C8 a1 |9 \' M+ N
  20. / g2 a8 M% P- ^5 Y/ k; {2 f
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)
    6 r+ g2 u2 ]8 C) ?8 n\" p/ V

  22. & w; L) P) r% K( N# U. [
  23.     for iter in range(20):7 K) f\" X- T# P1 t6 T3 G

  24. $ h2 b: o& T5 |# ^  o5 U- b5 c
  25.         # 1) 消除之前的梯度(如果存在)\" M5 X6 b' ~( r4 }; _6 X/ U9 ?
  26. , l! R4 }5 k! m4 \: m
  27.         opt.zero_grad()  [# }# L: L4 I9 W* a% e% F1 `* u
  28. , T) D$ {1 e. [0 w) S2 T6 s  L
  29. - y$ {% p9 ^0 b9 o/ z) O
  30. ) q  W) l2 F! N. P$ M' r0 A* u
  31.         # 2) 预测
    8 \3 [' C4 g. U8 h8 _# g
  32. 4 k. s% I$ ^; R/ i0 l
  33.         pred = model(data)2 D( m1 ?& r- U  C- q; L8 b
  34. 4 m# K8 j! w% J4 c, n
  35. 1 j& Q& B/ h2 w
  36. * `% N4 M5 Y+ ?\" K: h' E
  37.         # 3) 计算损失
    7 R% o7 ~. n2 l; u9 V

  38. 8 G4 I1 x( b6 X1 H$ G) X
  39.         loss = ((pred - target)**2).sum(): P6 a6 D+ b# F& A7 Y

  40. 9 c- M8 X5 y% K3 Z/ S
  41. - d6 }+ [  l$ p0 l- V8 J. V( }

  42. & f' }  ]' [3 K% C, y% L0 {! Z
  43.         # 4) 指出那些导致损失的参数(损失回传)' b) n\" I; e( y8 V; u( p% }
  44. , t: q$ I' p/ g9 H' \# _\" B
  45.         loss.backward()4 x/ g1 m& f/ g5 p

  46. ) y+ _/ T' i7 n+ D. x
  47.     for name, param in model.named_parameters():
    , f! R. g3 ]\" }\" m* i: e7 P3 X
  48. + D) ~# v, t5 Q! ?1 F: X0 Q1 D; q# a0 @
  49.             print(name, param.data, param.grad)
    ! W+ C4 q1 E; k
  50. 5 _  ^2 d5 J8 W6 m  F4 A: @* @
  51.         # 5) 更新参数
    + K; n$ {4 y2 G4 k* r  A
  52. . n# K& p' ^& a8 w# ]
  53.         opt.step(); G3 U. X+ C0 u8 c

  54. + o- J) K5 S- w/ F* P9 y
  55. $ a* K1 f\" p, d& G6 D* ~

  56. , K' W: w/ [+ L
  57.         # 6) 打印进程
    6 T( B' G4 g' ?1 W\" `$ w  p
  58. 7 N. F+ u2 F$ _1 }3 Y5 ~
  59.         print(loss.data)
    / ^# }# B/ x! ]  K* l
  60. 1 ^# T, h/ z5 F( T' H$ M  f, {
  61. ( N+ m4 T% c  K4 x# e/ [6 Y
  62. 4 G& H3 x\" c5 q( ?! ?
  63. if __name__ == "__main__":
    9 K, m: K+ L- j/ P: F1 S! Y
  64. 3 c6 r  M8 b) `- H: P2 x' N+ F* G
  65.     train(): i* h2 p4 Q; S* {
  66. 0 x( Y\" b; r7 `6 m
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。9 w) c! q+ w  H! Z/ Q+ O( b
4 ^. P/ Y4 L% R" c# j* t
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])
    9 U  e! e& x  r+ O2 M\" d; j, B. }
  2. 9 h8 U! x6 ^4 [. d* `
  3. bias tensor([-0.2108]) tensor([-2.6971])
    % C% q0 Q+ u! p. M, I
  4. * X: |. v( U9 X2 d, I% z: d( Y+ c
  5. tensor(0.8531)
    - A+ E: ^$ C) P: T- W$ I6 j
  6. $ L1 \& m4 k2 @1 ?3 G  d8 l$ O( ?8 j
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])
    ; n  p# M* U8 R- b/ f: b9 X0 M. y
  8. 4 _. R% V5 K+ I8 e7 p% m' q; g# K0 ?
  9. bias tensor([0.0589]) tensor([0.7416])
    $ B6 X; C- H$ B+ a% i

  10. . Y) s9 `; F/ n0 v# F9 M
  11. tensor(0.2712)
    7 {\" i( l% n7 D5 o8 n

  12. $ n0 S4 v; [+ I  Z' C/ {
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    % R\" `' ~$ _: s5 m$ D% x
  14. ! N3 g. M/ x3 q\" M
  15. bias tensor([-0.0152]) tensor([-0.2023])7 q3 A9 X7 M8 J0 `) g

  16. / e\" Q3 ^$ r4 W) t
  17. tensor(0.1529)
    ' A* I( `5 f, p8 n, h1 `
  18. \" Y  y4 a; o, U: l
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])5 m& e+ Y# r/ T$ O, R/ ?9 c7 M
  20. ' c9 _  d- g, {\" {% Y! H- B0 Z& E
  21. bias tensor([0.0050]) tensor([0.0566])
    # Z1 Z/ ^2 u  I4 H
  22. ; c! z\" k( M  P* z\" j7 ?. o' `
  23. tensor(0.0963)
      Y+ ^  W9 F. X% Y* i! ?( Q

  24. # F! A+ @3 z/ X, I- @* h( d3 C8 w2 }
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])* l8 r. ]; T& a. b: r8 \
  26. ! E) `4 g) `2 d4 X' |: Y
  27. bias tensor([-0.0006]) tensor([-0.0146])
    / b2 M( _& ?3 r
  28. ! z; U  F) \5 r9 O5 H6 T
  29. tensor(0.0615)4 \8 E; I& s- U! S, D
  30. 8 k5 G% {& j/ F5 g
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])/ c3 q- I\" k+ ^8 S. H  J

  32. 8 b+ J+ t/ J' {- Z
  33. bias tensor([0.0008]) tensor([0.0048])1 Z\" n5 n1 }( {6 d1 F3 A* k
  34. / M: v( d; f\" ^0 w7 O
  35. tensor(0.0394)
    ) b3 a4 r2 r4 t
  36. / A# b# \; D; N. Y( A0 A4 z
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]]), n) T% Z1 p% a1 A  ?\" Z+ `

  38. 7 c- L- M' B4 N, V, n+ e! F
  39. bias tensor([0.0003]) tensor([-0.0006])
    9 [) x+ v; l% t

  40. \" O1 x/ m, D1 g' Y
  41. tensor(0.0252)  u7 X+ f+ z8 E1 @3 w4 p
  42. - J) D$ _0 V* C0 J
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])# K  m  h6 X  m# K- L% J2 s, E  y\" a
  44. 8 c/ _- {! B\" r. U+ I; X
  45. bias tensor([0.0004]) tensor([0.0008])
    2 ]7 i  {( K. f' A; A0 [) N! t\" g% `
  46. 1 y3 K8 s1 H. H  \7 O7 Q
  47. tensor(0.0161)# p. a. N5 @& R9 q& M
  48. 2 l1 m) k% d- e7 w; X' K# M3 I' g\" L
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])
    3 G- K/ C7 m4 P! o
  50. , \9 M3 D  G( Q, G
  51. bias tensor([0.0003]) tensor([0.0003])\" q\" r. j0 C2 N+ ~3 b

  52.   Z$ Z( d9 d' |3 K: l+ s7 u( v! F
  53. tensor(0.0103)3 `4 O5 I' G& _- p. t: T' o  w

  54. $ w7 M4 }9 Y& j8 y' F7 `$ Z
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])
    ( u' J2 c  C  f8 }) _; x- Z
  56. ( R# k5 F\" z' R( |
  57. bias tensor([0.0003]) tensor([0.0004])+ Q1 a' x* _) h3 i( X% R
  58. / X2 K. H0 M8 b+ f: r0 N/ c\" U3 u
  59. tensor(0.0066)
    & L4 Y1 W2 {) f. Z- T
  60. ) Y& i0 }% o\" v6 r( l2 N$ z
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])7 L- x7 `, g; M2 C\" f

  62. ( I) v: b, K; N: U0 e) `7 h
  63. bias tensor([0.0003]) tensor([0.0003])2 r. ]' e: q7 d8 v0 A) Q0 S7 X9 i8 N; w
  64. * |3 R\" b$ c3 k5 r\" Z, \+ Z, a6 U2 h
  65. tensor(0.0042)  j6 X- g5 j. m# p# q; ^2 G

  66. & l# W4 a8 J( h+ F0 e$ @
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])
    % r0 c+ Y; m1 e: H
  68. 6 u/ e* x8 k0 u, p
  69. bias tensor([0.0002]) tensor([0.0003])& ^0 L\" C: [4 n

  70. # U4 }0 Z  n6 ]+ u/ {6 F0 f
  71. tensor(0.0027)0 ~3 _9 m& E4 B) S

  72. / @; U2 k6 ]+ d& h# S9 O+ }
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]]), Q  X: {+ G$ M\" ~\" Y* `, b

  74. . U' \; r8 q) |' z0 w6 L. A
  75. bias tensor([0.0002]) tensor([0.0002])& W% K! C  X  q, s/ K; N3 Q

  76. ! j/ O! ~1 c, l/ T( i+ G
  77. tensor(0.0017)9 [* o8 [' _* v' B0 k

  78. ) F3 d- _5 U4 t9 }; L/ V
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])1 Y+ p. x4 C7 S1 K
  80. 8 T5 U! ^5 F! V$ a$ I
  81. bias tensor([0.0002]) tensor([0.0002])
    - F2 z' Z. G1 b7 F9 C5 t; L0 x
  82. , P. M+ O: a7 x
  83. tensor(0.0011)
    . j2 d( \& `0 Y0 `3 X
  84. , A0 ]) W- ^& v+ j' ?
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])- m$ z8 M5 W; `' ?+ |\" l5 N0 \
  86. % ^4 H( f' G7 r0 Y2 _+ X
  87. bias tensor([0.0001]) tensor([0.0002])4 \7 G3 w7 f1 X4 T4 {

  88. ! F# M9 ~0 z6 P( k# m. q$ `9 M1 E
  89. tensor(0.0007)$ J% w) K, i' E

  90. & w& f) b\" Y2 V& Y
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])5 I$ D2 p. ?8 O; Z, \* Z; K
  92. - @+ f3 `- M( b! v
  93. bias tensor([0.0001]) tensor([0.0002])
    0 {7 O: M: Q4 r* V6 S6 c
  94. ' i9 i: [# x7 {6 H& ]. {$ x
  95. tensor(0.0005); }\" E1 W, \  v\" H0 C+ ]% d% G
  96. + a: a6 Q\" b' A, G( R% i4 m6 j
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])
    : v$ G' u3 D2 n7 Y\" K2 ]& g: V8 j\" n

  98. ' J- k/ M5 w6 _5 H( b/ I+ W& h4 C
  99. bias tensor([0.0001]) tensor([0.0001])
    3 y) @4 |! ?; d& V* a0 W- j

  100. 4 x) L- t5 h+ k* {3 Q
  101. tensor(0.0003)# H( L5 c' H\" d# U/ l8 [$ O% M; L
  102. 7 b9 M. o1 S: j9 p4 p
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])7 J+ H* L: X! k' I. T9 q
  104. ' U- ^4 P  C6 q/ P6 r( k
  105. bias tensor([9.7973e-05]) tensor([0.0001])+ W* x/ s& t/ B
  106. . |, C3 H4 i3 T; _0 z, L
  107. tensor(0.0002)# N2 m# @+ [1 J  P: u# @% l$ Y\" j
  108. 1 I& \8 d/ k' ~* Q$ M
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])2 [4 q9 |, j% ^2 d3 v6 o) p

  110. 3 b# J  b) D) f
  111. bias tensor([8.5674e-05]) tensor([0.0001])8 S$ i8 a1 J5 _9 R$ {9 c6 A1 v: l
  112. 5 o; v' {' O1 N  l5 \0 Q1 U\" g$ ~
  113. tensor(0.0001)5 G. [( V2 U# \( n

  114. 5 P- v2 }: f: s2 B% ?7 W5 p1 O% x6 ~
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])
    : d5 p; n$ R# I! x7 B% k2 U
  116. 5 e3 K\" i) b6 V4 J
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05]): y' L7 l6 B8 C0 U; Z' t
  118. + L% B& \$ l' M, K, P
  119. tensor(7.6120e-05)
复制代码
9 ?, j+ B" C& w7 X
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 02:41 , Processed in 0.838099 second(s), 51 queries .

回顶部