设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    ) U8 G" O& p) }2 N
    7 M$ e8 E" W- m2 V( N1 B  P* B7 H. Z为预防老年痴呆,时不时学点新东东玩一玩。
    2 @$ Y4 @5 ?) P+ D& C% HPytorch 下面的代码做最简单的一元线性回归:
    3 I6 B$ A) s9 j; F----------------------------------------------
    ! m, c$ [7 g7 Zimport torch
    * m& z9 b4 q* n3 C9 p% E0 Iimport numpy as np. ?: b- b: I+ C
    import matplotlib.pyplot as plt* Y5 D# }- S# L
    import random3 ?/ ^9 O/ s1 t( m
    8 e' ^; Q3 {2 v# ?
    x = torch.tensor(np.arange(1,100,1))* \4 @: m3 d( r, D
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    " i1 j5 o" `6 a2 N+ ~2 b$ ~/ z  C6 W# X& }6 Q: Y5 b
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b( o+ _* n% J" q- i) \
    b = torch.tensor(0.,requires_grad=True)
    * a7 A8 f% a! L' Q! _7 b$ h: K+ |2 `9 S+ A; M" b
    epochs = 1002 m/ C% p. W# ?/ {" I$ e9 Z
    $ S. f& U' a" E6 a6 |8 d) q4 u  }/ k4 ~
    losses = []
    ; k& |2 I0 m6 U/ Q. Ifor i in range(epochs):) ~5 b! N6 y; e, Q+ T9 |
      y_pred = (x*w+b)    # 预测
    ! C/ ~4 y; Y, K: r7 c  E$ f4 h  y_pred.reshape(-1)4 N- M2 {5 Y$ h# _
    ) k( F7 m9 \' r: P/ e6 K
      loss = torch.square(y_pred - y).mean()   #计算 loss4 l% [1 r" w/ U
      losses.append(loss)
    9 w: n2 S. L" |% Y  
    3 i0 i# Y. z3 j" M9 Q  loss.backward() # autograd& w# K+ b! `" c  H7 \
      with torch.no_grad():
    ( X# I/ U  B& m0 ?+ L: [# b( L    w  -= w.grad*0.0001   # 回归 w
    ; j) Y  v) U# ?" W1 w5 ^; D    b  -= b.grad*0.0001    # 回归 b 5 H1 d8 T6 C& C# s, y, v
      w.grad.zero_()  
    " a2 U5 G4 C1 C  b.grad.zero_()
    # D3 g; q6 o- U" E" R  o
    8 K. H0 `0 S7 t/ p% y0 o3 E' Jprint(w.item(),b.item()) #结果7 `4 w) l+ R. y7 |4 ^$ d
    # L( v* i. X* p: k- |
    Output: 27.26387596130371  0.4974517822265625$ C9 s- s5 K4 [7 Z
    ----------------------------------------------
    8 H3 C4 A/ l( y$ \: I8 I$ z6 [最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。0 i0 K+ d& ~% A) Z
    高手们帮看看是神马原因?
    2 O4 ^- K3 |' |

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 * M  f4 I' g( M( L( ~

    # b7 h! N! B3 C, a没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    8 H3 \! R. g2 m! p$ M-------
    ) k6 d4 |3 S+ {不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。; B. q: W& ^# K2 w- b4 K
    -------8 l3 l  f! d( A% h/ L. u" w
    算法诊断部分,建议把循环次数改为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
    $ c+ q: K4 I8 s9 A* i. O  O- M$ i没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    # B5 {2 u3 o/ }& P1 L1 B1 t. ]-------
    . c1 _" Y7 n0 t6 u1 P: q不好意思, ...
    " w, M7 h. ~( L9 a- \2 t/ ^% D
    谢谢,算法应该没问题,就是最简单的线性回归。
    + h' v8 q, y; F& W" h& ^我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 " Y2 C/ @5 ]: Y" f) f; `' w7 }
    雷达 发表于 2023-2-14 21:52
    ; l* {0 q8 K% S6 L& x) h0 s9 S谢谢,算法应该没问题,就是最简单的线性回归。- Q% i6 C$ i+ t5 Y$ ?
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    " Z5 F  H0 \, j

    , R" d. I( f# D* U) w+ p6 B5 a, y刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。0 t  q' H5 M5 n. ^( _' F

    . ]# g/ y# t; D0 }或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 1 x' y2 p) r  E8 ^% b0 F) |6 y1 \
    老福 发表于 2023-2-14 22:00
    # p% E7 V8 u1 e! k$ I刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。$ z  L% G4 }) A0 o( m

    ; t4 r/ ?+ @' F- \4 S4 C或者把b但的起点改为1试试。 ...

    ! z- _( U  j8 X5 h1 M2 V/ g* t- A2 D1 h5 s4 O9 z
    你是对的。
    : G1 e' j5 z8 X7 m去掉了随机部分
    1 `4 C; L, B, H+ h3 J#y = (x*27+15+random.randint(-2,3)).reshape(-1)( b4 l* J! o0 E' N, G+ j1 e
    y = (x*27+15).reshape(-1)& Z9 y" V2 w' y& g

    5 e2 _5 Q! e7 y9 Z( b4 Q1 ]% S6 ], j循环次数加成10倍,就看到 b 收敛了
    1 ^: P8 r& ~2 zw , b9 A$ n0 S" t/ ~- z! f) I
    27.002620697021484 14.826167106628418
    # P* F, Y" o$ I
    # p4 O. r! B. \和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-8 03:44 , Processed in 0.060210 second(s), 19 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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