QQ登录

只需要一步,快速开始

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

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

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

1192

主题

4

听众

2946

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么4 I0 L+ E7 a3 x- l. J4 S
SGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。
# F0 X& n* \+ ^2 ^8 L/ y; `  Q. Z7 r怎么理解梯度?
* ~/ o8 e1 e: k  ^2 Y" u1 k7 L假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。
1 J; R7 ?; L0 p+ q& q! P2 {7 B5 ^9 b+ G1 t' g3 a  x
在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。" x* U( m# Z9 V5 u/ `

( s' U; D: p8 S: |) f/ P/ H/ Q每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
1 R2 y/ V9 ]- J, j; \2 t& C& H) k; F3 L' M6 S2 i
为什么引入SGD
, U6 P; ~- Q: G0 s. d& _4 }深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。  W' i# V; U  A8 O; v" W, E

/ A1 v8 z) g' x" k怎么用SGD
  1. import torch, B# A3 H& X8 ~4 z6 b' {

  2. ; `  W8 D. x/ a$ h- l, d: k8 k
  3. from torch import nn
      s$ z1 J/ @& D3 l+ a. u$ q

  4. ; p$ L3 [7 h- X' D* c3 e9 C+ F# S' j
  5. from torch import optim
    8 j! r) E, k; |\" v; W
  6. 9 U6 Q: r; T# g# u& W
  7. 6 m1 K4 d( i& k* X- \9 g$ A

  8. ' u% x4 ?% t2 t: H1 q
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    7 P. C/ `3 L- z, t, W: H- J
  10. + Q$ ^5 r, S/ Z. W* C/ \
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)/ l4 y4 v. M6 V- h% u# y% C3 V3 K
  12. 7 q6 s) ?* m! m( @* B\" e6 z& k
  13.   Z5 f0 d% ]$ e

  14. # E3 M  b/ _5 `) ~: J! c* _) ]
  15. model = nn.Linear(2, 1)' N& S9 W2 Z. N) I: @9 M
  16. ) F  j% f\" \# F8 }6 T9 u
  17. # B# G7 `: J; l# d7 V3 ~0 E3 I
  18. ' C' A) w& S% D/ R
  19. def train():' B9 w5 D6 |7 n6 a% x1 K7 i

  20. : M: e) R/ n' A7 h\" I
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)
    5 g; a9 z( @5 _- S* L+ p/ E. c! b

  22. # b! ^; }1 H! r  ~2 d
  23.     for iter in range(20):: Q: R  q5 P% i  r

  24. , r: q& e4 t! V: w. Z
  25.         # 1) 消除之前的梯度(如果存在)
    ' z  G3 P: q% B8 [

  26. * |8 K) m. O\" H! a
  27.         opt.zero_grad()+ M$ p+ x# x  ]8 ?0 U, `

  28. 7 R, d; u  d8 {# S

  29. + k, d7 j# N- M! {

  30. 3 z$ s) u; o9 q  Q& D
  31.         # 2) 预测
    : N\" s  k  o1 e: ]3 c( M
  32. 6 q! B! E2 l6 |% T0 L
  33.         pred = model(data)
    ) [  h/ h3 s5 q% \# g
  34. * c! C1 i$ P+ C
  35. ! Q  c# Q% T+ E

  36. # r1 J; a& Y- w! f6 G8 t
  37.         # 3) 计算损失
    ; a% u6 P  }) s+ h1 A
  38. , Z) r, v& w# [' a
  39.         loss = ((pred - target)**2).sum()$ J) ]& F* G, H0 c
  40. + K  Q' @9 ~! A, }
  41. # k/ f9 I4 E2 o' }! Q5 [

  42. / L( \; e$ h6 {
  43.         # 4) 指出那些导致损失的参数(损失回传)$ ?% f8 l* X8 ^; {3 I
  44. ' c, z& W1 N) X1 |. J( I( v: Z# V
  45.         loss.backward()  f. ]# U& b9 Y* v+ q- j6 s

  46. # B! r+ R* n% ^2 h5 Q
  47.     for name, param in model.named_parameters():
    : t' S% o  X+ n. U; E
  48. / P* E; J# o$ S$ r/ z4 }) e
  49.             print(name, param.data, param.grad)
    * H: a- p/ c0 c) ^/ a- N
  50. - Z  }4 ?5 v6 A# z; U$ |
  51.         # 5) 更新参数
    : i; I3 f& T4 @0 O

  52. ! G: w! B, s) M4 z1 B0 `2 H
  53.         opt.step()
    ! Y. I6 A+ f4 G& R: Z0 R7 `

  54. % @2 n6 W/ d6 W9 Y

  55. ! o' b0 n! z: ?. m2 R! d  Q
  56. # I1 M. k+ S/ B6 ]/ @
  57.         # 6) 打印进程4 L! m+ U$ n% f$ I9 @2 J

  58. 5 _! H# G3 ~8 M' d8 q  k
  59.         print(loss.data)
    : V, {, ^5 V; c4 K0 ^8 x
  60.   w  c( h% j\" ~; K6 n2 j( @, w

  61. % z! U! b; D) ~; u, N3 I
  62. 0 P1 m% S* @! K: w0 J1 U4 k
  63. if __name__ == "__main__":
    9 ]* |. {( b* q/ N. q

  64. ' p! A1 A3 W7 K
  65.     train()! c/ p3 U. h/ X6 o\" U- M( h! ]7 D

  66. 3 y; O9 A5 ^. s+ C7 Y  m
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
1 ?. O3 z9 Z3 y: w: @
( }1 I9 P) j6 u: q. z/ V* x1 ?: V计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])\" _, z# |9 _2 w% r0 K+ T\" Q
  2. 1 ]7 f1 C+ Q0 M
  3. bias tensor([-0.2108]) tensor([-2.6971])  s# [9 K. D7 g& Q; l
  4. 3 ^) a9 |  q& J. i\" g. {* X
  5. tensor(0.8531)) u5 p! h0 ^! I( d! l& f6 Z5 d
  6. ! A9 L. f1 O, |4 ]; f. J' F
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]]), m! q' L; h9 l, k$ L

  8. ( G& L3 f& r- {9 _: D& K: z
  9. bias tensor([0.0589]) tensor([0.7416])
    ' d9 Z) I4 \8 W% O
  10. 5 m, s: Z\" g, y+ r* r7 f
  11. tensor(0.2712)
    2 R6 I0 M/ Q) Y; r
  12. ) k1 J$ p- |/ b  I! V+ f
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    + |! |% c8 K+ l1 f7 I

  14. 0 ^- Y3 f7 T5 Q) g
  15. bias tensor([-0.0152]) tensor([-0.2023])
    - e5 [6 p  A; l2 v, t
  16. # k& N. Q# k$ u6 S( n) K
  17. tensor(0.1529)
    . @9 V+ S$ l5 }* w  o: I

  18. - @0 F2 h: d% i
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])
    : c$ n. q  G# j8 B: y9 F3 ~1 o; i
  20. $ C5 E4 @6 W9 b1 C\" X
  21. bias tensor([0.0050]) tensor([0.0566])) M. f& ~) c% T: ]

  22. 1 m; C6 N; T6 \0 Q
  23. tensor(0.0963)
    : z$ q/ Z) K) i0 ]% T

  24. 8 [' h4 s+ r' H  T& L. r( i# j0 W' }
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])
    3 {* l8 m9 P' o) V9 u, m
  26. \" m) M- h. W8 l2 i% P
  27. bias tensor([-0.0006]) tensor([-0.0146])
    / G% [8 u\" V1 m! ^. m9 E
  28. ( O\" u2 G8 y' Q
  29. tensor(0.0615); x% W. A' ~4 C+ K

  30. / O3 R; ~# z3 y; w; M$ D9 @$ N
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])
    ( z: m5 G1 o, F# S

  32. , e6 U9 \# _! U; Y) {
  33. bias tensor([0.0008]) tensor([0.0048])
    2 m& C: X' g$ V# O. B! K+ X

  34. % A  |2 ?) ~4 l1 H. K/ x& r, l
  35. tensor(0.0394)
    6 G6 `\" H  {+ i  b* a2 b

  36. 9 K% H7 A2 W$ F; n
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])
    ) ]6 [3 w0 h- c- J/ f; a

  38. ; H3 A. o9 s. J# Z
  39. bias tensor([0.0003]) tensor([-0.0006])0 n2 e- S  b1 o& w
  40. \" u+ v: o* j: v- T. i5 w+ u/ w
  41. tensor(0.0252)& O  D$ |* U/ x* f( Q/ ]

  42. - {, G\" x3 F2 t3 m5 y7 m: V
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])5 J0 f. j5 ]% V* l
  44. 4 x# e0 j, m; C\" ]0 b
  45. bias tensor([0.0004]) tensor([0.0008])
    : {  Y, s* M* v' g; @

  46. 6 v& \3 s, q8 U4 l+ L
  47. tensor(0.0161)
    3 }\" |5 u% j& H% V\" _& I5 N

  48. 3 q\" [( v, Z: E; J0 S
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])+ z4 V, L; `* |! V
  50. $ n$ b% o4 x9 S# d0 T+ a9 B
  51. bias tensor([0.0003]) tensor([0.0003])
    5 a\" ]. x  [) `! i4 \
  52. 5 ^: b8 w  W9 P/ ^
  53. tensor(0.0103)4 O6 u2 N( ~- K

  54. 9 k) E& C% z3 R: |$ I8 ?
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])9 t! e+ X: z% y  Q4 C3 i  u3 Z5 c2 B

  56. 0 C! r5 _* G8 r( P  c
  57. bias tensor([0.0003]) tensor([0.0004])
    * L1 x3 w; r8 W( T\" r
  58. ! t. o$ z% G! [- K2 B/ L
  59. tensor(0.0066)
    . u: l/ J0 ^* H# j6 X# q+ }

  60. , ~# t: U' X$ D$ z1 O7 o
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])4 w, w) B. B% F- p
  62. ) B+ j# ]1 X3 q9 h\" T( C
  63. bias tensor([0.0003]) tensor([0.0003])
    6 k9 j' o2 _0 X( `& y( ^\" r, q0 Y

  64. 0 r0 U. [3 `\" h: |8 Z; P! m: E
  65. tensor(0.0042)
    ) M6 v+ t0 ?3 i

  66. / U* D, `  _9 F: n: m3 e# R( r
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])4 c! ?& p) `' |6 I) O
  68. & n5 [$ P( \2 D2 t* P- z6 R4 y
  69. bias tensor([0.0002]) tensor([0.0003])
    2 y% i5 [9 Y5 B# W
  70. & S4 x+ L( T\" `! d
  71. tensor(0.0027)
    , ~2 U6 }# z  V  q9 ^. x4 j
  72. ' f( b2 f( g4 `% D& V
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])+ D8 E  l0 h* {; D6 f% o2 ?% B  O
  74. + ?0 ^& @# V4 F+ s8 X
  75. bias tensor([0.0002]) tensor([0.0002])
      e# l! G3 H1 g, C6 t

  76. : P4 t5 `3 Q% N3 s8 x1 y! E
  77. tensor(0.0017)1 [! `2 N. [  x& i1 b3 J5 o

  78. & n9 K0 A! G, }) x& Z0 _
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])
    \" a1 n% ^\" H  h2 b

  80. 5 [0 q- ]8 `7 }1 B+ n! e& S
  81. bias tensor([0.0002]) tensor([0.0002])
    # w2 H: M7 u' H* Z

  82. 8 w: t; L/ u- S: x& i
  83. tensor(0.0011)
    + D7 M+ _/ P: j+ p& D7 H
  84. * [7 h1 \$ _; Z4 a
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])
    , `3 A; r. F2 D- ~: G$ o: g
  86. 5 {& _: y6 J& c2 k. j! t& \$ u$ M! J
  87. bias tensor([0.0001]) tensor([0.0002])( q) O; c5 D5 L2 H\" ^

  88. ! g6 l0 W3 v. s) T2 @
  89. tensor(0.0007)3 O/ w' q' @7 q

  90. / [4 r5 m\" o! r8 W
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])& `( A2 k# l# C3 \2 b

  92. . R7 g0 L3 B$ F- v: J5 S) d, X
  93. bias tensor([0.0001]) tensor([0.0002])1 R9 u( n; ?1 G\" @. B& y  v
  94. ; A- D5 f1 p! Y! X5 e
  95. tensor(0.0005). W$ T( X( D: m- f8 ]4 a. E

  96.   t, `) C' \; y3 h! g
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])  Z8 V7 M( f  n7 Q$ ~- U
  98. 8 ]2 u) a  G( p% t2 H
  99. bias tensor([0.0001]) tensor([0.0001])
    # z4 L2 X; y( }+ H, z8 E+ _

  100. + K/ I0 I. Z% z: [0 V\" {
  101. tensor(0.0003)
    % x6 F* N! P( i* T* N: V
  102. , T: i) T6 u1 P& z7 @
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])
    6 k2 \8 X/ b+ ?: I# y# }' @$ [

  104. # y% p* w: ^# q) h
  105. bias tensor([9.7973e-05]) tensor([0.0001])0 J) d- R0 I, b4 }5 z( k8 D0 O. _

  106. , G5 N( e5 l' D& w/ L& {5 u
  107. tensor(0.0002)
    # d; g& w7 O( n% s

  108. \" W; k/ ]7 ?3 _$ y% B0 @7 }
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])# j9 ^6 S2 F% H5 y' \. a

  110. 8 V\" S9 I; w) R* B3 N
  111. bias tensor([8.5674e-05]) tensor([0.0001])' e3 _% Z/ I2 _# v/ i1 f' K
  112. : Z# D2 L( c7 N  }2 z5 j
  113. tensor(0.0001)
    1 }8 G0 C4 M, a6 d
  114. ; w8 P- x% x\" d0 r7 n, ^% R
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])' S/ ]) c! p* C4 t

  116. 8 G: k* G% ^% Z$ O  L  t6 v) {
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])# E3 U  n2 o9 T8 ]1 H: S

  118. 2 a* u( f8 y: _6 z2 q4 x' W) R
  119. tensor(7.6120e-05)
复制代码
0 p& y* V& a* Q2 i/ u2 V+ z, f
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-8-26 05:21 , Processed in 2.125954 second(s), 51 queries .

回顶部