设为首页收藏本站

爱吱声

 找回密码
 注册
搜索
查看: 3297|回复: 4
打印 上一主题 下一主题

[信息技术] 继续请教问题:关于 Pytorch 的 Autograd

[复制链接]
  • TA的每日心情

    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    , s0 f$ n0 @$ O" n8 w9 g1 s/ P) [4 Y9 G1 R; ?- R% R
    为预防老年痴呆,时不时学点新东东玩一玩。
    . e8 y, ^! z8 P  mPytorch 下面的代码做最简单的一元线性回归:
    * r# K4 c8 T% c) T----------------------------------------------+ c+ m( K( D% N5 B
    import torch+ T7 W6 }7 ~& ?$ }# U; g2 r- Y
    import numpy as np! [6 J3 d! e6 w9 s
    import matplotlib.pyplot as plt
    ( F5 E# P) d- Fimport random& ~6 m. W1 h$ C1 P" c, U# G
    7 X5 a& \. m9 q  k% X4 _' R
    x = torch.tensor(np.arange(1,100,1))
    : g) |; M0 S. s1 h3 R* ky = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15; P( W. [& F) c8 A2 k; L  j: ^3 \" x

    & x; K; `. T  j4 r" k3 jw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    ' \3 e; G- E8 {1 Ab = torch.tensor(0.,requires_grad=True)
    % r" z, C" K. D7 T: i
    " R7 f5 F0 n+ q4 Wepochs = 100
    ) V: k- Z; I2 X0 w0 e& R$ b5 Q9 `+ @2 f/ l: F: o
    losses = []
    5 K5 w& t7 b3 L: x  [* J! @6 ifor i in range(epochs):
    & K; k) L9 x9 R( y. Y$ O) X  y_pred = (x*w+b)    # 预测
      a( F) `* |* Z, O; l! N  y_pred.reshape(-1)
    ! t2 b+ G# u: L$ t2 _8 F - }4 r' U! W6 v/ x7 J
      loss = torch.square(y_pred - y).mean()   #计算 loss
    0 L0 X) F- l6 n% Z# P  losses.append(loss)  C: e: M& c' \
      # A' S: Y" E: m
      loss.backward() # autograd& N5 B9 a' H2 a
      with torch.no_grad():
    ; S" N6 t1 ?3 E! i/ z) E8 i% m: G    w  -= w.grad*0.0001   # 回归 w
    ) W- _  [4 w+ d4 w    b  -= b.grad*0.0001    # 回归 b
    , T4 m& Q  ~* Z3 S% W  w.grad.zero_()  
    ( `% }9 m! P& w! n7 x2 j; l  b.grad.zero_()
    " q1 i9 b' N* x' N  I# U9 q  Q$ {3 d# k% ?* T7 X. C& i
    print(w.item(),b.item()) #结果7 T4 z/ U9 b" X+ K3 z" c2 j& T

    ) J: n) F4 `7 j9 UOutput: 27.26387596130371  0.49745178222656253 Y7 L, G$ B7 n% F
    ----------------------------------------------' j- l$ T5 b6 z) m
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
      \' w. J* a9 I; V高手们帮看看是神马原因?* q9 G7 u; r4 j/ Z* r; `" T

    评分

    参与人数 1爱元 +10 收起 理由
    老票 + 10 不明觉厉

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    # `5 O0 I$ ~" x$ q- O
    7 U) o7 Y4 y7 v. N, b+ }: a( i! K; L没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    9 E( i3 D9 z  d( }! _-------) Z5 c# X4 r* H' p( m/ q
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。5 G0 c( a0 J6 G( V5 S
    -------' ~, T9 [( d6 i5 T# i
    算法诊断部分,建议把循环次数改为1000, 再看看loss是不是收敛。有点怀疑你循环次数不够,因为你起点是0, 步长很小。只是直观建议。

    评分

    参与人数 1爱元 +10 收起 理由
    雷达 + 10 谢谢建议

    查看全部评分

    回复 支持 反对

    使用道具 举报

  • TA的每日心情

    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    板凳
     楼主| 发表于 2023-2-14 21:52:57 | 只看该作者
    老福 发表于 2023-2-14 19:23! |" u/ v* \5 Y$ u4 u
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?$ K8 W: r* b  }- T/ u' v2 z
    -------+ S( _* ?! ]. A3 V
    不好意思, ...
    ; y9 L- c( O) z
    谢谢,算法应该没问题,就是最简单的线性回归。0 P. |7 H8 ~8 x; S
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    ) C$ x0 p" A4 h2 h& p
    雷达 发表于 2023-2-14 21:52* @# n# e( t! J/ H4 B6 O. t
    谢谢,算法应该没问题,就是最简单的线性回归。3 b( i6 g4 o/ W
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    7 X" F7 K3 {+ U' L
    . m: Y. p) T7 ~8 a6 h
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。: t: \% K4 l1 p" L' H2 _3 y

    " p. @) }$ F( L& c或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    : S) g1 b: T& G1 n: F3 `
    老福 发表于 2023-2-14 22:00
    0 B! _( R: R' v0 ^% \刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。1 J' B' Y* s8 E' J, s/ v

    3 F1 K: S0 B  y) \0 {或者把b但的起点改为1试试。 ...

    . P. n% F9 v3 z4 A  y. l3 m" j' M0 f! y/ R, d' U( E$ C+ F
    你是对的。
    $ G- e3 P# f& _去掉了随机部分
    ) J$ n' m. ~1 U! i0 e8 w; T6 |7 m& @1 l#y = (x*27+15+random.randint(-2,3)).reshape(-1)0 {5 `# \0 i/ d0 R
    y = (x*27+15).reshape(-1)% Z! ]5 P5 }& M$ b0 W* v2 k; M  \

    ( \: m. Z; c/ a$ p' b& T* d循环次数加成10倍,就看到 b 收敛了
    4 F+ J7 |3 e& A5 _/ kw , b
    ; i) o9 h& K+ W+ V% H27.002620697021484 14.826167106628418, ?3 J7 `9 W7 k6 _8 D
    % ?0 P5 j# B7 H" h) l. P
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

    手机版|小黑屋|Archiver|网站错误报告|爱吱声   

    GMT+8, 2026-7-23 19:17 , Processed in 0.057505 second(s), 17 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

    快速回复 返回顶部 返回列表