设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
      u; b" E" j" W( E
    " Z1 }  A, Z  |& O& P6 A为预防老年痴呆,时不时学点新东东玩一玩。
    1 u; ?$ f" ^# A; VPytorch 下面的代码做最简单的一元线性回归:: s( t4 a/ [( c0 t( U
    ----------------------------------------------% c" n+ o3 k2 e
    import torch/ R- O' m2 N6 V1 @
    import numpy as np
    6 B/ S! C4 Z# {/ Ximport matplotlib.pyplot as plt7 ~' Q  g: v6 n0 Z- a6 b% a% S0 w
    import random  J/ z5 K. n" ?. x/ K; V

    ) A9 @5 R# e8 x, u% ?9 k! px = torch.tensor(np.arange(1,100,1))! E1 K) Z* W2 V" @/ a! `
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15+ M: s1 q8 J! T  l4 J; ~
    8 C  w: X" [; m1 S- T4 A( I) i0 W
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    & K6 w* m% I. Z& x7 a0 Vb = torch.tensor(0.,requires_grad=True)
    ( _9 [- k( s& p. o& e' J
    0 y* P7 j8 A/ Tepochs = 100* H# u6 X/ Z; ?5 x. ^" i' C

    & w$ N, y6 q: C8 _: Flosses = []
    / f' f7 d; D5 v& ?for i in range(epochs):
    + ]' O7 }4 ], \3 _  y_pred = (x*w+b)    # 预测
    8 n* A6 F" `+ o/ E! d6 H4 o* i$ ?  y_pred.reshape(-1)3 n. X$ k5 {0 {# k9 k' j% J1 r0 A

    ) g5 D$ T4 z6 ]8 p" \2 |  loss = torch.square(y_pred - y).mean()   #计算 loss: c9 d0 X7 H& J  \$ ?6 A* e$ b
      losses.append(loss)3 g' v4 A; r$ }
      
    2 d0 E' b! G3 L; J- o( v/ e  loss.backward() # autograd- K0 x. j  }/ |4 P% ]0 x
      with torch.no_grad():
    ! i4 v8 E  h  m2 @$ b    w  -= w.grad*0.0001   # 回归 w' @0 v3 p, \( b8 t8 E* M
        b  -= b.grad*0.0001    # 回归 b ; F4 h* U4 B1 B6 ~
      w.grad.zero_()  
    / k0 z. p  M/ S' \' s/ I  b.grad.zero_()( [/ _: w7 L& ?
    + ]6 z1 i; k$ P
    print(w.item(),b.item()) #结果7 a' U% p8 [. _2 A! M8 G

    ( i8 D; ?# r0 `) [# B4 v/ n9 n1 ^5 @Output: 27.26387596130371  0.49745178222656252 Y! p; q4 n1 U1 ]- U  u$ A. R
    ----------------------------------------------3 Q. z5 ]3 {6 b- ?% d2 ]; m
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    / |. e. Q2 E4 c. s高手们帮看看是神马原因?3 F% P5 q7 R/ V: y

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
      E) y" {4 `& S1 H( h
    1 k9 L+ b- K$ n0 C/ G  o没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?! B/ r( @; u; A
    -------
    ) |# {2 ~4 I- Q* D: o& P# ]不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。! q4 W! s3 v/ ?) e  g
    -------; i; J  }+ C( u8 _. g
    算法诊断部分,建议把循环次数改为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" Z, O0 d9 |0 r' l7 }/ @
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?. U, t# I+ _& e- a& s
    -------3 T/ a3 ^  v6 ]
    不好意思, ...

    # n7 [3 H* C  E: E6 q谢谢,算法应该没问题,就是最简单的线性回归。
    + n  A+ g7 G# F  n/ p4 J我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 % T$ e& H  {& T/ B: _
    雷达 发表于 2023-2-14 21:52
    * c3 k6 e$ Q: P, b+ X+ P谢谢,算法应该没问题,就是最简单的线性回归。
    # p9 G9 W' `5 h* j* }我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    ; r" Z" Z  ~$ W4 H) z
    5 t0 y( c8 Q( W$ x) b% Y4 ?
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    1 D4 h4 i: p  l8 X( X/ X1 d+ |* V+ \! c- A# z; O1 a
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 $ A6 ^+ i8 u: ~8 W' d7 I3 p
    老福 发表于 2023-2-14 22:00
    0 j3 K9 E% X$ f4 U7 f刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。+ |3 v5 |2 S. O8 o  Y! e, j, c* W
    / L1 o3 ^3 W. `& _/ L  k6 B
    或者把b但的起点改为1试试。 ...
    ( z4 a& O4 |, k$ j: a% i

    0 {. a+ T1 M( W% {你是对的。
    0 t. ]  r: v7 v+ c; g. M去掉了随机部分
    ) {' w2 d2 _3 G9 W9 I/ |& Q. \( |#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    % e3 F! N2 ]2 ~. f$ by = (x*27+15).reshape(-1). `' Y# C' V. Q( B0 K
    * f( {% I0 ^; D
    循环次数加成10倍,就看到 b 收敛了: i; q: ^$ @) \8 s" m6 M: D
    w , b) U0 q* R2 w1 ]
    27.002620697021484 14.826167106628418
    - u" h% Y+ e; i8 e* n7 W# w  s
    8 l( V9 t% v+ Y- q. b和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-17 10:29 , Processed in 0.072662 second(s), 22 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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