设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 8 l, ~# X9 {0 Q; d
    $ b/ F7 h: I/ Y
    为预防老年痴呆,时不时学点新东东玩一玩。2 x2 p2 ?8 j8 `. e/ w
    Pytorch 下面的代码做最简单的一元线性回归:& o; }) k/ d7 S; [4 p
    ----------------------------------------------
    1 {. Q- X) m2 d' r- m2 ?- Himport torch
    2 R, i' m+ Y7 E) eimport numpy as np
    + W" y( L  R; b3 ~import matplotlib.pyplot as plt
    0 Q" @& }$ R7 @# \! rimport random
    1 D# m: i( U& P' L- W8 E4 f7 D/ B& t% W% V
    x = torch.tensor(np.arange(1,100,1))* h" v# A/ l& H
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    " {+ b3 }" N% Z& D* t! T! }
    ) r- u( N- a' ~8 G  {w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b% W5 D& W" O3 |8 ^0 f% G
    b = torch.tensor(0.,requires_grad=True), q' s3 b  ?, \- k# Y% \
    0 d- g) \: z4 w/ Y+ F
    epochs = 100
    ' a# i) \3 v! F$ u# `* [% i' A2 ^* z& I; ^
    losses = []( ]' U5 O- x& R5 C
    for i in range(epochs):$ G4 a$ G; \# b! n" P
      y_pred = (x*w+b)    # 预测8 C6 P( [: q& Q
      y_pred.reshape(-1)
    4 C5 l2 g( ]* S+ k: j/ p
    + ^6 ?! D. m/ G% e/ j  loss = torch.square(y_pred - y).mean()   #计算 loss8 p( |, s6 F8 w, I" j& u2 ~) u
      losses.append(loss)
    & s2 f$ A% E; G, z( Y- i  
    7 ~9 N. z8 U, Z' k6 y6 s  loss.backward() # autograd
    6 i9 [* H: N& f  with torch.no_grad():
    1 \+ R$ D2 h3 Y% _; t3 @    w  -= w.grad*0.0001   # 回归 w8 @1 m" l6 W: {5 h
        b  -= b.grad*0.0001    # 回归 b
    , {% k  F0 q& w0 s; R  w.grad.zero_()  
    " Y1 S1 k/ A8 V" [. W$ G1 y$ P  b.grad.zero_()
    9 _: S( x& b" m, A8 w6 G7 J5 S% v# [" u  N
    print(w.item(),b.item()) #结果
    + `& l' H, G" H9 n" g  C2 _: Q/ {9 v/ n6 t& O
    Output: 27.26387596130371  0.49745178222656250 c! g3 m- [- c5 L
    ----------------------------------------------& b- |& |/ K  y
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    + K% \' S6 K! ]$ N6 d$ q高手们帮看看是神马原因?2 E3 d. S- B& _* ~7 s& R* h" Q( l

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    6 G% s% e8 Z' C- b
    6 Z* }: X5 L) r$ o. y7 S+ o3 q, i没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    / c: `! r3 V: ^-------
    5 ^; ^- h9 z: s6 {- a. n/ h: k不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。) H7 n' H  [- C; T
    -------' ~6 Y8 ?0 ~* _  s7 q- r* }; 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  v7 o/ j( F% U9 O/ i0 z+ e  E# n
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?9 }/ o9 m  G+ @
    -------9 g6 A1 @6 I- ]1 C% P8 d3 y
    不好意思, ...
    - b3 \6 J4 X; f
    谢谢,算法应该没问题,就是最简单的线性回归。  n) y) C( R8 o6 }. c! B
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    % t, ]6 T9 \  Y; T' c1 Y
    雷达 发表于 2023-2-14 21:52
    / {3 x1 U$ z2 X0 p3 f5 c( |谢谢,算法应该没问题,就是最简单的线性回归。, m) o% x' _& g1 ~$ w! K
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    ) g4 ]+ {  L  L  _5 p( e- [+ a0 ]
    0 z- N  |. {) a+ U4 ?刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    ' }  d0 {- J7 P* @) R1 _
    : y; ?/ }9 N2 V或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 3 G+ r; b. T1 z4 d  P& z  P8 M
    老福 发表于 2023-2-14 22:004 H. w$ c5 X7 Q$ j* |
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。, ^6 j9 _! U0 d' C$ R7 l9 p

    * A0 f. w1 t) y4 U1 i4 Y或者把b但的起点改为1试试。 ...
    4 [5 g3 o1 p# @; ~# r$ n) [

    , H# ?# E; u8 @# X你是对的。+ z- L( c$ t- v
    去掉了随机部分
    1 I: f( y* E  ]: n8 S1 b; g#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    # }" o! j) T; O5 Z9 |y = (x*27+15).reshape(-1)
    . v; D! p$ W/ d( F. T
    5 f# D3 W7 q& {* P2 \; [) X# u+ ]循环次数加成10倍,就看到 b 收敛了
    ( o: X# _: R. B* R0 c9 L/ L3 @w , b
    7 E, D* I4 [# g  G/ @27.002620697021484 14.826167106628418
    . v* f3 r$ ^4 v( k3 U% X, P( g8 F2 V  }3 H' w9 Y- u6 U7 @
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

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

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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