设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 1 r: w" Q" R4 D
    9 O6 w( d7 K6 Q
    为预防老年痴呆,时不时学点新东东玩一玩。0 ], y6 O5 H( B5 v3 F' O
    Pytorch 下面的代码做最简单的一元线性回归:$ q. S4 G" e3 a/ o
    ----------------------------------------------
    0 W8 V9 F7 ~8 F5 s, l5 m+ yimport torch
    3 s5 @& H, p. ximport numpy as np
    ) m% C+ O$ }) {, n% |import matplotlib.pyplot as plt+ p+ K4 e& Z( |4 m4 k
    import random6 `6 l' V' D# Y3 m

    2 u) y, B. X" \3 C. `& Fx = torch.tensor(np.arange(1,100,1))
    3 c$ n7 m" _) h: \y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    3 L  _1 k+ H! Q9 k2 p: n: w
    9 N/ G) ^3 O9 J5 G$ f8 G5 ]w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    % i! V' S  ^; Lb = torch.tensor(0.,requires_grad=True)
    * G& _4 _' S- G& D
    , V. b$ _8 W9 M0 a, U" Depochs = 100
    ( u$ G, o( H- ~" W* S% w0 W2 h: n
    + s* m6 t. X7 Y* u) |2 slosses = []
    ) l9 g6 w9 t' @: P- F. efor i in range(epochs):% ^* A0 q/ P* b( C  D* |
      y_pred = (x*w+b)    # 预测. ?& f' A* M+ y3 J" ]- s
      y_pred.reshape(-1): e9 ~1 f1 |/ i$ X+ u6 n$ T. s+ g' Y
    7 U* b2 R/ f; s! F
      loss = torch.square(y_pred - y).mean()   #计算 loss0 g! d0 W$ J- r( v4 r3 N3 m
      losses.append(loss)+ p' ?, q; Q, P, u, \, ~
      ! s5 ?; S' y5 w: @' }5 q
      loss.backward() # autograd1 L( v7 Y" s# y! r
      with torch.no_grad():9 e% @8 x+ f& m, c  O$ T" f
        w  -= w.grad*0.0001   # 回归 w' x6 o, X) j3 d4 S
        b  -= b.grad*0.0001    # 回归 b
    3 z0 r6 X. t# |! S3 z; F  w.grad.zero_()  
    % M- L+ ~  b4 m2 E9 k  b.grad.zero_()7 H# Q$ U: X( p. i0 J7 L$ b
    0 T; `# o0 C8 }* g( J0 h; x6 H! q
    print(w.item(),b.item()) #结果
    6 U0 P9 y; s( x( [: M# f8 P( O  v2 y
    ' ]! T( U% f5 o8 IOutput: 27.26387596130371  0.4974517822265625
      o6 q  \$ `  C, z----------------------------------------------+ }2 s$ G+ w/ V. H9 O* k, p- l, V
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。) i7 ~7 Z; Y3 a. m- g- R) O
    高手们帮看看是神马原因?8 w, v: j( T$ Y0 N" f+ j

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 8 y+ l8 k: Y8 ?% M3 X- ]' a

    , H5 w, N9 Y! d没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?# G; u8 |! J1 \
    -------
    ; I* x# t. L' D- o0 M1 `* f9 [  E不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。/ j! e( C$ a1 f6 b
    -------- T0 I6 P4 F: ]0 i8 \9 o
    算法诊断部分,建议把循环次数改为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+ _. K. E+ ?6 t) }
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?) b% V& j# C' u; L9 k' H/ t# n) k
    -------
    ) f2 _! B! G: B: n$ z不好意思, ...
    / W' p5 [- a. H6 Y' r0 o
    谢谢,算法应该没问题,就是最简单的线性回归。
    5 U2 P9 h+ l1 q7 R4 q我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 1 `8 @9 I3 f4 r
    雷达 发表于 2023-2-14 21:52( @8 h1 b1 J9 j) ~" _8 c1 V5 {
    谢谢,算法应该没问题,就是最简单的线性回归。
    " r1 O8 w: E8 K6 `$ Z" [我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    8 a1 Y3 I4 x) {- Y
    ' O2 s  y9 \( j( g6 B1 s: \5 B刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。. y2 N/ t/ ]" l' \4 y
    : t, W% s; J) [9 S; V% }
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 ) Y( J, O4 x! H4 T1 s
    老福 发表于 2023-2-14 22:00
    & O9 d/ f6 F  o/ a刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。' N1 K0 ]% @' u

    $ s" p# c" r( a6 V- i或者把b但的起点改为1试试。 ...
    4 A" J  I0 A4 {' q8 H$ u% A

    1 s" q7 q# g& u+ v" n. a你是对的。- ~6 v# {/ E* {: m; N) ]+ E
    去掉了随机部分) I; h7 n0 U( l! w$ j
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)
    * s+ l, L: d5 d( }. ky = (x*27+15).reshape(-1)
    - x) w: f6 B, G  s- Q0 m' Q+ C2 \+ e9 K( _) b8 I
    循环次数加成10倍,就看到 b 收敛了
    3 u  s# E' x2 Mw , b* I8 Z5 y7 F( C
    27.002620697021484 14.826167106628418
      [+ y+ w# W, u2 G  @. ]# r5 ~* y" j7 [, s( j1 U  e3 i$ a
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-22 16:46 , Processed in 0.058562 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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