设为首页收藏本站

爱吱声

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

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

[复制链接]
  • TA的每日心情
    怒
    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    : D) e& j' f8 R9 S+ y% I. A* m3 H
    为预防老年痴呆,时不时学点新东东玩一玩。
    # @- }5 }* u: {Pytorch 下面的代码做最简单的一元线性回归:
    ! U8 j' ~4 v1 I' g1 W# ~----------------------------------------------$ k% A0 m& [! e7 y& z& h2 R1 z
    import torch
    / e9 |1 e% S! S/ Z# d  simport numpy as np% j; }  @2 q! c- s
    import matplotlib.pyplot as plt
    7 y( \: t+ w! R3 ], qimport random
    0 C7 b( W) `7 G- E0 U' U4 Q, @) z! H9 F$ o) ~
    x = torch.tensor(np.arange(1,100,1))
    " {7 |- f0 J: x- l! h: l( ey = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15- R' R3 U% c) ~+ v

    - r. Y- Q; j. `* l: G7 d0 T! Fw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b, `' j. t5 ~) P# k
    b = torch.tensor(0.,requires_grad=True)5 a5 S, {+ b3 F9 t1 p# V/ e  M0 @
    : t: z; q3 D2 p/ f0 L
    epochs = 100' S$ N$ i" W% t4 U

      c0 F% {% w( e( s1 M7 ~losses = []
    1 H: M& {, n* f7 b* q6 n1 ]for i in range(epochs):
    # W- F! ]) c6 v5 d0 e6 A9 K  y_pred = (x*w+b)    # 预测3 _( `% ]+ j3 L5 w: F- u. V
      y_pred.reshape(-1). g# A" H3 E+ t+ H

    ( j+ Z7 ~3 }9 x& B1 _: t  loss = torch.square(y_pred - y).mean()   #计算 loss
    8 E9 S0 a, E! ]/ E! B  losses.append(loss)
    0 H8 I) g6 h9 ^0 G* e  # Y6 w9 }+ ], f6 a0 _- u9 v
      loss.backward() # autograd* q# W1 C$ v/ X0 {2 G) w
      with torch.no_grad():+ E% l9 G8 B. e+ p
        w  -= w.grad*0.0001   # 回归 w5 ^5 `4 }4 l( ^. ]
        b  -= b.grad*0.0001    # 回归 b
    : n& `+ J* X; w6 I) y  w.grad.zero_()  
      r( Q' {5 ?% s9 D  b.grad.zero_()+ |- ]6 d  M$ T2 e

    , k1 l* b+ ?' [: o$ \print(w.item(),b.item()) #结果
    1 H1 p  m  n3 ^
    7 v% k, J6 Q8 gOutput: 27.26387596130371  0.49745178222656258 K! {2 V( r: E/ E- U  Y
    ----------------------------------------------3 x8 L' A7 a' Q6 G
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    * d3 n/ t& N5 ]# ~高手们帮看看是神马原因?
    ( e! Q. T# Q4 D) d5 |- e; ?2 Y( F$ }/ R

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 8 |: f, W/ Z5 r4 Z4 ]& h- b5 H
    4 o; Z& u  O2 N) {4 }; t
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?; o- |. l  q0 M* q2 S1 I3 ~3 B& C
    -------5 l/ p8 k' `) b  z- r
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    . O8 k$ }$ N+ ~' {) D* O) `  t, w-------
    5 {# b/ d/ N2 S4 r' R算法诊断部分,建议把循环次数改为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
    ' p! d8 j' a9 H5 z0 \没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?; i! A- p- p2 l* r
    -------  K0 `% R. {) v: M) u" R' z; \
    不好意思, ...
    ) ~$ P1 E7 E( d  F) o0 x" z3 a
    谢谢,算法应该没问题,就是最简单的线性回归。
      z/ z+ u: A- G我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    1 J$ b4 j/ M6 O! M; V' j( O( Z
    雷达 发表于 2023-2-14 21:52
    6 q2 A/ O3 X% Y1 \谢谢,算法应该没问题,就是最简单的线性回归。5 v* N* _) ^" Z  t
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    2 U1 q, k4 l- M1 _) C, B, D

    : Z  }* a0 @+ _8 l5 m' e5 `' q刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。: X; X) U! k; j9 y6 a

    ) K% Q. N" e) J或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情
    怒
    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    # k- E; X5 `4 F" `( S6 C! ]  n
    老福 发表于 2023-2-14 22:00
    9 y5 u6 g. L2 w刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。% m& O- G* ?2 p0 H6 d0 B) \* d

    , f+ Q  P. K9 B0 W. [, m& v或者把b但的起点改为1试试。 ...

    0 ^' R9 J" y% N+ Z9 c& s& m* w' j! t: l0 b, X* a+ J% Y2 o
    你是对的。, g: o& i1 ~: D9 _, V% a, H3 B% B1 T
    去掉了随机部分8 u. l: m; n* a$ _
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)
    5 J# _$ D  c4 D( p( Qy = (x*27+15).reshape(-1)# U, `* x: `7 _/ e1 e+ ^
    * F9 S# P/ n( p, ~& V1 G( P
    循环次数加成10倍,就看到 b 收敛了
    7 {- G: |$ v( W# hw , b
    2 \/ m7 F; ~" E4 |  t& ~+ p( Q27.002620697021484 14.826167106628418
    : l8 M" {& x9 |) i0 e5 c: _5 w$ \& h9 E$ v/ s% |- t6 I4 ^( f, ~& M
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-30 12:37 , Processed in 0.057732 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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