QQ登录

只需要一步,快速开始

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

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

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

1189

主题

4

听众

2934

积分

该用户从未签到

跳转到指定楼层
1#
发表于 2023-11-28 14:57 |只看该作者 |倒序浏览
|招呼Ta 关注Ta
SGD是什么
- A% k" W6 R1 lSGD是Stochastic Gradient Descent(随机梯度下降)的缩写,是深度学习中常用的优化算法之一。SGD是一种基于梯度的优化算法,用于更新深度神经网络的参数。它的基本思想是,在每一次迭代中,随机选择一个小批量的样本来计算损失函数的梯度,并用梯度来更新参数。这种随机性使得算法更具鲁棒性,能够避免陷入局部极小值,并且训练速度也会更快。7 I4 w6 M% y9 G
怎么理解梯度?1 J' P$ V5 v. V: [' ~8 Q
假设你在爬一座山,山顶是你的目标。你知道自己的位置和海拔高度,但是不知道山顶的具体位置和高度。你可以通过观察周围的地形来判断自己应该往哪个方向前进,并且你可以根据海拔高度的变化来判断自己是否接近山顶。0 z8 \+ H4 t9 b8 y: Q; p8 c

# x. a& E% E$ k' p' M/ ~# f: H2 j在这个例子中,你就可以把自己看作是一个模型,而目标就是最小化海拔高度(损失函数)。你可以根据周围的地形(梯度)来判断自己应该往哪个方向前进,这就相当于使用梯度下降法来更新模型的参数(你的位置和海拔高度)。
- d8 P1 r- x& a- Q: W; p2 g' H$ U9 M, l1 r% s- ^  w( q
每次你前进一步,就相当于模型更新一次参数,然后重新计算海拔高度。如果你发现海拔高度变小了,就说明你走对了方向,可以继续往这个方向前进;如果海拔高度变大了,就说明你走错了方向,需要回到上一个位置重新计算梯度并选择一个新的方向前进。通过不断重复这个过程,最终你会到达山顶,也就是找到了最小化损失函数的参数。
1 z! x/ H4 ^. f5 t9 w: v  b+ g7 ]: b* j8 c2 |$ q
为什么引入SGD2 g) P, k# j5 O0 D
深度神经网络通常有大量的参数需要学习,因此优化算法的效率和精度非常重要。传统的梯度下降算法需要计算全部样本的梯度,非常耗时,并且容易受到噪声的影响。随机梯度下降算法则可以使用一小部分样本来计算梯度,从而大大提高了训练速度和鲁棒性。此外,SGD还可以避免陷入局部极小值,使得训练结果更加准确。
  a- p8 L% @; ]& ?* A( f5 m; W8 @1 f. A7 {- v
怎么用SGD
  1. import torch
    ' i+ D# x1 {\" D. P! c8 t\" w' u: S, a
  2. 9 U' j0 t9 Z* t9 a; v- H7 J1 g
  3. from torch import nn( o0 c# j& o- |- d4 n
  4. / G( |5 {0 S1 e: o- J; ]' G
  5. from torch import optim
    ( U9 N) O+ ]7 n+ T/ G* l
  6. 6 c9 O! r% z4 P4 d' _$ z6 C4 H; n; z: [

  7. % L$ p$ W& D( r. S

  8. 2 Q$ {* j$ K9 y2 V: h5 @
  9. data = torch.tensor([[0,0],[0,1],[1,0],[1,1.]], requires_grad=True)
    ( X0 ]* q* u: _% O/ f) A8 O
  10. 0 u  G& I' `. V: k7 z
  11. target = torch.tensor([[0],[0],[1],[1.]], requires_grad=True)$ F; Q0 Q+ l0 I  W0 v- Y

  12. 0 I; W. v5 @  X1 U

  13. 9 [8 ]3 k6 A# K

  14. 6 r1 k. t9 }5 J7 x9 [  v
  15. model = nn.Linear(2, 1): j* |( O; `% ^; B9 p' A$ c; f  C0 ?

  16. + _& j/ L2 E1 n
  17. 6 S2 O8 F8 O- d8 C9 r% }5 }
  18. ! b% k2 h$ l+ f
  19. def train():
    $ O/ Z! u  v; l# l/ @- A( v& L: o9 p

  20. . R* S\" k( C+ f# C
  21.     opt = optim.SGD(params=model.parameters(), lr=0.1)  e3 w9 Y2 p1 P2 \+ l. L
  22. 7 u: U5 O  |# ~5 F: I& E: h: V
  23.     for iter in range(20):$ c/ F! H1 h9 m% c! h

  24. , [9 n  u( m+ T7 G4 l2 d0 P+ j+ N, v
  25.         # 1) 消除之前的梯度(如果存在)
    % R4 D7 s- n( m2 C; Y/ o

  26. * \( k5 |! \\" ]3 p1 Q2 [$ b% n\" J
  27.         opt.zero_grad()
    5 S\" a+ r7 G$ g3 |, o. t

  28. * Z. i+ s. {! o6 J% B+ m/ o/ d
  29. . }+ u7 @9 Y4 v0 C

  30. 4 R7 Q4 y+ U' q# ^9 G
  31.         # 2) 预测
    % L& g. Z  {. h\" x7 o. h
  32. 3 s2 j. C) W) L1 ~/ G9 {& k
  33.         pred = model(data)7 O+ A( ]1 h$ G( J; W, o. G+ ?$ x) u

  34. \" N; ^! Y2 e) y3 S' o; o$ g

  35. 7 A& `) }& Y) K! ?\" |$ X+ O
  36. # K6 @' C1 r; z' F
  37.         # 3) 计算损失
    \" m. n6 N4 r& ]! T

  38. % v& h( ~( Y; ^$ Z& G5 r
  39.         loss = ((pred - target)**2).sum()
    5 X4 r* z* z: z  x

  40. 2 x% e3 B% ^3 k( @
  41. $ }& r0 Y. J: M$ |\" a) ^* j

  42. $ |% c* a3 h, R3 T# X7 O; C
  43.         # 4) 指出那些导致损失的参数(损失回传)
    * S, `) C# p& z2 v* R5 X/ {
  44. . c( ?\" p\" |* {1 R) U
  45.         loss.backward()
    8 j3 r. b* c! [- P0 G; C1 o+ l$ P+ n
  46. ' f- j9 D3 F\" f
  47.     for name, param in model.named_parameters():
      Q% F7 l$ X\" T! ~' o; B

  48. * d  A- ^% y; I* [1 n- x) ~
  49.             print(name, param.data, param.grad)( N+ p; b# T2 j8 O2 G

  50. & b+ U2 _* f; R8 Y3 t
  51.         # 5) 更新参数) A+ i, W# M4 G# Q

  52. ) O7 X7 V! X9 Z- i0 K
  53.         opt.step()
    ( u6 L% e) V; r4 x4 [, Y
  54. , G8 X8 g3 ?/ _/ k7 q, y* @( ?

  55. : t0 @3 R) G% o* R) d
  56. 1 G. @$ D\" ~* r: r
  57.         # 6) 打印进程
    9 R# M2 D7 s$ e( U\" r4 T
  58. 5 J& A3 E9 @; j2 P$ R
  59.         print(loss.data)$ f& o$ e1 x9 m; t4 v% j

  60. - H  R; U7 T+ E( X% R! t. ]

  61. ! \\" j$ Z* w; U

  62. 5 l0 g* z( s  `$ c. j- f6 G
  63. if __name__ == "__main__":
    # r% r3 d\" u/ _# ?
  64. 3 w. M' L( _\" s; `
  65.     train()
    ) I' Q& ?' @7 ?7 ~; B( |3 M* A
  66. ( D/ |1 x3 }9 x
复制代码
param.data是参数的当前值,而param.grad是参数的梯度值。在进行反向传播计算时,每个参数都会被记录其梯度信息,以便在更新参数时使用。通过访问param.data和param.grad,可以查看参数当前的值和梯度信息。值得注意的是,param.grad在每次调用backward()后都会自动清空,因此如果需要保存梯度信息,应该在计算完梯度之后及时将其提取并保存到其他地方。
! x! S; K( j0 p6 h2 J9 ~( k4 L5 }) C* J" f3 z! h" y; H8 o7 K% D
计算结果:
  1. weight tensor([[0.4456, 0.3017]]) tensor([[-2.4574, -0.7452]])
    + u. S, O) z2 A* A8 O+ O

  2. 1 E; H1 x+ }. S1 M5 V4 H7 [
  3. bias tensor([-0.2108]) tensor([-2.6971])
    4 ?' z. j# D* h& r

  4. ! ?2 a! H9 Z\" o\" b6 |
  5. tensor(0.8531)
    ! Y9 H4 I& I$ `; t8 A9 L# I9 a; g8 j

  6. * h( S, i3 c( {2 y* c8 w7 |
  7. weight tensor([[0.6913, 0.3762]]) tensor([[-0.2466,  1.1232]])
    : Q/ q$ E. ^) q1 q1 S

  8. ' c# n3 P2 S$ Q& Z
  9. bias tensor([0.0589]) tensor([0.7416])2 I6 c# s1 A! g1 _

  10. \" B8 V. N+ S3 H8 ]9 L, ?
  11. tensor(0.2712)& d* H4 t# ?* I- v) k\" w8 R

  12. # z) i  f3 i, d
  13. weight tensor([[0.7160, 0.2639]]) tensor([[-0.6692,  0.4266]])( t$ u\" C' K+ q+ M  W- O9 \
  14. + D4 m! k9 T( \/ G
  15. bias tensor([-0.0152]) tensor([-0.2023])0 z  j* r4 X5 M: S0 i. L* I

  16.   P& c- m/ q) Q+ k/ J
  17. tensor(0.1529)' ]9 p/ r* x4 L
  18. ( V/ [4 T4 G& T( r# O
  19. weight tensor([[0.7829, 0.2212]]) tensor([[-0.4059,  0.4707]])( A2 u\" Y  R( Z4 a7 g+ w
  20. 4 j3 z2 J5 K9 i! m
  21. bias tensor([0.0050]) tensor([0.0566])
    1 h. ~& p3 c/ e
  22. 8 }# T+ C- ?7 @8 x+ h
  23. tensor(0.0963)6 [$ A/ Q7 d4 X\" G7 A( P
  24. & S+ `8 }3 [5 I0 ]$ a) u) D
  25. weight tensor([[0.8235, 0.1741]]) tensor([[-0.3603,  0.3410]])
    : c1 |7 N' \2 T# B6 r3 e
  26. 1 B. H* ]- l' y) a
  27. bias tensor([-0.0006]) tensor([-0.0146])
    % y1 ~: C# G8 l  M
  28. , h2 ?7 y: j! p) ]) I
  29. tensor(0.0615)
    - U9 N) s3 ^% B: h

  30. ; B8 t% P& J) K: N
  31. weight tensor([[0.8595, 0.1400]]) tensor([[-0.2786,  0.2825]])
    8 Y7 n( ~3 l6 K6 l! z# j
  32. ! k' [  h6 A. m3 t0 r1 L5 Z
  33. bias tensor([0.0008]) tensor([0.0048])$ a, f! |- S) Y! ~* A  z3 t
  34. & _# x% S# d8 R/ f4 h
  35. tensor(0.0394)
    + G0 W+ J6 i& t6 |/ A

  36. + R0 x* {4 N% ?4 X* N( V
  37. weight tensor([[0.8874, 0.1118]]) tensor([[-0.2256,  0.2233]])# U9 |& i% q. X- r9 Q
  38. 0 @) @  Y: ]' I
  39. bias tensor([0.0003]) tensor([-0.0006])$ t8 E\" H$ p9 |: }5 g9 w2 k
  40. 4 O$ R1 s: [2 Z1 e* ^
  41. tensor(0.0252)8 c) Z) P5 j3 H+ {  x. a5 j

  42. 0 _% r7 \( {, P2 O9 O, b
  43. weight tensor([[0.9099, 0.0895]]) tensor([[-0.1797,  0.1793]])
      w& ]+ @; u* B
  44. 0 u! i* B; {. s: i& r: [
  45. bias tensor([0.0004]) tensor([0.0008])
    7 b# W9 E0 D* V+ L' {7 ~) y% n& r

  46. ' f) T+ p# C/ X) c! y
  47. tensor(0.0161)/ a+ Y8 b! Y; g5 [

  48. ( M7 d8 U: e# G: f. p
  49. weight tensor([[0.9279, 0.0715]]) tensor([[-0.1440,  0.1432]])
    4 y- Y4 p8 i& q% }% z, [0 f\" ]\" F) I
  50. 2 \6 e' q. q5 u& T! R\" r' O9 t
  51. bias tensor([0.0003]) tensor([0.0003])
    ! r* i7 U1 c7 _\" E9 N. [

  52. 0 A/ i1 N, K% j
  53. tensor(0.0103)2 q$ \; S\" }* r% R. x+ O' z4 U
  54. * A5 `2 K* M2 D- s
  55. weight tensor([[0.9423, 0.0572]]) tensor([[-0.1152,  0.1146]])
    . X1 A0 c0 o% P

  56. ) h6 }* c1 H6 S. C7 w
  57. bias tensor([0.0003]) tensor([0.0004])/ W2 n, T! T7 t( h; S+ h
  58. % M\" m. J) W$ Y$ ?8 A$ e\" w
  59. tensor(0.0066)6 `5 L5 p8 d+ h6 n, H  E
  60. 3 d4 G3 d- i3 p, U* j0 g
  61. weight tensor([[0.9538, 0.0458]]) tensor([[-0.0922,  0.0917]]). ^) o\" U% O1 O0 h3 v\" Y
  62. ; l3 w# E: h/ s' [* f' k( q
  63. bias tensor([0.0003]) tensor([0.0003]), q8 g' ^9 a/ V6 a% l3 y) \9 p9 M

  64. $ v\" X& t0 w\" ?$ e
  65. tensor(0.0042)2 ^6 @* s8 M\" ~- \, W9 |
  66. , @: H$ o: S+ D\" ~6 Q( Y: {& V
  67. weight tensor([[0.9630, 0.0366]]) tensor([[-0.0738,  0.0733]])
    1 X' A: m9 N* c) `4 p\" z

  68. 4 \9 J' d- h  p
  69. bias tensor([0.0002]) tensor([0.0003])# w; f$ Z& `$ o. W

  70. , ]! ~) Z1 x0 U
  71. tensor(0.0027)/ `9 I/ ?% P3 S7 ]# L& t; }6 Q: y, j1 N
  72. 1 W9 X. w0 r& h+ c% Y
  73. weight tensor([[0.9704, 0.0293]]) tensor([[-0.0590,  0.0586]]); d2 x( r0 g) c6 S, P# Q
  74. + p2 `4 m# t; ?6 H' P1 U
  75. bias tensor([0.0002]) tensor([0.0002])! T4 ^8 Z, U2 r) b: Q7 q- B
  76. % J4 U, A/ R- f
  77. tensor(0.0017)
    2 q# j+ |& \$ p& }$ V- \
  78. : {7 V7 A0 K( b% p\" p1 W
  79. weight tensor([[0.9763, 0.0234]]) tensor([[-0.0472,  0.0469]])
    ' W4 m4 G# y+ O\" ?
  80. $ F. j$ }+ A* ~9 A
  81. bias tensor([0.0002]) tensor([0.0002])
    0 g8 n( f2 T\" P$ \
  82. ' i* J7 e& W: T' [\" t; L
  83. tensor(0.0011)! n( Y8 B7 A: x' y: k1 {* q

  84. + H3 c* [) ?5 K
  85. weight tensor([[0.9811, 0.0187]]) tensor([[-0.0378,  0.0375]])
    . O9 {4 f3 F, R' a4 A/ p7 {

  86. ' h) q- b! }6 X\" }& G8 S& o5 N
  87. bias tensor([0.0001]) tensor([0.0002])0 q$ v2 u1 [# G# o% b
  88. 0 H+ X* B6 w9 O) a, A
  89. tensor(0.0007)
    + N: A. x) N* K0 `% O/ V: U$ i
  90.   v# Y  e: M0 Z5 M3 m) {
  91. weight tensor([[0.9848, 0.0150]]) tensor([[-0.0303,  0.0300]])
    # r, f  w0 C4 w2 k9 u6 x+ X
  92. ; @) c3 A% b) G( g# m% E
  93. bias tensor([0.0001]) tensor([0.0002])
    2 Y. X3 |/ _! Z. I5 i  M6 N, T8 o

  94. , F+ [+ T( _6 v8 R$ c0 o- \4 f5 F
  95. tensor(0.0005)
    2 f4 o+ d: F! c. \3 ], Q7 w

  96. : @6 @( e7 F1 ?. Q4 T& G
  97. weight tensor([[0.9879, 0.0120]]) tensor([[-0.0242,  0.0240]])
    3 x0 ^, b/ I# F8 s- k6 ?
  98. 7 C  f2 W8 h- B9 R- [5 r
  99. bias tensor([0.0001]) tensor([0.0001])
    6 B% a( x# p6 S2 E6 I9 l
  100. ; o) L& N% ]0 }7 Y5 q
  101. tensor(0.0003)
    , w% [+ C5 \/ F0 m5 J( d

  102. : J0 [2 y& B: _
  103. weight tensor([[0.9903, 0.0096]]) tensor([[-0.0194,  0.0192]])
    ; Z! E! N- z9 Q, |9 |, n

  104. % X' W, G$ M9 p, k0 V
  105. bias tensor([9.7973e-05]) tensor([0.0001])\" G. h% ]7 |( R9 u
  106. * ?( e0 k8 ?% d! E# F. y
  107. tensor(0.0002)# ]: d7 W+ h6 A
  108. & ?2 v\" R, `/ Z% ^
  109. weight tensor([[0.9922, 0.0076]]) tensor([[-0.0155,  0.0153]]): s; X2 e6 B$ J, j! H  q1 S

  110. 9 W) |; A3 j) S5 l2 V
  111. bias tensor([8.5674e-05]) tensor([0.0001])
    + v. q; H/ |; R4 N
  112. * |& q6 t! d) C# {9 x
  113. tensor(0.0001)
    1 |! \$ z4 J. g7 _7 P% ]0 j' h. C# U

  114. * [. e  P7 T6 u% ^9 k8 C
  115. weight tensor([[0.9938, 0.0061]]) tensor([[-0.0124,  0.0123]]), ?2 D! O9 _3 w- L
  116. . E2 Z' r( H% @% Z! w4 ]4 g5 E
  117. bias tensor([7.4933e-05]) tensor([9.4233e-05])& {% X6 n( ~$ G* {. m9 D

  118. # y/ l$ H8 y. l9 @4 a
  119. tensor(7.6120e-05)
复制代码

: b6 H1 R* |/ D1 i. d9 E& G
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-7-24 05:28 , Processed in 0.424273 second(s), 51 queries .

回顶部