设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    $ \6 i" G2 c; X( M* O; k
    8 o7 ]# V9 @; v' a+ Y为预防老年痴呆,时不时学点新东东玩一玩。2 e) }0 K7 g' @8 w0 x- X
    Pytorch 下面的代码做最简单的一元线性回归:
    : ]  M9 j1 e5 K8 j2 d- b+ ?----------------------------------------------8 V; G3 H( {  k9 S8 U! c- ^
    import torch. t) x% `! r. e2 y, ]
    import numpy as np
    % y, o+ y5 r( E+ r. wimport matplotlib.pyplot as plt
    5 b4 |: a) ]# {8 `! qimport random. |6 u) T5 l4 W) B9 D( g; T7 I
    & g( V3 A& U2 K8 G: }( G
    x = torch.tensor(np.arange(1,100,1))
    ; I  V/ s: s2 q0 x! E) r, v7 yy = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    8 S$ K( ?, I  d3 N/ J/ c3 q, a# z% i( P: g0 {( A
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    8 E# f. _% g, O1 h: Eb = torch.tensor(0.,requires_grad=True)( W; g3 g. ?* q: @- e2 d; i6 f8 e9 ~
    $ X; q% m( O0 `, M
    epochs = 100
    % |+ P- J; b/ r/ P) `2 x8 e- d9 a+ H9 `
    losses = []
    9 _$ D9 ?, e/ l$ ?% Ffor i in range(epochs):
    % @+ y/ [5 q) i" \$ k+ L  y_pred = (x*w+b)    # 预测
    , V4 x. x3 v. u) I, I/ Z+ E  y_pred.reshape(-1)
    # u( h* S- z) [$ b. k
    ) U' h1 x* q+ l) r/ A- Q. t) n9 I" H" ?  loss = torch.square(y_pred - y).mean()   #计算 loss" T; T2 `! a* D/ b6 w
      losses.append(loss)5 P2 o& I7 H- T# ^& A
      8 a7 q/ E) j" ?- w! z6 i
      loss.backward() # autograd
    4 j/ Z, j, R" |+ H% W+ y- I: q  with torch.no_grad():
    5 [) D' V) t" w  z4 K& X    w  -= w.grad*0.0001   # 回归 w% E4 `) Y' X! ~4 U8 H+ E
        b  -= b.grad*0.0001    # 回归 b " ?. m6 X( C, E7 z! V( V% \+ D# r
      w.grad.zero_()  4 {9 z1 ~; V3 r
      b.grad.zero_()
    7 p" Z# d! b  ^, o+ S. G4 o; r( n1 S% p. A$ C: N5 E" {
    print(w.item(),b.item()) #结果
    + m% X0 R8 h$ |: j2 j
      u$ {, X' l: t  c. L# rOutput: 27.26387596130371  0.49745178222656258 c' [3 H  k! O9 O
    ----------------------------------------------% ~$ W8 l/ u  s) `
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。( D/ X2 @4 a; A4 J3 E
    高手们帮看看是神马原因?
    4 M  {4 R/ P& M: @

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 ( L/ \1 W* }* f- v
    " k( W7 W7 y9 ~& X4 u9 @
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?/ q, Z( M/ ?: A/ u( u! B
    -------
    9 [, d5 w, T9 \5 O; z' d2 O, f& Q不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    ! }2 X; u- V- b! Q( K4 Y-------
    , m+ Q- s# {, Z% i' b- }算法诊断部分,建议把循环次数改为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+ q( k# U; k  @" m' [2 v5 X1 W
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?4 ~: ~9 n0 z' y: |& v
    -------- L! f" a; G7 V% A
    不好意思, ...

    : k  H+ [8 j- }, _/ b6 J- g谢谢,算法应该没问题,就是最简单的线性回归。
    & D% W1 R# y4 x我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    + e1 q% ]" E9 E& G, T
    雷达 发表于 2023-2-14 21:52
    4 ~$ y; |  e; i3 I3 t7 {! i谢谢,算法应该没问题,就是最简单的线性回归。
    & V% i" W6 p& o! _2 i4 J8 K2 E我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    $ a7 x  J3 b- n% P! \+ [! F* v& N  Q% |
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。9 m2 Z2 _2 X9 p$ e

    6 P* g/ m6 \( ]: d0 w2 ?或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 " N# i+ `- e! ~9 V2 \0 Q) u
    老福 发表于 2023-2-14 22:00
    : E2 _5 {$ b7 ?刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。0 O! u; w2 u, u1 f9 v+ }( S: @

    6 c. [  Q" h7 H# t+ A4 I% ^% g或者把b但的起点改为1试试。 ...

    ! W$ L( f  T! W- r* B3 n) a1 X" ^
    ! L2 ^- j; V0 l% ^2 @. B0 z$ B$ ^你是对的。
    ; _8 ]6 ?" d5 ~! s去掉了随机部分! ]7 c: q, b) O  P. q; E
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)6 R; s. C! V1 Z& }
    y = (x*27+15).reshape(-1)& f3 G4 O% \# e

    5 @, [( N. S9 @9 t2 Q循环次数加成10倍,就看到 b 收敛了
    0 g' e" l% c+ l5 _w , b; ]! ]" W1 \; B  d1 k+ W- T" L. B
    27.002620697021484 14.826167106628418' }9 d. D0 x( {& ]9 Z
    / N8 ], K$ Y2 K% t( M; V3 ~
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-28 17:42 , Processed in 0.056899 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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