设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 # j4 ?9 u- T6 O

    2 D' X0 ]9 X* q为预防老年痴呆,时不时学点新东东玩一玩。  X4 }6 o( n5 V; o7 I! f
    Pytorch 下面的代码做最简单的一元线性回归:
      d  d: z0 c% ]1 z3 U----------------------------------------------! E: J0 ^2 k; o, I5 @+ ?# d: b
    import torch
    ' U3 s4 d6 X4 vimport numpy as np+ n' g3 F+ P6 {+ E9 c9 v* o
    import matplotlib.pyplot as plt
    * y" E7 B( e- N! fimport random) {: z$ n, k$ c" a* J7 K3 u1 T

    / G. E0 r8 v  O9 [4 O3 w1 Ux = torch.tensor(np.arange(1,100,1))2 \% S1 L6 I8 |; P1 e
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=151 N( X% }- C+ ]3 ]* E* z; v4 d
      h4 J. ?. k/ B! Z+ Q; ~0 E  p6 R
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b
      D# R" h) N9 i- q' fb = torch.tensor(0.,requires_grad=True)+ ~8 M0 r( \' e( k- a  R) \/ k

    ; ]  [, Z( e& w7 m! W+ Mepochs = 100$ b% p; g. _2 {& J4 l' w, V
    ! {% |/ {6 S! h0 I9 m7 n
    losses = []3 D" o' b' v. W) e9 r; Y7 s! p$ d5 e
    for i in range(epochs):3 @' r* ?8 m' ]' p
      y_pred = (x*w+b)    # 预测
    4 l1 _8 c3 g# [, h  y_pred.reshape(-1)
    3 s, U+ B( ?' C: S/ v0 X % y- ~1 H. ^9 C3 G
      loss = torch.square(y_pred - y).mean()   #计算 loss
    ! L0 g/ {1 q9 ^8 ^; K5 O4 w1 P) r  losses.append(loss)
      K: [. G' v" c: V  + a7 h0 ^! D4 s$ G
      loss.backward() # autograd7 y6 a3 k: U* E2 g& R/ K7 }6 @
      with torch.no_grad():, E; Q' A- _6 _% p" g9 t
        w  -= w.grad*0.0001   # 回归 w& F$ v. X  w) ^
        b  -= b.grad*0.0001    # 回归 b 0 j/ D3 L% q7 B: _* A. Z6 t8 T
      w.grad.zero_()  
      u* r& g% K- @0 x7 n  b.grad.zero_()$ V0 u5 D5 X& w. T

    + K8 l" A) w- A7 S0 vprint(w.item(),b.item()) #结果
    9 R2 K# J- R, c
    5 E3 S8 P4 M2 \# w# }7 OOutput: 27.26387596130371  0.4974517822265625
    " n; A2 @3 J, ^5 H2 k! A% }6 G----------------------------------------------3 Q* M# n9 @5 c
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
      j) b! g- k# S1 v高手们帮看看是神马原因?
    6 g% ~$ p* R; r! g$ `

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 # w: }  h0 {9 N( c4 U5 _

    + i2 H7 T. t  j没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?6 e; C& N1 b% D) Y6 W1 X' R  |
    -------, W$ K: R8 t. S: U8 ?
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    " O! l5 F4 k7 M8 y-------
      Z+ B, A& e# U; p! Z算法诊断部分,建议把循环次数改为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:232 D. N# a$ B& T" d1 K4 L, D
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    : I+ _$ n5 ?9 {+ f' {# a8 I0 v-------
    2 s% i- }  \% A' }不好意思, ...
    5 |8 p+ k( y( h5 V2 L
    谢谢,算法应该没问题,就是最简单的线性回归。; a) [2 T! x% e) d/ d
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    4 {$ q( y1 g3 L
    雷达 发表于 2023-2-14 21:526 r  G+ N5 @& E7 H7 A0 X
    谢谢,算法应该没问题,就是最简单的线性回归。: j% _% a% I" ?& p# U
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    8 R2 s0 t  |/ @. |% i& l
    0 P$ S5 O5 r5 a  V! I9 t. w
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。4 m0 E/ j5 p  {
    9 E" T3 g0 b# O* P6 ^9 k
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 % E: l% f' @. b+ E, a+ B% A
    老福 发表于 2023-2-14 22:00
    ! D0 B0 m2 X, m/ G刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    2 s+ a# t2 c: \
    , Z$ ~) u' n! ~9 V  `* \或者把b但的起点改为1试试。 ...

    . T9 U9 `# P" g* c
    1 J4 A, H. J' z0 B0 ~1 m$ [8 \你是对的。
    ) N. z* W) S& I去掉了随机部分* w$ L! E2 a: m- B; S( T
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)8 a% t- h0 P2 ]6 a8 m+ N
    y = (x*27+15).reshape(-1)
    + t! \- v8 b+ O4 m9 C' P- ^! u0 V* g) l. B/ D
    循环次数加成10倍,就看到 b 收敛了
    0 M. u# E; k5 s. m8 ~w , b
    2 w: Z6 @" e3 @* }4 @27.002620697021484 14.826167106628418
    , y& n1 Z: w1 O2 C- R5 G0 c/ B2 ~, w& Y% S
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-23 22:02 , Processed in 0.058148 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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