设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    ; R7 T1 b& _/ K& l
    0 D( N1 c' g0 u6 Y) O* d为预防老年痴呆,时不时学点新东东玩一玩。
    " ?/ e* a/ x; SPytorch 下面的代码做最简单的一元线性回归:
    / F+ |% I; I* r----------------------------------------------
    ( U7 ?% M  m) A. Bimport torch
    ) w6 ~$ q* X: ]$ Yimport numpy as np
    * h  z- C0 k, K) _, |/ K* o1 }import matplotlib.pyplot as plt& @( v1 @) }& u/ }. l# M  t
    import random, ?4 k0 X7 K9 g  @+ a1 C
    , e0 K) \) G) p+ h  X
    x = torch.tensor(np.arange(1,100,1))' {& }  ^$ H" u- Z' S; O! z
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=156 S  n, J/ `& {2 Z/ ^8 _: y7 i
    , \( a3 P  ]; j. @3 j, R$ ^
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b* I6 x; O3 h8 N6 d4 J
    b = torch.tensor(0.,requires_grad=True)
    ! o( c# q  V6 Y, D
    5 \5 z+ b5 G4 E- X/ S: Yepochs = 100
    3 t: k; I# w0 h9 j0 h: L1 r8 X8 Y" h7 ]  D- P
    losses = []+ W% v' I: B% M- Z3 X( {3 l
    for i in range(epochs):
    3 J! V' p6 A, _% n  y_pred = (x*w+b)    # 预测
    : h2 t5 H3 ]9 G* F8 p  y_pred.reshape(-1)* v8 {) ^1 _& o2 L6 D. G
    , y3 i" J5 b3 O/ Y
      loss = torch.square(y_pred - y).mean()   #计算 loss
    + M) X1 m! _! q) C6 n  losses.append(loss)
    ! \, ~& \+ H0 h: y  1 d& G3 }% y/ [) R
      loss.backward() # autograd6 e& L: v  D/ i5 S, O/ w' J
      with torch.no_grad():
    2 W: p' X- I& Y& C    w  -= w.grad*0.0001   # 回归 w9 L/ L5 J9 K" j& W
        b  -= b.grad*0.0001    # 回归 b & T% }2 C8 u6 {( D. I; e& P- h
      w.grad.zero_()  
    * `) t5 E: W9 }* L! ^( ~  b.grad.zero_()
    , A* I# Z& O7 G7 t8 c4 I7 C$ J
    : ~: P' V! i9 S. N- G4 kprint(w.item(),b.item()) #结果
    # W2 q) Q2 w8 q' s1 N4 I
    " U2 C5 }9 [0 Q! V8 x, Q8 K0 mOutput: 27.26387596130371  0.4974517822265625
    1 N7 c3 k. w' @! F----------------------------------------------, h- N" C. x3 \! o, D8 F) q
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    : v, V) K' Y; R: [8 Y7 v高手们帮看看是神马原因?+ w0 z* Y8 t* I

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    2 c* M7 ]: l( E* `) Y) e8 q6 L& ^0 p1 H% W* u1 w
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    3 R6 x$ y9 q) _4 a6 n/ R' \-------) j5 u) z4 D- X! [4 T# G9 R5 d- b
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    & m; c5 Q8 M+ ]- S) ]$ K- Y-------% K( s; v/ T& h! ?: V+ `5 o
    算法诊断部分,建议把循环次数改为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$ c7 O& `8 p$ O- M, D' O0 B
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    $ B! g9 n' o0 V/ [4 _) b& |-------# b' _& {& i# M- p7 f3 ~# ^" m
    不好意思, ...

    . x' ?1 z& b$ c2 C谢谢,算法应该没问题,就是最简单的线性回归。6 n$ w! \& o* Q6 L6 t) [/ w
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 4 M3 p7 z, T  \8 r9 {1 ]2 w; |4 [
    雷达 发表于 2023-2-14 21:52
    ) d% D2 z: k7 j" K& C7 F; Y谢谢,算法应该没问题,就是最简单的线性回归。
    7 ^, |) ^  V' ]6 O, |) z, [" q我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    9 O6 `: M& m- p

    2 m: M7 f  W0 B4 p: y7 e9 \刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    4 [% [0 W2 P# a5 K+ C) @# ^. c
    6 p3 w2 }5 L1 T, S) b$ s% M- n. n或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    ' ^& N" }9 e. R4 ]
    老福 发表于 2023-2-14 22:005 j8 T: n1 X! D  e0 u
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    * G4 M$ t- H  k, e$ x
    ) @' H0 A# i9 c! L或者把b但的起点改为1试试。 ...
    / b* c3 K5 G4 C
    + P" g, l- ~3 u8 {1 L
    你是对的。$ r, }" H  H# V( M  u$ V
    去掉了随机部分
    . R: u7 D" G: [: [8 c( ~- k#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    ' M& p1 {) y0 zy = (x*27+15).reshape(-1)
    ' A- i' Z5 c% \, j! m- R
    , J* ~$ E9 g" x$ g循环次数加成10倍,就看到 b 收敛了* e8 T& v+ P, C5 o$ C7 o
    w , b% u* R3 U$ g0 ]; q6 Z
    27.002620697021484 14.826167106628418
    ; X/ c8 h( m: M2 Z# h/ d- _/ u+ q
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-11 05:50 , Processed in 0.063673 second(s), 22 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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