设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 + z4 {. H/ n) T2 p3 ~
      X2 n9 U  {7 f! T, V- P
    为预防老年痴呆,时不时学点新东东玩一玩。0 r& t" S4 v9 ]
    Pytorch 下面的代码做最简单的一元线性回归:
    ( Z7 q/ F& t2 P8 V- W% V9 ]----------------------------------------------
    * ?( l7 v. H+ Z0 k1 W2 ^import torch- E! w8 E6 Y1 B% H% \4 k
    import numpy as np
    " h: y/ \# q5 Q7 ~import matplotlib.pyplot as plt
    ) N) n8 L  |  I( \% L$ |5 fimport random4 }2 r% t" S1 L/ }6 l4 S
    * U4 I# F0 @6 @3 ]
    x = torch.tensor(np.arange(1,100,1)), p4 Q7 y+ H$ X6 P. q
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15" q5 b: L" [8 K( |

    " i: \6 X. V, I9 T: [, uw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    3 d' I9 w8 O5 j, O& z: y' eb = torch.tensor(0.,requires_grad=True)) }2 g$ {+ \/ |' R6 F6 D
    - R8 R3 Z6 @0 F3 x4 C' {: }) H
    epochs = 1008 ]# g0 B7 G. O6 |8 Z0 p0 a( _

    2 O% N1 p& s+ Dlosses = []9 W& g! X& c: @: O9 l
    for i in range(epochs):
    ! f4 Y3 O( j6 ~7 e* W  y_pred = (x*w+b)    # 预测
    0 v8 \$ J* ^9 g: V  y_pred.reshape(-1)
    2 I- }( A- X8 r% I
    3 X8 p9 o1 p9 u  loss = torch.square(y_pred - y).mean()   #计算 loss
    $ k8 \+ G2 B* N  losses.append(loss)
    ' v: a& A8 Z, n( X9 Q0 O3 @    p- U) j) c+ L) G9 H
      loss.backward() # autograd1 q+ O* X$ v+ t2 w- b
      with torch.no_grad():3 p+ I" v' u6 g; V0 l
        w  -= w.grad*0.0001   # 回归 w
    6 b0 E% \* o# J% _0 p8 K* @    b  -= b.grad*0.0001    # 回归 b ' y  c" G+ T% z/ b( T
      w.grad.zero_()  ! }: g! n, B, k3 u8 Y
      b.grad.zero_()
    / J  ^, ]8 `' f  ~) w8 [$ a0 B* x5 ~
    print(w.item(),b.item()) #结果- ~4 H4 H1 B3 p5 B  C/ B  C

    7 x2 j; ?+ h$ x% o' u* }4 X2 d2 lOutput: 27.26387596130371  0.4974517822265625
      z9 b# P1 z* e" F- f: i' q1 x----------------------------------------------
    ! D. V: v0 p  R/ G' K8 w4 W最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。' I, H3 @$ Z; P" v  M
    高手们帮看看是神马原因?" j; e5 j8 a# V

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 * I- j- g0 z6 [8 M$ K. U
    6 G: ]! |- k% w; r1 v% Z
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    . P6 E$ J' V0 v$ B, z( s-------
    ; o  s; ^+ M7 t+ }不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    * R  X% D& e+ D3 z; `5 ]-------! X& j4 x, I6 s9 ~! S
    算法诊断部分,建议把循环次数改为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, l6 f- I5 W- J% B2 T8 U9 |# f6 A
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    * d( [8 C! K7 M/ q2 o, I/ k# J-------
      z) k3 i/ R4 \/ K/ P% s# \" x不好意思, ...

    2 [, z. \4 F: g! Z6 r( N* W谢谢,算法应该没问题,就是最简单的线性回归。
    : V( E4 e+ H! s& l我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    ; r! p) u! E1 m  C) y3 A! p+ G
    雷达 发表于 2023-2-14 21:52
      X5 T. s3 \6 F, p2 m谢谢,算法应该没问题,就是最简单的线性回归。0 u) }5 y4 e3 C6 z5 ]
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    " y' ~$ ]/ g; y
    4 e) n* T8 G- o( _6 O0 q3 E' P
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    5 P& D0 V+ j. v. ]. s5 _( C' w/ q: U
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    + L1 X1 F) _* {+ A
    老福 发表于 2023-2-14 22:00& g( q1 b1 D6 h5 Q# }
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    / y) q! m" D6 ]
    ) u9 z) @$ V4 K0 L! H4 P或者把b但的起点改为1试试。 ...

    5 I( \# z9 l5 c- R1 j3 `* O) z
    你是对的。
    8 D1 B+ N9 S9 W去掉了随机部分/ S2 J* J6 g+ O! w5 y% }' c4 z6 Q
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)
    " F: _1 U3 U8 r- {y = (x*27+15).reshape(-1)
    * f; Q3 s! E! T8 C/ G
    0 }0 R1 A" ?" L  O循环次数加成10倍,就看到 b 收敛了
    - [& S) x8 v; Q+ e- Ew , b
    , {  X. s5 ]2 U3 _27.002620697021484 14.826167106628418
    3 i2 f7 U7 Y" U* G- A8 S
    0 G7 U7 t" W  d9 Y和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-7-25 08:03 , Processed in 0.057548 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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