设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 4 |7 z, J7 o1 D1 q/ h, F

    1 T0 K. V, d7 G6 [9 N4 x( b6 Y7 _. [为预防老年痴呆,时不时学点新东东玩一玩。  B+ i* l: L# r# d7 N; t" S  X
    Pytorch 下面的代码做最简单的一元线性回归:
    & L0 j/ U7 @& n----------------------------------------------
    5 _; @& o/ ~. Dimport torch( O3 [% Y6 |0 C& |
    import numpy as np
    ! c( R, ?& ~% G  m& g4 h, [6 Fimport matplotlib.pyplot as plt
    ' h; s8 q, B7 |' q% Cimport random
    : Q" k: x. P3 K8 d3 o0 [
      E6 r8 o) c0 W) z" }+ G5 R* F# Lx = torch.tensor(np.arange(1,100,1))' W6 d  v5 v" |0 a- ]# z+ W
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    . T1 @# l) H- l; a0 \/ u) ^7 e
    - B0 X* K) y8 w( e" \w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b3 k( u* A. Q5 Q# q
    b = torch.tensor(0.,requires_grad=True)
    / K/ C5 w9 `4 D1 w$ U& `( b; p' D2 q, K
    epochs = 100
    1 p& F4 J( M; S4 x- e! o% L
    : A  \9 `. U7 K3 }/ wlosses = []$ V1 B* @2 E. V4 F
    for i in range(epochs):9 a4 T- u0 @1 r  s
      y_pred = (x*w+b)    # 预测  r4 {1 h2 E. N6 p/ r7 o+ h' T
      y_pred.reshape(-1)
    6 N0 g/ `! y8 T1 I. Q. U ; |; k* C* O0 c$ Q5 g
      loss = torch.square(y_pred - y).mean()   #计算 loss' a- W4 Z- I5 K) C- L( ?
      losses.append(loss)
    . Q% |; n0 ~8 L- V" C3 ^4 I  
    , J' W+ g) R, j, T  loss.backward() # autograd
    ) i" N; ]6 A: E3 q7 x  @  with torch.no_grad():
    8 B  o% W9 Z; J5 o: o4 ]8 u) H    w  -= w.grad*0.0001   # 回归 w
    % ~9 @: ~3 D2 g1 T0 Q8 L: j: L    b  -= b.grad*0.0001    # 回归 b ' G% D* F$ k7 O. x1 ]
      w.grad.zero_()  * y; n( i1 v; s5 s- p6 y
      b.grad.zero_()8 H& m2 |6 J, b

    / O( [" M5 b  ~5 Sprint(w.item(),b.item()) #结果
    . X0 _' m+ @' ]6 L/ F) J" z
    ! G6 ]+ K& }# Y! w3 j. Q* GOutput: 27.26387596130371  0.4974517822265625
    8 m+ T9 `  B% N) r5 Y0 _. d7 g----------------------------------------------9 t% {( H  ^% G6 I
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    - o, t7 f% j# R+ \. {! [高手们帮看看是神马原因?
    * z$ d! o. {0 \1 y

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    9 c# H! W5 n" l9 m6 q- r" C" z8 ]* F
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    * x0 Q0 A6 v" {" n. k3 h-------
    * S) M, e, n: D1 Z不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。$ o  h5 Y7 }/ Y2 ]) W# }$ s
    -------0 }( }! D' Z# 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
    " I" E7 a2 @( k% h; Q没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?9 K; c$ d) ?: s. T8 G
    -------  h" ]& A) ?: q% X
    不好意思, ...

    $ e2 ~( C! U8 l8 z* Y! {谢谢,算法应该没问题,就是最简单的线性回归。5 K7 _0 z/ y+ p2 r2 Y
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    5 U! a, i) `9 D: x
    雷达 发表于 2023-2-14 21:52
    & m" S! L/ ^2 f- L; H- d3 ~1 N谢谢,算法应该没问题,就是最简单的线性回归。, \' i6 ?- U+ K
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    3 o9 P0 r2 r2 h% z8 M4 K' i( h+ W  v! @, f0 i+ |) J* d* z
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。% j* d/ n% }% Q) w2 e% z6 k( b
    9 z8 q8 @9 k4 C) @
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    & B, F/ P& [0 T! Y1 M: q
    老福 发表于 2023-2-14 22:00! [* x5 _" J' k) ^2 {
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。) l: n3 ?4 _2 o
      L: d6 Y" v( D6 ?$ v, Y1 [4 P* }
    或者把b但的起点改为1试试。 ...
    & ]" p4 s) Z5 M' W: E! b

    ( }* z- i$ I; H" \6 h你是对的。
    ! ~& @* V& {+ r$ I7 \去掉了随机部分
      P0 Q, s' Z5 T0 S#y = (x*27+15+random.randint(-2,3)).reshape(-1)/ O/ t1 P8 k1 Z
    y = (x*27+15).reshape(-1)+ t) B8 c  m! e2 z, Q8 I" E
    1 T; A5 f, I. Q3 Y2 A( g; L0 u! z8 e+ T
    循环次数加成10倍,就看到 b 收敛了: y, A- z4 u* l+ |! [2 E6 q7 A
    w , b
    . C1 `$ v& Y! F7 z/ K& {" [7 V27.002620697021484 14.826167106628418
    9 [* T1 `9 K3 {: T  v- E
    7 y9 }/ ?% B& n) V  Y; y4 b1 b和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-5 14:43 , Processed in 0.062643 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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