QQ登录

只需要一步,快速开始

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

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

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

1192

主题

4

听众

2946

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么
) }: \+ `+ h& b( u7 gSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。
$ l4 N# \% t) Q4 a6 `7 g8 f% U5 L怎么理解梯度?* t) a% t0 Y/ V, P6 t1 U0 ~- w
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。$ k' a1 L' m: L  _+ `  w: p# j

! A2 `7 {1 g) _# a+ ]1 p在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
8 s& z7 k8 Z8 P: ~& x" r, A/ }7 p. ?7 K( r: ~
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
+ I& S7 ^& D- y5 [' X( f* l
( E  {$ J9 w. {" s: l7 N为什么引入SGD1 }! |9 k; P6 b. Z; ]
深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。% N' ~& \& E; V! A, @1 }' S

1 V, N* A7 T0 e$ ~& N怎么用SGD
  1. import torch
    - N6 F0 W9 c  g. m* y
  2. . {8 d  O, A7 i- K/ C  w
  3. from torch import nn+ u# G\" g, w\" ]# D+ h

  4. . c/ Q* U2 v  k+ ~. E1 R
  5. from torch import optim/ `, W9 D6 c% _3 H3 c, ^* U( p
  6. 3 G9 ?1 |8 V3 u4 _

  7. $ k8 n: a  y$ U/ a0 |* Y

  8. * n$ I9 M' K: M0 G! h0 F
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    \" r. S1 `$ O1 O( {8 _, D2 e
  10. 0 o# k4 p# p$ g3 X0 v
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
    % v# V) R, @; G3 J

  12. ( s. q& r! l' b0 M; T
  13. 5 O2 o' @3 _- n# J, {8 ]7 i  r$ {
  14. \" k4 W7 h1 p! F$ P7 t, {
  15. model = nn.Linear(2, 1)
    & m2 c. C7 m7 z
  16. / Q& \7 |  _; x* U( U

  17. ) u. A3 b. Y9 ]- ~9 Y/ a* z
  18. 5 A\" c0 S* K* c4 n
  19. def train():
    0 J6 @( [1 P& h5 `; G

  20. + F8 s' L2 ^# \. l0 x8 ^& h3 v7 R) B
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)2 ~/ _$ @% C  Y\" R
  22. . c8 n8 M2 P8 h
  23.     for iter in range(20):) k' Y# G& D0 ~, \! ]( L( O

  24. \" n& Y; f' Y/ a# q: K6 r# b$ O
  25.         # 1) 消除之前的梯度(如果存在)
    6 @( E  a6 M' q9 s0 {

  26. % [+ D4 B8 B- m( r
  27.         opt.zero_grad()' ^1 a9 p6 \) t

  28. ( d0 |8 e; `& I! }5 A. R0 |+ A1 p
  29. / [6 [' K; U* H& o- P
  30. 4 t2 Y% f! K8 Y6 j0 g% r
  31.         # 2) 预测/ J1 G( {' i# {8 V

  32. 8 e) N8 X8 [, Q9 u5 g
  33.         pred = model(data): V5 H7 A, C! j# H- F

  34. / b: F7 ]- E) p! X1 V\" x' D

  35. : ^2 M4 k4 Q- z3 A  i0 g

  36. , ?* L: s' N6 n- ^
  37.         # 3) 计算损失
    ) i  ^; q; y: u2 T; }) E/ m0 Z
  38. 2 r  u# p6 S, L0 T
  39.         loss = ((pred - target)**2).sum()- L7 H) y$ e8 P! N/ ~& i' e4 y

  40. ( n& V4 v) i% |0 o# A- J, T( J4 G$ ~0 N
  41. 0 G% x+ ]' e- f2 p  c+ z2 w
  42. 5 Q7 ~$ p* @+ a( K$ k; y) e
  43.         # 4) 指出那些导致损失的参数(损失回传)# _* \. J4 `: \7 Q1 _

  44. 2 K, e; v$ [# M
  45.         loss.backward()7 e- F& m  C7 t\" J# I, T
  46. , w* i4 x; t- o, ?
  47.     for name, param in model.named_parameters():3 ?\" t6 e* _+ k! u5 J; A; `/ s
  48.   g5 o. V( L2 F) S
  49.             print(name, param.data, param.grad)
    3 `% \) \; ^* ^
  50. ( v  B+ p( r  F0 Q# K/ @* t+ s' [) `
  51.         # 5) 更新参数
    9 ]. k$ F8 \! S' [  Q
  52.   z# v& u. S; M0 V! R2 a& c
  53.         opt.step()2 b6 [; ^% G7 w; ^/ H

  54. 9 O; W9 ]. O( `* d
  55. ' R. r# I: E; K/ O' D8 u. P

  56. + k9 T% R/ e\" |0 q
  57.         # 6) 打印进程
    / w$ p  }4 y: m8 K0 V) @

  58. 2 _' O  ]7 u/ `: l- X4 e
  59.         print(loss.data)+ y\" C! q\" p) E* _
  60. 7 K: U\" A& X\" [

  61. 0 Y. T& x$ {$ f. q
  62. 3 r7 x9 v: b0 s/ U7 I
  63. if __name__ == "__main__":6 L6 z- f# l3 b; `
  64. 4 H  ?% J( S# ]3 Q# }6 t
  65.     train()
    $ k, s* J3 b( ~7 I\" H9 d* r) L/ s

  66. 0 ?9 n+ \4 a( Z  D& Y
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
+ s% P) m/ @# j$ @4 @6 p) E& Z( q* G3 P* S' P. N: C
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])
    3 t( g$ g. v: N

  2. ) |6 |9 }1 G7 E5 v- L
  3. bias tensor([-0.2108]) tensor([-2.6971])
    - t3 E  y0 l: P: T9 w% R

  4. 8 ^7 M1 h7 Q! _; I1 f: C0 ]* D6 s
  5. tensor(0.8531)  ]( C7 k; i( l* ^$ i  J% ^

  6. ! n* e& r  I6 C
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]]); F* p2 h; w: k) [8 I( N0 g
  8. ) G- q2 m/ N. s- M7 w2 g# z! e
  9. bias tensor([0.0589]) tensor([0.7416])
    1 H% k* n) T% r. k0 S* Z

  10. \" \& n9 d* O1 o/ Z6 S* |
  11. tensor(0.2712)
    7 t. R9 _1 G- o, p& M6 f

  12. 7 v: a7 L& x' ?6 @7 g: r
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    + k! |* _; R! f- ]: }

  14. 8 ~9 r- u$ Z) q6 E) P
  15. bias tensor([-0.0152]) tensor([-0.2023])
    7 Y* p. H( Q- M; ?- @8 D3 ?
  16. # v5 S( ]% P( p6 x5 Z' V. v3 q
  17. tensor(0.1529)
    # h0 Y2 P& |& h1 u3 r1 u
  18. 6 s' A$ E. q/ X3 \# [- R- E
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]]); _) g: J9 z+ J& v7 k

  20. , [2 f, w0 d, v  ]8 Z
  21. bias tensor([0.0050]) tensor([0.0566])
      F5 Y9 F: C7 G* x6 k# T
  22. % h8 m, N/ l( @2 s& W
  23. tensor(0.0963)
    5 h, W0 W9 O2 y3 K# T

  24. 2 o8 G1 r; R& o+ {
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])5 ?9 i& u+ \* O( C/ V- j+ t

  26. , q( \' ]' Q1 M8 N# I
  27. bias tensor([-0.0006]) tensor([-0.0146])
    ! l2 l0 L' ?* w# R5 ~
  28. + M! j4 h, P( ~2 t& b: v0 M4 u2 P
  29. tensor(0.0615)( y$ Q- e: X6 j/ U  A7 ?: m

  30. 4 f\" V4 ]! w% M2 z1 f) R, P
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])% t8 H1 z+ v: @/ S. h2 }

  32. 2 X& g( F0 o' e/ t$ ^# [\" a
  33. bias tensor([0.0008]) tensor([0.0048])) r3 C7 v3 k6 g5 I

  34. / t3 c7 x; u: U$ y. i
  35. tensor(0.0394)) d5 o4 ?0 y$ O4 Y) @

  36. \" u+ t! F# u\" b3 _! n3 y. K, u8 ?! ^
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])
    / [7 d! y+ ?, {* t, K

  38. 9 Q# b\" m! z$ {5 m
  39. bias tensor([0.0003]) tensor([-0.0006])
    0 n) S3 `$ p! \) t# ?  n7 b
  40. & T; S2 V6 c- E. G; [/ }
  41. tensor(0.0252)
      O9 ~9 x5 m5 S: K4 a: ~3 k

  42. , L# j# g; K2 A0 ?. I% c  t, z; H/ }
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])! O9 Y. n& x0 n5 `$ O

  44. / h\" C$ Y2 b. E1 s7 A
  45. bias tensor([0.0004]) tensor([0.0008])
    , X6 i, K. u; Q# P
  46. ; \1 q# i7 H- d( U
  47. tensor(0.0161)
    % \0 o; T0 v. f8 j( w& ~/ ^

  48. ! n0 g0 D2 U0 f! e- g; V* \
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])$ h: ?+ V. Y6 n( t% n
  50. 8 X\" q- W/ U8 [! P8 }) W6 F0 N7 `
  51. bias tensor([0.0003]) tensor([0.0003]); i/ B( o( L0 X! a7 o

  52. $ H  S$ D2 O8 w* _$ Q: r8 ^
  53. tensor(0.0103)6 c5 _# t$ p0 N, r- z# z. M; w

  54. ' S! \- q9 _2 d/ d
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])
    8 I& z% p! ^* u; Y

  56. % h6 p) p* B  U+ m/ i/ \8 u
  57. bias tensor([0.0003]) tensor([0.0004])3 ~6 Y6 P  K$ ^8 N* r

  58. * O; g) B) ?; x! f
  59. tensor(0.0066)
    ' e# H! Z1 l1 H+ _

  60. ! I* D5 A3 L; H  {, X. I% g
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])
    8 K. p' ~6 |. i8 w; r$ D8 s$ b
  62. 9 J9 V. X: q' A- \
  63. bias tensor([0.0003]) tensor([0.0003])1 e/ T2 l3 T4 h$ g- T$ H  J
  64. 3 z& z\" |1 L% a% F. q
  65. tensor(0.0042)
    7 {& @. v' V: p# [* G

  66. , g. k4 O/ _! j# ~
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])& D( y0 ^& T1 Z

  68. + y+ n( A2 K8 J; @4 i
  69. bias tensor([0.0002]) tensor([0.0003]), G) u\" M0 ]0 E; ^

  70. 9 j$ p+ E8 g3 G0 ~7 ^, h. U
  71. tensor(0.0027)* x; x' _+ _: l' g5 x
  72. \" Q( S% y- H\" b; n& g0 @0 y
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])
      z* q\" }, |2 h' |

  74. 1 q3 t' i; v6 ?1 \) N$ {, I7 O% e* [
  75. bias tensor([0.0002]) tensor([0.0002])
    $ Y* v; c7 N6 H  h* O: o) }

  76. 7 p! f* {1 X, C: P; o/ J
  77. tensor(0.0017)
    & c2 w$ u! g) _3 j- `  K/ k
  78. & e) i/ P0 u1 m+ o
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])' [2 u. d( C' U4 H& s  @. b# x% i

  80. 8 L; @# y0 o# {4 T
  81. bias tensor([0.0002]) tensor([0.0002])
    * C! O! ~3 _$ C4 y8 k' u
  82. \" F  q. z0 x/ ]
  83. tensor(0.0011)
    : F3 \; J2 s+ X' ~
  84. ( W' Q9 J( O# V4 w& U: `* g. a
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]]): L( d( |3 `( h  I9 Z( w
  86. 2 t5 ?0 w$ H8 m9 P( D4 V% V0 k
  87. bias tensor([0.0001]) tensor([0.0002]), X- \7 K: ?* g( h
  88. # g+ d! b7 J& v6 f
  89. tensor(0.0007)9 c. f$ n# p6 r\" @# w( f
  90. 4 _7 Q4 x& N0 B. X* y% Z: Y& d
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])
    % n- N' Z; J. w7 b; w0 W: C  O

  92. 7 d/ d( T8 I1 \
  93. bias tensor([0.0001]) tensor([0.0002]): p# [+ o; i  _' ^

  94. : S* }' X8 W5 d, ]
  95. tensor(0.0005)$ W; C+ G$ \% n8 d6 p
  96. : [8 m# d\" d' A1 A  f7 l$ O
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])8 {5 e9 s& c) {7 X& U8 Z, Q

  98. 7 C- y\" V  R. D0 }! m& L  ]! x8 b9 {
  99. bias tensor([0.0001]) tensor([0.0001])
    ) F: R3 `  @4 G/ V: d9 }& v
  100. 7 Q. @3 m! a9 ?
  101. tensor(0.0003)4 m  R. g* v: `7 }5 g9 ]
  102. $ L/ O: c* l; b% ]\" q2 p3 h/ ]
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])
    , C, p+ |) d\" E; E3 r\" }
  104. ; f2 ]( ]# R3 n4 s4 z
  105. bias tensor([9.7973e-05]) tensor([0.0001])3 o( Y' {4 Y$ t7 y+ r

  106. $ ^3 R. l/ c& @
  107. tensor(0.0002)% C7 i4 _0 Y2 d
  108. / I# _( n* o( n5 m4 j  {4 F
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])
    ! K, K  A' X5 e# B  X& J+ Y

  110. * a5 j, y$ r8 c9 U- A1 x) @
  111. bias tensor([8.5674e-05]) tensor([0.0001])\" m) P. v) K! d! B) t
  112. - [7 K4 M/ v' \  S3 t' y
  113. tensor(0.0001)/ ^2 J5 k/ k  n, @3 U; X0 ]
  114. + i, ^$ O6 }' D. N( I3 g% q
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])$ x! s4 C# k/ E8 a+ e

  116. - T8 ]$ C8 G9 J: Q, i% a
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])
    0 e  U% @2 m) z# j* s5 E

  118. $ F( `9 s\" I# \
  119. tensor(7.6120e-05)
复制代码
) W* j# O, o" C# P1 s. M
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 00:39 , Processed in 0.516295 second(s), 51 queries .

回顶部