设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    : l% J  s# J" z' n! _( J7 O  g9 L, s7 h/ E  _
    为预防老年痴呆,时不时学点新东东玩一玩。7 H1 O  z+ R8 Q; d2 C
    Pytorch 下面的代码做最简单的一元线性回归:
    # I( J8 Z) U: V8 I----------------------------------------------
    8 B) u& D" r- y: U0 |1 n0 {import torch
    ' I4 Z2 e: x7 `3 G4 B9 o6 M% oimport numpy as np( C. K6 ^; K/ g- v" K
    import matplotlib.pyplot as plt
    * K0 ?0 u+ h4 Z% Pimport random6 j* Z: b- e; g8 A2 n$ d
    . F( w, T" d9 n
    x = torch.tensor(np.arange(1,100,1))
    : z! Y. K: e6 e% K# C6 H8 P* vy = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=159 s& v; R" Y0 _* P9 K

    ) p: T& d9 e( m  x; u& `2 Zw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
      |* o2 B2 G0 Wb = torch.tensor(0.,requires_grad=True)
    7 }" Z% r9 z+ S( ]+ _
    : m5 N9 z; m, W( T4 H2 Sepochs = 100- G0 e8 G7 a) [7 z
    - \/ V! \+ V. r. ^4 E6 ^
    losses = []; u5 X( t5 r: Y: `! l: v! d6 A
    for i in range(epochs):0 D7 \' ]# u# e" P
      y_pred = (x*w+b)    # 预测
      c7 [) s& C8 x) v  y_pred.reshape(-1)/ Q5 i: h8 E# G- k
    9 _# i( X7 y- Q# j7 _, d  E
      loss = torch.square(y_pred - y).mean()   #计算 loss
    4 f0 f" _/ ], z  losses.append(loss)
      {* k5 i* G. \  
    , G3 V& P$ G3 k% X( ]  loss.backward() # autograd. t* e, N7 I+ I2 I$ g0 c
      with torch.no_grad():
    $ s4 k2 ]- n, i' Q0 ^7 W; B9 e: v    w  -= w.grad*0.0001   # 回归 w+ e2 ?7 k  I7 k0 {! d
        b  -= b.grad*0.0001    # 回归 b 1 h9 U# A6 g  u2 q$ J
      w.grad.zero_()  
    8 ~, [$ V0 d4 x$ h1 j$ K  b.grad.zero_()% N3 T7 s1 ^+ W
    : L$ {- N) v0 m3 Q. b- I
    print(w.item(),b.item()) #结果. T; o; F; H7 L- i
    . o( `& e/ ?9 B5 g- B
    Output: 27.26387596130371  0.4974517822265625, m9 C  K5 c% D- b0 u% X" E  f. {
    ----------------------------------------------; c* |0 n. Y1 `2 U, B
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    8 B" Y* o" Y+ H# K高手们帮看看是神马原因?; A) Y  ^$ ~) Y; O8 m

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 . b# }# J! b5 ?

    * `' O- e% U2 l! O7 Y+ x; \6 ~( D没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?" \. E$ P5 w3 a' C/ D9 k
    -------
      e9 |% O  ^7 c& \不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    ' C8 x5 g. r+ K- D- ^) `-------
    - f# `% d/ X7 j  f' X1 C% Q算法诊断部分,建议把循环次数改为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
    / V' l1 c, ]* r& s, Q没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?* r- i7 }* A# X) x# n( O8 |
    -------
    * ^: F" r1 t" J7 W不好意思, ...
    $ A/ D  N) L$ A6 U+ [
    谢谢,算法应该没问题,就是最简单的线性回归。
    & F/ ?$ j3 I7 {) B: f1 u, H, J我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    # p! |1 r' ^, l1 A& i* S
    雷达 发表于 2023-2-14 21:52! q2 B" R6 d0 J$ Q* l1 E( `
    谢谢,算法应该没问题,就是最简单的线性回归。6 x$ j& g0 U3 p6 z6 {
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    + J# s; S$ ~1 Y) ~+ y# R8 D& C* Z6 ?0 K1 A( e1 J) M
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    ; @) x- O4 t# P8 Q7 Y! S- P& B
    & l1 ^' K: U9 B1 \% v或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    9 s3 A$ x) r7 m3 {* [& q: }
    老福 发表于 2023-2-14 22:009 D8 J" y/ ?3 d- M+ h
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。: q2 f$ u) \$ X$ d9 {, H9 G
    7 ]" M% F7 U) b; h8 j0 p0 E
    或者把b但的起点改为1试试。 ...

    & a7 F. k: X2 y7 m. Y1 N4 s0 ~# u3 C  k+ c+ r- h
    你是对的。  ^$ p: L' C; C7 `
    去掉了随机部分
    - g& l9 f0 ~+ K! o! d/ F! k#y = (x*27+15+random.randint(-2,3)).reshape(-1)1 O0 ~( `( S% K/ ]
    y = (x*27+15).reshape(-1)
    * @* K4 G" k+ {; S8 n6 |
    # {! z1 U" D/ e/ m- L) G循环次数加成10倍,就看到 b 收敛了- X; G4 @+ @+ P0 c, M4 ]
    w , b% \- e: \4 T& M9 d3 ~
    27.002620697021484 14.826167106628418
    ' F, e/ e3 T, I9 P$ Z
    5 z4 A' a) _' `% X# U1 o9 ]' Z和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-13 15:57 , Processed in 0.061767 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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