QQ登录

只需要一步,快速开始

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

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

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

1192

主题

4

听众

2946

积分

该用户从未签到

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

  4. % A5 z: ^) c: e% R7 M
  5. from torch import optim9 X2 D6 ?3 B) L

  6. / s1 A) o\" q5 y! t$ z' b

  7. 4 Q  o$ p5 k. R

  8. 7 E\" j% B) u\" e! a, ]9 K1 }
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True); x$ ^. g\" y- M- M) L9 P
  10. , B\" _2 T& e( g+ e# O
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)
    8 f: |4 M9 U7 r8 t

  12. 8 G- k6 L5 r3 W4 w5 @5 w0 O

  13. , A9 X' b! f) j; i: A
  14. % G& M/ M- Q9 k6 ~# M2 p4 m
  15. model = nn.Linear(2, 1)0 C$ `8 ]8 e4 m7 F/ y+ M\" s

  16. / Q  I9 y/ G, N# x

  17. % }; |+ f, \& `6 E

  18. + }: G; c# o. Z8 v. F. S% h
  19. def train():
    $ t' q+ v. M% B. v
  20. : X+ s3 S4 x4 V  s+ c/ }0 H
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)0 S7 Z8 z  L7 I) k5 v6 x
  22. 0 p* a3 R4 V; j$ c. t
  23.     for iter in range(20):
    ( ^4 A- {$ D5 s0 i
  24. 0 D# l! f) X! \  |+ {
  25.         # 1) 消除之前的梯度(如果存在)
    : `5 H- m% _. `5 K; h

  26.   A) [9 U- A$ o4 i+ @
  27.         opt.zero_grad(); R4 F9 O! w+ p& p\" W- A  j
  28. 4 u) a5 ?) P% Y- Q- K7 I% ?

  29. ) M$ {7 V3 W( p\" p' l. i* K  y3 E* @
  30. 0 U* C4 T6 F# N; k$ R! E0 y0 I
  31.         # 2) 预测
    0 B8 f  m7 s: n7 T

  32. 1 ^5 z1 Q& R$ d+ c\" V( z
  33.         pred = model(data)4 m3 X7 _* G% l

  34. 8 s0 Y' W& U* h9 w

  35. - x/ P# m3 m9 ]9 j; i7 R9 J% P

  36. . [8 @' F9 x  v
  37.         # 3) 计算损失
    % J  K. g$ F! }# \

  38. ( U) h5 b\" e0 J
  39.         loss = ((pred - target)**2).sum()3 K\" ^: M\" W/ O5 x
  40. 4 L  G8 g5 L) O2 Q: Z
  41. ! N5 `, g& H( {/ t

  42. 3 i7 I. {. _1 E; l% U2 ]
  43.         # 4) 指出那些导致损失的参数(损失回传)
    $ W% Q, s/ I% W4 P2 c
  44.   G* ^' B1 `+ C+ C. \% ^
  45.         loss.backward()1 y3 q1 A3 C$ m: U- |% x, i7 t( U
  46. 6 E7 v2 V6 p& X
  47.     for name, param in model.named_parameters():! \9 P3 y- t. c& I; F\" E& T, m

  48. 9 i1 Y/ D, p5 |9 z+ ~; d7 @: ?/ f
  49.             print(name, param.data, param.grad)
    5 q; Q3 d7 k, R/ K8 @2 V( A
  50. . M4 ?9 D0 X) |: a) J
  51.         # 5) 更新参数+ q$ t( i' ]+ F. y+ p. Z8 k6 S1 q

  52. ! g$ D, w0 R' f' M+ ~
  53.         opt.step(). _1 w# H$ ?8 L8 J. \! f7 W

  54. 5 i\" f1 ~5 i+ `
  55. 1 |) ?3 z! r, t9 E: T1 t! j/ x

  56. 5 n$ F8 Q6 h- r
  57.         # 6) 打印进程7 n' C$ H0 r. K

  58. 6 Z3 W: {, u( Y2 Q5 a% {
  59.         print(loss.data)% G0 `  O/ A- e\" l

  60. & h) b3 g9 }3 W8 {# j- I0 M
  61. 5 B1 ^6 `. y, w+ ~\" u2 }\" ]

  62. 2 D8 r: S, o! N& X3 Z# Q
  63. if __name__ == "__main__":
    0 x9 v  ~' b  Y( T' h. H- ]

  64. ! c# y# z5 Z$ X5 H
  65.     train()1 {1 E7 K9 x7 B& C\" n) W

  66. 7 g) T$ ^$ g0 j2 _4 @; R8 p1 W
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。& x8 ?: R! i" ]$ [" b
+ M! V4 }5 {- i, ]7 ~0 i
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])' C) L, t( ~% D0 e, y* A# _

  2. $ M* I4 Y+ \; b7 f7 I  x
  3. bias tensor([-0.2108]) tensor([-2.6971])1 r\" r& t, m  o% \/ l4 x
  4. 3 m' B+ [; I/ W3 V/ I5 C
  5. tensor(0.8531)
    5 |0 _3 Y% {5 G% U  T
  6. ' c/ V  |# ?; U1 r0 {  J
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])
    0 L/ [9 [\" p( v& }
  8. ; w) N& T0 s3 j9 Z6 h) q' U2 R
  9. bias tensor([0.0589]) tensor([0.7416])
    9 A0 A$ D6 ?3 y. [# N! G( o) u

  10. * |0 J' o) {% ~' t: E\" J
  11. tensor(0.2712)
    + v/ N5 V, r! H) g4 u, z

  12. $ Q7 F. W5 w0 K3 a( W2 x# t
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])
    $ J& h! _& I2 L; i8 w
  14. & m- K* ]/ @2 B9 t0 M\" ^
  15. bias tensor([-0.0152]) tensor([-0.2023])/ A2 E- f6 \2 q# ?! X3 U
  16. 5 b8 ~\" Q8 D\" }6 t# W
  17. tensor(0.1529)1 K: t# ?  I( y1 a# {- p0 J: R

  18. 2 ^4 {+ S  {3 T. ~5 W: ?
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])' u. S5 y! F4 f1 Y/ l+ }
  20. . L' t( z; j/ H  z
  21. bias tensor([0.0050]) tensor([0.0566])
    6 r8 f9 ^4 D6 o3 J9 l* w
  22. % c# c) T  J2 j( _8 M, O. G1 c
  23. tensor(0.0963)
    # J6 c' E! M; b9 ]5 s
  24. 2 K$ d  E  J0 O\" Z9 v
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])
    ! ]/ p; N2 l; a3 o1 z( d
  26. ( E1 ^; h1 ~: d7 Q# I+ X
  27. bias tensor([-0.0006]) tensor([-0.0146])
    7 @- N7 `\" H0 {( F6 }% K6 y& V
  28. 0 |0 k9 l2 i* q4 b: w
  29. tensor(0.0615)
      {# J0 x, i9 Q( D8 j. r1 \8 V  S

  30. 0 y9 X7 D( Y; a  T/ o
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])( C! |1 J; a8 y$ `2 m* g6 z

  32. * P1 m* E* S! k
  33. bias tensor([0.0008]) tensor([0.0048])4 E* k; w, R% x; w9 U

  34. ! p+ ^; p, T\" D3 B7 [5 T2 Y9 P
  35. tensor(0.0394)
    ) t5 M4 z6 u6 p

  36. 5 Q* l. t- s& y0 M' x, s* P8 _8 N
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]]): M4 Y4 m4 h& J2 t6 k

  38. $ t9 z$ q6 y& ~6 S
  39. bias tensor([0.0003]) tensor([-0.0006])/ }; {: G- [- _4 d
  40. 6 n) K' J4 h& c  }! g\" J
  41. tensor(0.0252)
    , N! q, g; Y1 c5 Q& K! A
  42. 6 `8 J6 Q- _' H% a
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])
    . F\" n( P. r3 s# p% n6 O) g1 D
  44. 9 L, c% s( i, U: O$ x' R
  45. bias tensor([0.0004]) tensor([0.0008])
    \" O( c& G6 A9 U( A/ O- {, f

  46. 4 E  Q5 }( L/ N; P  n$ p* |- F7 ~+ O
  47. tensor(0.0161)! u: K9 k1 f+ E. V: R

  48. & }- o: u( a5 e' W' @8 y\" c
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])0 @, R; G: O3 y8 m6 ]

  50. ) {4 e8 _9 @; R$ R+ ^) Q# u1 c
  51. bias tensor([0.0003]) tensor([0.0003])
    $ H1 B7 }+ {% ^6 R# M* f# L3 M2 u5 T\" v
  52. - D1 |! M2 A! N+ C$ ]$ W- @4 Q
  53. tensor(0.0103)
    1 Q1 q+ f1 b  X0 b0 p# Z
  54. 8 X2 r/ z$ i+ P. C: A2 }0 ^% N
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])\" X* G# l8 m4 f
  56. # X! J4 U6 E+ t
  57. bias tensor([0.0003]) tensor([0.0004])! h- z7 b6 ?6 h4 o/ V* B

  58.   \( G+ a: D/ R8 |\" M4 j: w) P
  59. tensor(0.0066)4 Y1 @9 M7 [% J2 h) C
  60. $ w' ?+ ]5 Y. N9 P
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]])( l' J, r# n* I# Q6 |

  62. * L2 P* H7 y8 N- \; s
  63. bias tensor([0.0003]) tensor([0.0003])
    + t0 Q( d% f% G! j4 l4 s

  64. ) ~7 t\" G& H  g
  65. tensor(0.0042)\" v3 L7 Z; J& @* D% o
  66. . V& m7 N! w9 M
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])2 S4 Y6 F3 [7 {
  68. 2 ]: t3 ?4 L! F
  69. bias tensor([0.0002]) tensor([0.0003])) k4 R1 ?\" j' t2 K- k# g
  70. 9 J% [6 r/ e. y1 T/ i+ ^$ K) x
  71. tensor(0.0027), a# q2 P/ D. ]; u7 @3 p: v

  72. 9 o6 ^$ ?$ r; S; Z
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]])\" m) `4 C! T8 B: e
  74. 4 u% u1 j1 m) C; J$ d$ L
  75. bias tensor([0.0002]) tensor([0.0002])
    : s( |  N4 [% F: m3 ~+ R5 Q% L
  76. 4 J( s. M& g' m& z$ _* j# }
  77. tensor(0.0017)+ s1 z5 x/ @9 Y

  78.   B, |& ?  L9 S' {
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])- J\" @/ w4 c9 a

  80. 0 E* L% R, {9 x3 O& E! I
  81. bias tensor([0.0002]) tensor([0.0002])
    % Z3 E% _; c4 P
  82. ) f0 n; d+ k; C5 t\" q
  83. tensor(0.0011)% s, K$ Q, k9 z4 _1 \
  84. 6 v3 V' \7 T% G# Z4 V
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])
    4 Z/ J\" m1 m\" P- t. v; R! E
  86. ; C! G0 C/ I  [\" d
  87. bias tensor([0.0001]) tensor([0.0002])
      H, Y9 z  O& J, t' i

  88. ' t1 r1 f1 w9 Z7 L) U, [8 ]
  89. tensor(0.0007)3 Z2 Q- J' l3 V8 W. D% l& \( U

  90. $ W( [  L1 N0 I: F. G' e3 I5 T
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])# X4 Y; b) j4 X) v3 N+ Z

  92. 7 n) L, P8 s\" H; o5 V6 `# C# F
  93. bias tensor([0.0001]) tensor([0.0002])3 x: d2 J3 g! l( c0 R

  94. - a* \( q# C' j( a3 a: t7 s6 _+ L/ W* a
  95. tensor(0.0005)
      {% Q6 M# u2 _4 @) Z8 J. h

  96. 0 P+ s/ q! o: n+ U# v\" f) J
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])
      r) x: ]2 G$ P, j$ w

  98. % f( d; W3 ]4 f\" B/ }, c& a
  99. bias tensor([0.0001]) tensor([0.0001])( o& a+ d( Q% p( D3 D* V' D+ j\" k
  100. * D4 P% E0 k8 m6 T
  101. tensor(0.0003)9 t% W\" k1 W- r\" B1 M. ]
  102. 0 ]6 q\" l# ]0 D. |; ^
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])
    , o- ?. V; [6 D2 d- ^
  104. / H& `5 r7 I2 [  @3 e3 X) ]
  105. bias tensor([9.7973e-05]) tensor([0.0001])/ p% a1 y1 L9 a, z5 j
  106. ) T/ n. A& B3 j' }# t( a$ L
  107. tensor(0.0002)
    + v/ F* h9 G' Q2 r; l
  108. \" W' Z5 d3 r2 Q, t, ~7 \! a# |1 O
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]])  R  `+ s. Z+ g  v+ w# e; p8 A8 e
  110. 0 z5 @7 Z# w% @: T$ C
  111. bias tensor([8.5674e-05]) tensor([0.0001])) l& \* u- r6 m3 y
  112. - U0 P' R# c* z
  113. tensor(0.0001)
    8 H2 d7 z$ Q8 y) a( a  ]
  114. $ w5 W8 C6 L7 J; T2 v. }
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]])0 {9 T/ [9 S! \( S9 ]3 r2 S
  116. 8 k6 |5 M: t/ d, h, Z1 h: M% z
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])
    1 s\" |) [. |) ?1 i- |\" T
  118. / e1 X  W0 g! u) \
  119. tensor(7.6120e-05)
复制代码

, |6 P1 M0 L* G/ g1 k9 z; I
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 02:40 , Processed in 0.383124 second(s), 51 queries .

回顶部