设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    : m- ]) ]7 G' W2 C
    , h4 H& [1 Z6 M0 `, r! O% {为预防老年痴呆,时不时学点新东东玩一玩。4 M# x; V/ e2 `) f+ W* M
    Pytorch 下面的代码做最简单的一元线性回归:
    * k; _2 |6 K5 I$ t----------------------------------------------8 z' U" A, j* p2 V) D" a8 L
    import torch
    7 e" g6 }; x& l: _import numpy as np7 W3 L; h" C8 i& M  {& n$ _' N
    import matplotlib.pyplot as plt
    8 I' n# z4 h6 O! j: Nimport random* x% g6 q8 @2 r7 M/ D
    7 N+ X" i1 F, h; a
    x = torch.tensor(np.arange(1,100,1))' J8 C+ y3 B7 r# [; C
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    0 g0 a8 j: o; g. b9 h
      t3 k6 V4 w4 i* N+ s4 Lw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
    6 m) Y( l' |$ ]0 [  o) z. {" Qb = torch.tensor(0.,requires_grad=True)& p* G$ S* l$ y

    # H7 j+ K& M! p8 C' N4 J: N/ S/ \4 Jepochs = 1004 m) b  R+ ?, F- T
    4 \7 \, }- M  T
    losses = []2 U- M2 W8 w. H- I* x8 B; K3 J5 T5 p
    for i in range(epochs):
    : j; C) m% _  `  y_pred = (x*w+b)    # 预测
    9 W; D$ j; W* i  y_pred.reshape(-1)$ e; w! n3 t" b

    + v! w, d% X/ A0 m. n& o  loss = torch.square(y_pred - y).mean()   #计算 loss7 d: |0 V. j. |
      losses.append(loss)# ~. D7 h) C) G3 f) T( j
      
    ! ?3 q4 p; g* L9 T7 p  loss.backward() # autograd( u; R# [1 T( E# |4 E
      with torch.no_grad():) y! [: c' o' f" A9 s% B  F7 o) {% E
        w  -= w.grad*0.0001   # 回归 w8 z" J  t# E- P5 H9 g
        b  -= b.grad*0.0001    # 回归 b ! ^# D$ M4 I6 g7 k+ L
      w.grad.zero_()  3 A# @3 {) W: Q5 k4 ^) M  |" `
      b.grad.zero_()
      W" A+ ^2 k& g. r
    0 G9 E# i% b, W5 J3 n1 X) zprint(w.item(),b.item()) #结果
    ) b, i6 ]/ [. o) @4 e* F4 s( ?& I' L, h* m$ {
    Output: 27.26387596130371  0.4974517822265625
    2 J) S8 _, @9 R( m2 T----------------------------------------------1 X. G) g; u( c) f# V7 H
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。2 h6 f# r$ T3 |/ K, ~
    高手们帮看看是神马原因?
    " ^7 Y2 x9 `9 J+ |9 f

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 " @9 j9 b# `$ z! G

    % m7 ?8 G/ ?& ^/ w: i+ X) W没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?$ @& w8 k: [% }9 k/ G$ [
    -------$ K+ [  A+ c/ {% y7 U$ ?
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。8 J  m% x6 I7 r# W2 g+ R+ C/ X
    -------
    ( t2 ?. _, V2 T" h% e算法诊断部分,建议把循环次数改为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; Y9 I- e  P! ~) B, b, l: }
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?! Q) E  C+ ?) Z
    -------9 w) }  a5 O& {8 u+ n1 S
    不好意思, ...
    ; }  I0 l/ C) G+ R
    谢谢,算法应该没问题,就是最简单的线性回归。% O; E% C3 l- D6 y! |
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 2 J7 W3 i0 @0 l, v2 p8 K5 K( [1 b" g3 y$ j
    雷达 发表于 2023-2-14 21:528 i0 U3 M2 I5 x, K1 o
    谢谢,算法应该没问题,就是最简单的线性回归。
    % R( G7 y9 ^% U  E* o' h8 i我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    1 D* |7 [  I8 l) R0 @0 [
    1 b# a  [4 t0 R
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。" L6 ^" @6 z& y: r: l

    ' p$ K1 r) i  _或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 0 o6 u) J$ }5 X9 U* ?9 ]/ t
    老福 发表于 2023-2-14 22:00
    4 [2 z- P. O1 W* ~刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。, V# Y+ f; Z6 T2 L& t+ z+ ^1 U

    , G" \/ j( Y- P% H0 V或者把b但的起点改为1试试。 ...

    1 _6 m% b! h* g6 @: W1 J
    3 V0 p4 i+ n' d( |你是对的。: n8 c3 Y, J) Z3 w6 Q* z' T
    去掉了随机部分% u( A1 k$ a) t5 U& W
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)
    - }3 `! u9 k1 m6 j9 c4 Ny = (x*27+15).reshape(-1)
    % D( m; ^: L" N# x5 r$ h7 L3 j/ m& s  Z9 @6 ]9 A2 ~; o6 g* n9 f
    循环次数加成10倍,就看到 b 收敛了
    . m  E. z0 I, ]; Aw , b
    9 b& x8 y' K# m! x27.002620697021484 14.826167106628418
    ) I7 O0 b  U' _  n8 j2 o! L8 H% c  a# t0 `& U- n* y' a- Q$ Q, M
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-4 16:47 , Processed in 0.060127 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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