设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 : Y6 a2 ?& T1 ~7 H
    , _& _3 T- x2 @. M
    为预防老年痴呆,时不时学点新东东玩一玩。
    $ T& v7 f* K( H3 D8 `( YPytorch 下面的代码做最简单的一元线性回归:. l# n/ b8 g1 v: X! y6 @
    ----------------------------------------------
    * h3 p6 I% P; [import torch
    : \4 a9 s" q0 K: Qimport numpy as np  J2 C+ Y% \5 g4 `
    import matplotlib.pyplot as plt
    $ _) D$ S6 J- t& D3 Gimport random+ L5 q! U  p8 C# u# V. X

    * h+ f( M8 `4 V3 o& K# l, Qx = torch.tensor(np.arange(1,100,1))8 L4 Y$ s9 P) B; M/ _
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    ( _0 ]) g2 ]7 Z* b* k# T7 L4 v+ n) e. U1 i8 `) I  J* u0 Q3 z
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b, q  x0 A+ _8 w- t$ L: B
    b = torch.tensor(0.,requires_grad=True)$ H2 y1 Y4 t/ Y3 t+ Y* Q

    3 \9 Z( O2 N2 d5 f: K/ C; n" iepochs = 1008 U, x5 d5 [6 L& Q7 o4 R

    6 z/ U( y- ~5 \' U3 olosses = []7 r2 Z6 l1 o7 T1 Y, O! D
    for i in range(epochs):% M4 ^" y2 b9 O
      y_pred = (x*w+b)    # 预测
    : Z: n6 W2 e2 }  O0 ^  y_pred.reshape(-1)
    + J7 O* t# i. R' X8 D9 R7 c# T( \ # G7 N) s, Q& z
      loss = torch.square(y_pred - y).mean()   #计算 loss
    9 U( B* S1 k; r3 U- z  n  losses.append(loss)! y* b. B* T2 g0 m, q) g/ f9 Y
      
    . r! z) L& ]" _  loss.backward() # autograd
    ! b# z. @  [* V( d) T  with torch.no_grad():
    4 h% L; _, t5 J    w  -= w.grad*0.0001   # 回归 w/ b: a9 k9 Z& i# m7 `
        b  -= b.grad*0.0001    # 回归 b 1 I8 r- P! L9 U/ ~5 T' @
      w.grad.zero_()  
    / N6 L- o7 P1 `7 r+ ~1 h$ U9 W  b.grad.zero_()* q* c% h* `$ [: W

    : M6 P7 m5 ^6 W% _- bprint(w.item(),b.item()) #结果0 J! r" ]# f+ u) Q* ]. ?

      |6 D* \; r- q1 I3 l( FOutput: 27.26387596130371  0.4974517822265625+ s8 r8 i$ ]8 Z* c2 L) y& q. W
    ----------------------------------------------
    % p" y/ C  H  k( o# O最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。1 `8 ]8 v* T" |7 D
    高手们帮看看是神马原因?$ k: u% q9 B3 F, ]

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    , W. i# Y9 d( t& j$ z
    2 k  \$ @$ B& t1 H$ L( M! Z& F9 M没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    - V9 \# q& T/ q-------
    ( X; t! b2 u- q' w4 j. ^不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。8 G" U) d  n& p% n
    -------0 [+ ~9 s. m  s4 y9 w) n
    算法诊断部分,建议把循环次数改为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
    & b5 l6 C* Y7 J# a# I& Q  B没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?- Y8 O# k6 o' j. N2 B; o
    -------
    " k9 k4 F8 A$ p, L/ M8 f* o5 E不好意思, ...
    , d; X9 t0 ?4 g2 m% i( G/ m
    谢谢,算法应该没问题,就是最简单的线性回归。. ]0 T- f% Y1 S) ~( a3 c1 L' ]
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    : i3 B& {: ]1 U4 m3 C3 m4 p  o
    雷达 发表于 2023-2-14 21:52
    + q/ L: x- F9 P* {谢谢,算法应该没问题,就是最简单的线性回归。
    * \/ T6 U. J* v. ]8 G我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    $ Z6 }/ L2 k8 Y0 V1 p! Q

    " z+ }5 o; m% k( ]刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。- H0 t6 e3 }* V! m5 j: c' N
    2 ], p# d& w4 a
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 # t9 k$ X) V" `1 r' y, z6 L: ?
    老福 发表于 2023-2-14 22:00/ l6 d9 _( r( q" b: l
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。# U+ N3 e+ ?% H) }, J) {, w
    , D2 R) W8 H) X% V
    或者把b但的起点改为1试试。 ...
    % e8 e7 M# F  v/ W. D3 u

    % l: Q' U5 @$ o. K你是对的。% ~$ s6 I! K; I# F# M
    去掉了随机部分
    ( Y0 L2 R# |6 h6 Q: q4 y, @#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    7 P9 u' o. |8 Y/ W7 k) qy = (x*27+15).reshape(-1)% e6 V. e5 G9 B0 Q

    ' z) k3 O8 g% o循环次数加成10倍,就看到 b 收敛了$ I7 D' o2 P3 O5 f% `' s
    w , b$ C' n, W6 u* m+ u& m* Z0 p
    27.002620697021484 14.8261671066284180 G$ ?2 G9 x9 e) i9 p: f

    , W8 E" w& [8 e8 S+ U* ?和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-5 02:09 , Processed in 0.060919 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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