QQ登录

只需要一步,快速开始

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

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

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

1194

主题

4

听众

2951

积分

该用户从未签到

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

% P& H8 ]# Y6 [# O  V怎么用SGD
  1. import torch+ A* ?6 `2 a: F

  2. 7 O; m1 e! j+ b+ F1 M9 n0 j
  3. from torch import nn
    8 Y4 G0 h! _\" @* K4 p+ p
  4. ! x% ]6 Y& N; r8 H
  5. from torch import optim
    5 ~: x* T! J7 s5 L$ ]; j

  6. 3 s. O7 L: L! N1 D
  7. : X4 Q: K. n2 t& A& e
  8. 5 ~! G- i* ]* R' U+ ~& P
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    2 @. n$ @4 q' t+ u\" K3 W& y: Z  |

  10. ; [* m- V0 \; q\" G
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
    ; `! |, M5 p2 D

  12. 4 L6 @2 V1 \0 t% x2 F
  13. ' q, E) r0 m5 J( c; I' t* t

  14. % {9 p9 ^3 {. D; g- V% v9 [# N* i
  15. model = nn.Linear(2, 1)
      U  T+ G( ~  z9 f) {
  16. 6 g$ ~- Q9 u/ C  \: z; t

  17. 6 h3 P9 _( K, i1 j8 s

  18. - B! K6 ?0 A* u& |. f
  19. def train():
      _% j' D2 h  z' r7 O* {

  20. 6 L) ]+ }& t6 S8 l& t  E: o
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)1 ]9 `5 I% E7 V. d
  22. 4 A. m. R/ D5 l; u2 w0 I2 q6 r
  23.     for iter in range(20):
    % E! q+ I; z5 u% F0 B, w\" [

  24. ) s/ e3 O% z: P/ s) d2 ?
  25.         # 1) 消除之前的梯度(如果存在)
      U6 g& ]& O+ @# ~2 Q3 Y/ H2 L; l, R6 ]
  26. # |  b+ Q+ F1 j. M! R
  27.         opt.zero_grad(), D$ v- h; b) M. `9 J
  28. - g# B9 X- n, \$ `( {& a

  29. # k3 v' N2 i; U( _2 j
  30. : `( x5 i* @' z$ A9 ]6 s
  31.         # 2) 预测
    3 v  E\" f) n) Z. _! M2 @/ I

  32. , I) l. [$ h% d8 K
  33.         pred = model(data)5 D! M\" h9 m# q2 O2 t6 O1 X
  34. 0 b( i  o8 z+ o\" g7 w
  35. 3 c# m& J& h- I, |# Q3 ~/ |% B2 w

  36. $ d) N) i/ q, ^# G
  37.         # 3) 计算损失; e\" W1 r; d) P- T

  38. : D, n, Q6 a8 W3 W9 S4 I
  39.         loss = ((pred - target)**2).sum()
    * t2 X4 X\" T, _$ j5 m# ?
  40. ; O2 j1 j0 s: \5 h
  41. ) m6 h1 |% y  Q( T1 {# @/ Q

  42. / E5 P; E, J, S3 ^% r( [: w' J$ _$ K& t
  43.         # 4) 指出那些导致损失的参数(损失回传)8 k) m- A' u, N; h
  44. 8 L2 @3 Q3 I, I\" I
  45.         loss.backward()
    1 F9 @, m; J& l* Z/ I8 W

  46. & ?; ~8 q. O$ `
  47.     for name, param in model.named_parameters():
    6 {; r5 h& h9 k( v; S
  48. . |2 Q, I1 g! T
  49.             print(name, param.data, param.grad)4 N$ a0 `- W8 t
  50. 4 b/ H$ T- G% ]
  51.         # 5) 更新参数
    - @\" m8 n  e2 D- p- }* h
  52. 8 }5 n2 q9 x- W$ k8 X+ k1 F0 O  I
  53.         opt.step()
    / \( ?) L7 x9 x

  54. \" e) W( X5 z7 r( }3 s

  55. % [  R. X4 a) ?: Y% B9 J: l6 m

  56. 9 H1 r6 L1 S: a0 q9 F  }
  57.         # 6) 打印进程9 H( k+ q* u4 o! i4 w# R
  58. \" F  |7 h6 Q  j% e
  59.         print(loss.data); d0 u# _2 z7 F! j3 ^# d9 \
  60. & j& q1 I/ S! `9 P

  61. ; B& G. u0 `\" R1 }5 ~

  62. / @1 n1 \3 L3 x+ j, l9 b9 _/ K# o
  63. if __name__ == "__main__":3 j- d. j6 s; ~: m
  64. \" M\" [- y; m\" N; U  k) O
  65.     train()
    5 h. i' x; K( t- y6 ~
  66. 8 q% f6 D- p3 {1 Z5 g
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
, v9 a( V/ L0 z2 F+ T. a, C& b5 X/ d. {1 `  n% q6 A
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])
    ; V3 M+ ?; a/ h/ t. I

  2. + Q9 W) w/ W7 r5 p9 m. I! \' L
  3. bias tensor([-0.2108]) tensor([-2.6971])
    7 K  t% I  A/ `* B, e# a

  4. 0 s2 l& D% p, n# {, O* [
  5. tensor(0.8531)' |, \8 J, z' X) f+ T) h
  6. \" b! ^! a# Q0 W3 W8 ?% X3 i$ k
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])/ a. t, @: |) k7 _9 t! e
  8. 4 E9 A/ W* v1 M7 s! J3 J2 ^! s
  9. bias tensor([0.0589]) tensor([0.7416])- R; K1 Z* }4 S
  10. 9 k; P  s! y9 T\" l! _0 H4 q
  11. tensor(0.2712)
    * S' Z4 j) x) y0 U( F8 \  s$ e

  12. ; {  I# z% _4 W. }, Q
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    0 q/ e; A% q0 a# a\" t, \6 {
  14. - I3 j6 R+ O7 @+ D
  15. bias tensor([-0.0152]) tensor([-0.2023])0 P/ p- E\" d& ~# ~7 v

  16. # I$ I& ~# Z2 k+ w. h* p, @
  17. tensor(0.1529)3 O: X\" ^; N' b# p* ?; k7 l

  18. % Z) n& r: U7 L6 R2 z2 g7 {$ ?
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])2 g4 F9 \. o2 o0 b. ?
  20. % b' ^1 Y& X6 d7 t+ I
  21. bias tensor([0.0050]) tensor([0.0566])& X( {/ v0 T' t0 G
  22. ; [\" S: t5 _! l& _  l
  23. tensor(0.0963)
    9 k5 f  i9 {4 Y
  24. + Z# r0 i& S$ H  C$ Y! q$ V* Q% H
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])
    . z# D* L6 q0 V/ N8 j  J1 F+ G

  26. % j. K3 T8 I+ A$ ~; G  G- j( t
  27. bias tensor([-0.0006]) tensor([-0.0146])9 O! D\" G  l3 z$ H! g, U5 ]; g& @
  28. 0 Q! y& |# b; U! ^0 m- D
  29. tensor(0.0615)
    $ B9 a5 E1 j; T7 q1 A8 G) d2 Y/ @

  30. : V: d. p  O; c) n) a* j! j
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])/ N, f1 Q* e, h6 p! j

  32. ) ]5 D8 _/ q1 w. \
  33. bias tensor([0.0008]) tensor([0.0048])
    / Y! ?+ ]8 k6 u
  34. 2 k$ ~/ y8 x2 M; l
  35. tensor(0.0394)' u! _  Y/ i2 J& c

  36. , l! J- u9 z. n+ L1 l
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])
    7 ?& b  t; r4 Y$ J9 A, B* R

  38. . v; g+ }* a: Y1 _2 Q* i- h
  39. bias tensor([0.0003]) tensor([-0.0006])
    & W6 r) C8 R* J% @

  40. 7 {3 z/ F& u& S
  41. tensor(0.0252)
    , S. m- j8 b% q% j% Y
  42. ) g# Q4 ]; I6 q/ Z% M5 r
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]]); N$ ^! ?. ~2 Y! D
  44. - N+ N5 ^! G+ u: w0 F1 P. u7 ^' P3 i
  45. bias tensor([0.0004]) tensor([0.0008])6 `) e. j: U8 W# {9 a  H

  46. $ s! s- o2 s1 X' x1 l  Z
  47. tensor(0.0161)
    ; F/ B0 P: z) l# H
  48. ! Q1 e\" T; d- O' ^\" L
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])
    ' O: m$ Q' i; I/ i5 U% @
  50. : H8 J# ~/ J0 C) u+ J
  51. bias tensor([0.0003]) tensor([0.0003])
    & Z* `0 r5 {) r8 x) e1 A. ]7 a
  52. $ A$ c& Y& Z$ S
  53. tensor(0.0103)! C6 s2 ~\" h\" K  i& H  U7 \
  54. # g- G7 u5 P& t- J
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])
    8 [( }: X& P\" i' L# ^2 U9 t7 j
  56. : Q' l$ Y2 T9 V
  57. bias tensor([0.0003]) tensor([0.0004]): Y8 }& G; J1 Q$ s) o5 K; k
  58.   s2 V% W* m& @& F: O1 S7 d
  59. tensor(0.0066)* W, ^\" G\" \' C: S
  60. 7 v9 Z% S, p3 @/ A8 F  q, F+ c1 V( e
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])
    - `$ Z) {* @1 [; o3 |' Y

  62. 7 R, G2 `, ]: p2 W/ X
  63. bias tensor([0.0003]) tensor([0.0003]). s; P& ]0 @\" g' ?, V) Y- J8 p# J0 B

  64. ( a& F% r* ]5 R- ^0 I: c0 V2 X
  65. tensor(0.0042)
    / @% O3 a7 y\" F, h9 u& I# i
  66. : l. `# \/ D: J) r- M
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])8 S; t# c3 M9 Q5 t

  68. , I/ M! {& Q6 T9 m  l
  69. bias tensor([0.0002]) tensor([0.0003])
    5 P9 R\" V4 T  h' h( g- Z& R
  70. 5 I; L% f6 r& @! {; f: u7 ^$ n
  71. tensor(0.0027)2 E; W5 {0 l0 F5 k$ d9 ?\" {
  72. $ n: ~6 `/ P& V( F  g: ?- y
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])
    ! `: P% Z5 F' }  E; f  [+ `
  74. 7 y( j8 q; u# c- e
  75. bias tensor([0.0002]) tensor([0.0002])& {4 j5 l4 I3 b! k\" k2 m# c2 m

  76. ' a- Z! a4 \/ W( R
  77. tensor(0.0017)7 d\" A2 ^8 A! u* @/ z% i1 A
  78. . x& I5 D/ P+ V( P
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])# ]. G. `. d$ G5 [2 o( @3 Y* _

  80. 2 G\" R0 {3 v- {, v' T
  81. bias tensor([0.0002]) tensor([0.0002])8 i\" G0 J; D9 |; A, s
  82. 9 U. T! v) Q& v2 p/ L
  83. tensor(0.0011)
    , V! h5 O: _% \* n- _& E/ K
  84. ( d% t  U5 O4 e\" E  V# A. X
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])
    % G: u% J( ^3 B

  86. 1 b2 @  H  V( \/ i  ?
  87. bias tensor([0.0001]) tensor([0.0002])
    & D% C: d( Q& H\" Q3 @0 e' B! n
  88. 6 r) N( z7 t$ `: ]& _$ _, I! p$ K
  89. tensor(0.0007)
    $ _8 i$ _* _  h6 F. I
  90. 6 L\" G\" o& @0 E' f0 |
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])\" c( \$ D& {3 Z/ Z* t$ G

  92. 4 J0 p& b3 k- i) \+ h9 [$ N
  93. bias tensor([0.0001]) tensor([0.0002])
    3 p\" [/ A6 ?7 T/ G( y
  94. , k1 @4 A& ?, Z4 U$ V
  95. tensor(0.0005)  A7 ]\" Y2 o# Y, k5 K/ [
  96. 9 U1 T! `) r2 {$ k& }
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])7 j8 B2 o: c9 \' Y* W( P1 D1 {& @
  98. 7 P4 J' p6 _4 N$ C1 q. ?
  99. bias tensor([0.0001]) tensor([0.0001])% j1 E1 j7 {) i4 v% K9 g1 j
  100. * X' V, `, V6 u7 B6 O3 f$ O
  101. tensor(0.0003)
    3 M2 ]' t0 \/ U8 L' |/ k: s

  102. ' k0 V4 P) q6 `' v\" R4 v
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])/ G) ~! J1 k& ^' Y# d

  104. 6 S* z* K% `. G
  105. bias tensor([9.7973e-05]) tensor([0.0001])
    # T7 l7 j7 Z5 S0 y  N& e) [
  106. - ~/ s4 r8 i* ~\" N: Q3 S
  107. tensor(0.0002)
    ) Q6 v- R1 F5 r* F% K! s
  108. * _* b& [# }  U7 \5 W8 M
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])
    ! g, g8 C# {; ?3 D: a0 _+ a

  110. ; v+ y; C6 P, E2 w
  111. bias tensor([8.5674e-05]) tensor([0.0001])
    4 C- A8 `6 g* D: o0 P8 }# `3 A
  112. \" w\" x1 k) r\" \9 V: }  }/ e) ]6 H
  113. tensor(0.0001)6 t6 y9 I9 G5 o8 e

  114. - R' A( F3 y. |
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])3 k) E9 q% R3 O8 a  h* ~6 v1 \# d

  116. 1 t8 ~8 d4 F; g; d- F2 s5 r. }\" i
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])
    2 U  N: A$ Z2 d& `7 c

  118. , ?1 A, ]; ^5 ]% M; o/ g: U! ~0 Q
  119. tensor(7.6120e-05)
复制代码

' Z; |$ ~8 z  E+ N
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:24 , Processed in 0.463391 second(s), 50 queries .

回顶部