QQ登录

只需要一步,快速开始

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

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

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

1194

主题

4

听众

2951

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么4 k1 e) O( S) p1 |5 ^
SGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。
  K+ x; |) n0 V7 K  e' d7 q2 s8 b怎么理解梯度?6 @) @# S8 r  T! p, A: u4 Q8 A
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。) W# p- t$ K$ R& p
  V& O+ y% C) @: _( s
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
' S; D4 s7 u7 u) Z3 }& o; x- r$ c; _5 |0 q. }
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
* V7 H+ h5 y* K: \% C( L
, Y5 n8 p$ h$ j% L. z4 ~; L为什么引入SGD
2 v, v: m/ a5 f- E! }深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。
3 J9 @( t* d& w! `1 _# n" S( `0 L2 ?# n* j4 ?
怎么用SGD
  1. import torch- g* j8 d% E4 R! U

  2. . e1 N  e) X! V* M+ [
  3. from torch import nn
    1 X0 q# D  [, G

  4. ) o0 r8 ?% k7 e. j) K3 Z7 L\" W  c
  5. from torch import optim- G! B- Y5 k0 q% N( C
  6. 5 v/ K$ c; i7 v: }* O- P* t
  7. ) {5 a# F2 _2 C' J/ l* ]' o: T

  8. 4 h# [) x9 T; u
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)6 A; V6 u0 O# \8 y/ p
  10. . p+ W! ^: y\" M
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)9 t4 D; L1 X8 y/ `+ v% f

  12. / \8 ~3 [\" p- y- G\" ~

  13. \" c( H8 L- Y2 V# I, o
  14. : m# g/ r( l. f5 x1 }* Z
  15. model = nn.Linear(2, 1). y) \6 x( G1 _- u

  16. : w6 `9 {4 ]$ l2 \. a; n

  17. & M/ q! z3 I+ y* Z

  18. ! k6 A. [, T% Q8 t
  19. def train():
    ) ]% a1 l  v1 d  O! h' m. \' [& p
  20. . k' m6 m3 d4 M: R9 q
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)( x$ U, h( x+ X

  22. 2 R  X' R) L' J0 J
  23.     for iter in range(20):
    \" p* s2 U% s! L6 P' ^& c3 s

  24. 8 b7 S3 `; N4 r. J/ z5 N3 m5 t
  25.         # 1) 消除之前的梯度(如果存在)
    $ i+ S5 S1 q% x1 B( [
  26. 2 q; ^- E+ K- J8 y
  27.         opt.zero_grad()
    . S- \* ]% l. j\" @
  28. 5 D; e5 u. ~7 h8 d  t; M1 ]1 s. N: Y
  29. . s! d7 V- Z\" ]' b9 c
  30. 8 B6 I$ H5 X+ J2 e& ^+ C6 D* x, @1 t
  31.         # 2) 预测
    , g( Q- u& U! ~3 ?; n. a: {
  32. 8 M8 E6 y: Z; |! P
  33.         pred = model(data)
    $ B2 D9 n1 u! W* _

  34. $ C\" U( D( w; O+ X
  35. 1 b1 E. |& {) H/ B1 D

  36. 6 m3 N* }+ ?5 Q2 h6 M4 w1 h
  37.         # 3) 计算损失7 W6 b2 ~' p  j) Y8 p; W3 z
  38. * L. Y) x4 N# b) K4 L8 T
  39.         loss = ((pred - target)**2).sum()3 E/ z' w# ^1 P( t- K3 y

  40. * }- V2 @, C0 A% p3 E0 E
  41. * D6 s: }8 o+ a1 o% a0 ?
  42. 5 o8 {2 c0 n1 o7 z: L4 F
  43.         # 4) 指出那些导致损失的参数(损失回传)
    / M  A# c8 {% I: R; Y
  44. 0 }* F9 p+ W  t\" E/ U. W4 X
  45.         loss.backward()
    ' J. k2 s  c  ~4 i  Q; h\" S

  46. & L7 n. @  p; t- V
  47.     for name, param in model.named_parameters():
    0 G& f  w# x( S. e$ e9 H3 |/ M1 x) }

  48. & n+ s9 Y- o! p+ I\" A\" U
  49.             print(name, param.data, param.grad)
    + A2 p, f9 \8 S, d; j3 E% T
  50. 9 D  ]7 G! c2 l+ N& z  v& r
  51.         # 5) 更新参数
    / ]! @3 f, p) W& w

  52. 5 e4 ^+ b\" ]- N, S1 q7 D
  53.         opt.step()  W' X) I9 g1 \; _# L& o+ q5 [9 C

  54. ' M: b9 @. I, C\" a

  55. % u: r7 G/ w5 T

  56. 9 @+ I' q' ~$ n& R  S+ ?
  57.         # 6) 打印进程7 |; }0 z3 S0 h' Y3 [( A! p7 u
  58. 8 J9 i# [* l1 l2 {4 f
  59.         print(loss.data)
    - \4 u4 ]$ H. F. |1 R0 c$ b& f3 F

  60. 8 S1 a4 \- P& I
  61. # P9 {# A; v! B! ?2 L: n  }# l1 u% `

  62. / ^! ?' i5 T8 X* o7 N: H
  63. if __name__ == "__main__":( ^) ?% @9 L# T: K! x

  64. $ Q  _, \4 ?4 @+ T5 e' I
  65.     train(). {* t# Y% a, M2 Z\" h0 q* B3 e
  66. 5 A! M9 J' y  a! }9 t
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
: ^( V4 f% e) F: c8 F
, F- ], ?2 p& M& V计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])1 h/ U& t4 n' |; v: E9 b

  2. ( ?' Z+ G/ J\" g, T
  3. bias tensor([-0.2108]) tensor([-2.6971])
    / }+ W0 r, K3 U- G
  4. 8 Z! O4 n* f% g, c+ H
  5. tensor(0.8531)0 a- J/ A& {6 L  I

  6. 1 f+ h7 j. [9 P( c7 P% Q9 a6 B
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])2 ~4 x% T5 m* e
  8.   u* d4 g' q  C
  9. bias tensor([0.0589]) tensor([0.7416])
    0 N+ x2 K' z- J
  10. 6 S7 Y# |2 E0 A2 F' W
  11. tensor(0.2712)
    , y  a  M- \! c\" L* D
  12. 8 t7 A% R& B\" l+ s/ p6 |1 M
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])\" b* o+ b\" Q; ^\" {3 u\" x0 I4 S
  14. $ S. }9 r9 E) q5 q\" x& S+ g, @* y+ b5 r6 ^0 n
  15. bias tensor([-0.0152]) tensor([-0.2023])4 F5 w5 c3 J& d5 \5 r; N( z+ d

  16. 7 W/ R7 i1 F& K& ^+ p& k
  17. tensor(0.1529); D' f: B# f6 b$ ?$ v9 J& w
  18. ; s+ M0 q0 K, Q9 Y. b  N
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])
    - C3 y4 E  b) `
  20. * [; `' w* s6 R% C. L, ^6 z# e
  21. bias tensor([0.0050]) tensor([0.0566])
    2 E+ G  q% D8 d1 u& {9 E

  22. 5 N% k\" f& o+ R9 Z  t& W9 X
  23. tensor(0.0963)
    \" t0 J/ x( ^( |/ _\" {5 k& M, |
  24. 9 P4 O& M+ m) K5 @, x
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])9 s( M; \* `- G

  26. ' [- ]\" g6 x+ ?  Q% O! a: ~
  27. bias tensor([-0.0006]) tensor([-0.0146])
    0 J% ?2 _  G9 Q  ~# C4 J: e( R. {* J
  28. * n7 k5 y0 c# c' k! m; ~% C
  29. tensor(0.0615)6 t% O6 ~' u4 {, H
  30. \" h5 \) D5 z\" v8 v1 {) \
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])
    1 O' V# f5 U+ k3 t$ m% |# S* N; j
  32. * [. F1 m* l- r$ }' a
  33. bias tensor([0.0008]) tensor([0.0048])
    2 m2 h9 q& x) z6 K: q! G% B0 e

  34.   i4 v$ R! g0 |* J
  35. tensor(0.0394)
    * n9 L7 d# E, l' e/ a2 A# X8 r: Q5 f/ o9 h

  36. ' B  {6 p( f) u/ ?& l9 f
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])4 g! |& r+ n- |
  38. ' E  j\" d9 {4 }1 |. y4 E
  39. bias tensor([0.0003]) tensor([-0.0006])
    5 c9 h6 ^3 E- v  D; d3 a
  40. + g4 I6 c\" H2 x/ D2 p' l- h
  41. tensor(0.0252)
    ' p: e! h4 [; s  O8 h2 g

  42. : g, k& V, X1 c' h0 ]3 I7 {; q, V\" Y
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])
    # u6 f3 Q/ |1 Z

  44. 4 h! z* N* C) W% E
  45. bias tensor([0.0004]) tensor([0.0008]): q1 z' a- |1 y+ r5 z

  46. + m  X9 w3 B+ o
  47. tensor(0.0161)
      R& [% A7 ]+ `4 v. P' s

  48. 9 n, L) A- f- H$ D
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])& |, D6 t( Q9 u( _  M( q% A

  50. % u7 y8 `# I\" U( c  M7 u
  51. bias tensor([0.0003]) tensor([0.0003])
    - n- t  t6 f& y\" x- z) l: a: D) f* _

  52. ( o8 r5 |\" l, B; v5 ]/ A
  53. tensor(0.0103)* p4 o8 S3 G( f: o

  54. 0 E' d3 ]; ^$ T5 u
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])
    ; S9 p4 v9 r7 S. F; a% ~

  56. , [: E  @- w; S& s9 n8 d3 Z
  57. bias tensor([0.0003]) tensor([0.0004])( b: J1 p* G+ }0 [

  58. . @1 B5 h+ V. M) P# R9 m- I# O
  59. tensor(0.0066): g6 s% r0 Z0 S/ Y7 ~2 P, h5 M( L

  60. 1 W% F# P3 D. H$ R; l# \
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])
    ( s1 u' I8 A) l- m4 r

  62. , @$ j1 f8 z6 X# t! j8 Z) O+ Z; B
  63. bias tensor([0.0003]) tensor([0.0003])
    ( {2 z9 e% d. M
  64. 8 t/ X! d+ i* E4 `- i! j# \+ I, b
  65. tensor(0.0042)- R$ P, m9 H% ]) _8 d

  66. : O& f: M7 a. z, U6 a# R
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])/ K+ v9 R' i- q/ B

  68. $ a9 D7 d8 F: q( }- ^6 ^; P1 J
  69. bias tensor([0.0002]) tensor([0.0003])+ x& x1 f. w5 n& B
  70. ; Y7 M0 J/ C0 U* Q( j, G0 g
  71. tensor(0.0027)/ l' }1 [+ p- I# D% L0 _0 M9 ]

  72. \" J3 z* F, \+ e
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])2 `, k. a! X4 M2 m. n! L3 g

  74. 2 A/ c\" S# G, |/ R
  75. bias tensor([0.0002]) tensor([0.0002]), \. _# B\" R$ m\" D( @
  76. 3 P+ O' ~; a/ [3 S. J5 D% t
  77. tensor(0.0017)0 G# l% }  f; J6 l- N

  78. : D) S: p- `1 v# n  d
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])3 S) S$ L/ G$ P. x' U
  80. & v6 I4 `: |# w! K
  81. bias tensor([0.0002]) tensor([0.0002])  \; F, z3 _2 P

  82. & z/ s! g0 U# t
  83. tensor(0.0011)/ L% g8 v# X0 M, s' a7 y

  84. / c; ^6 K! {8 @! e7 r7 \) ?1 V
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])& u: `# X, l5 T6 n: r6 B( [- w

  86. $ ^' b, O3 x6 U. N+ Q9 r; O0 b
  87. bias tensor([0.0001]) tensor([0.0002])
    # U7 r* p5 U\" I7 z1 S, v

  88. & i: `1 I' |3 G6 \- ]7 s
  89. tensor(0.0007)) @3 }6 |\" U3 Q5 h7 Q- ?# h, Y# W2 _

  90. & P$ b& w\" j4 P( o1 L
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])! e- c- n- D8 V* [2 B

  92. ; {6 `+ Y6 ^) @9 a; O8 ~, c
  93. bias tensor([0.0001]) tensor([0.0002])
    ; w1 l. C; [% q2 b9 M5 A

  94. # R2 p; W, a/ U8 Q+ E% B
  95. tensor(0.0005)3 ~  [8 v) {2 }$ }

  96. - ]2 y\" A1 [- x) Q\" q
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])
      Y( b- r' ]( r: ]
  98. 0 j. f4 Y' g: R
  99. bias tensor([0.0001]) tensor([0.0001])
    . T  m# h+ o6 O

  100. + E9 O% e  l  T0 @7 J6 F5 b
  101. tensor(0.0003)
    , \% E( E, e  P9 d6 k% Y

  102. / N% M! U' `- T: z\" d6 j
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])
    ! v1 ?, y# Z5 c: e1 J) @% y

  104. % _8 P/ r/ \( a' J8 t8 x2 b6 G- a
  105. bias tensor([9.7973e-05]) tensor([0.0001]). t% \& v, _) I- w1 Q  U% V
  106. 8 ^7 W% J1 `; u
  107. tensor(0.0002)$ g7 |' Q( v9 Y, L# ^5 c
  108. 1 u. H4 _' ~\" s6 R8 r6 f- d
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])
    \" @0 b% O, k+ \/ d# @

  110. \" i& Q- S7 v8 ?6 H* N
  111. bias tensor([8.5674e-05]) tensor([0.0001])
    4 c: d0 l, f; s5 u9 x+ J9 [

  112. 7 N\" h& k1 h- W* U; S7 S1 d* s
  113. tensor(0.0001)
    $ G, r. ]; E+ |, ]) o0 g# H9 c6 w

  114. 9 j( `4 g8 W+ ^. x7 e% Q/ [
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])
    ) D2 R9 }/ m- m8 }1 ?
  116. 6 ^$ K4 q; ?: E
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])2 V9 |# Y! P1 e3 K' n/ n
  118. + O' X* \8 L( u+ b) N3 a  K\" e  l2 o
  119. tensor(7.6120e-05)
复制代码

& i" [3 L' I$ N/ I7 c# H$ ?' U
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 16:21 , Processed in 0.429680 second(s), 51 queries .

回顶部