设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    . b7 |% d4 D2 `- P0 B2 T9 Y. l% n' C, ~$ S
    为预防老年痴呆,时不时学点新东东玩一玩。, J# ]+ U& l3 X( m( U+ f0 n# O$ z
    Pytorch 下面的代码做最简单的一元线性回归:
    - V1 q& p+ a& l----------------------------------------------
    5 X' M$ U* V4 L5 p: o4 R( e- Qimport torch
    ( N8 U! Q+ c( ~/ ~9 o- |import numpy as np" w* c9 H* w8 U- y
    import matplotlib.pyplot as plt! d  X8 _' ^* S+ s  c- _
    import random
    6 ^( p' ?( L1 t# X2 k( M
    * c6 `7 N3 v+ b9 B1 ~x = torch.tensor(np.arange(1,100,1))
    5 e) r* S* s$ dy = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15
    . _& J8 M' O: N  ]8 n4 I1 Q) y% x. Y  d1 r/ x0 r, _+ z
    w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b  @2 u2 s: f5 I
    b = torch.tensor(0.,requires_grad=True)5 ]" [; d0 R5 p$ k9 d

    ! Y( H8 A& {' k. r- eepochs = 100
    : ~" P3 R% H& l5 N) d& d
    6 Q: K5 |% N5 d! D* B  b9 a' Ilosses = []4 d( q. g5 Y7 H6 D# D% u: M
    for i in range(epochs):4 ?& n7 @; t( H& N: l' n+ S
      y_pred = (x*w+b)    # 预测
    # B: B" c3 _1 P- O- r. o  y_pred.reshape(-1): G* R" a) F& R7 ^! {! e* i3 |
    6 U, s) {& P( O2 |' w# P
      loss = torch.square(y_pred - y).mean()   #计算 loss" Y0 O% T' ?9 a
      losses.append(loss)* G& z4 S4 b; @* p" Y
      7 S! z8 b( s! q! G: C
      loss.backward() # autograd
    % \! g7 }- c# V9 x- t. ^  with torch.no_grad():
    " ?# G/ |7 V5 e0 o# I% ?! F    w  -= w.grad*0.0001   # 回归 w: ~  i. o# U2 O; N4 I
        b  -= b.grad*0.0001    # 回归 b 6 I1 z* W  ~7 Y; g) h
      w.grad.zero_()  $ D( O/ x6 F1 V
      b.grad.zero_()
    1 x" ~% g# h* _% h; v, w
    % s* ^  N+ H6 d8 K3 fprint(w.item(),b.item()) #结果( m' H5 J/ _  u5 I
    / |8 Q1 r: Q$ Z4 p9 C/ a
    Output: 27.26387596130371  0.4974517822265625- `( U/ c  v; O  l6 E" z5 c
    ----------------------------------------------
    * F! Z' C8 l( ?+ L, H0 Q0 O% L! {最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    % w+ a4 s7 J5 [6 O高手们帮看看是神马原因?) g5 P+ e& W0 S+ ~

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
    0 O% t8 C: q3 U% x* A+ K% J! y# ~1 y- _$ i& ^. q+ z' Q! F: ?
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    2 q/ G* i1 v$ }- J-------7 L' E: X! e+ u) j0 }
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。/ ?  y$ r( O& q" m- C
    -------# \, ?1 A8 H( y
    算法诊断部分,建议把循环次数改为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
    # x* F( v) G# d- z4 s. j没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    2 _) q9 h  x# a$ Z+ P9 G-------
    1 n1 G5 q0 N* T" D! m不好意思, ...
    7 Y) E* }6 t3 L& ]5 q
    谢谢,算法应该没问题,就是最简单的线性回归。7 l: X( t6 U, S, ]" F. |- j
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    ( x+ |) ^: [/ R# f
    雷达 发表于 2023-2-14 21:52
    0 M: G1 ]( H1 }7 A4 ~谢谢,算法应该没问题,就是最简单的线性回归。  F3 {& J2 V1 g$ [9 y
    我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    - k8 x2 D) X) E3 x; R4 o7 N

    2 \. p5 c0 _( Q$ a. V刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。7 A: K* d5 b  o- z
    # S) i" V+ x' r8 k& I
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    , Z* G( G/ I; `. N) j
    老福 发表于 2023-2-14 22:00
    9 |3 G/ m# l, u% h6 @刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。1 H( k' i- V1 d* b3 t2 Z
    7 g$ [1 r$ J) m! t' L' c
    或者把b但的起点改为1试试。 ...

    4 O: ~1 d3 b" Q- _$ ]7 Y( Y  l' }: J
    你是对的。0 k# m, e* y7 H; i  p% q
    去掉了随机部分
    ; @: i2 X! H! U( B#y = (x*27+15+random.randint(-2,3)).reshape(-1)+ _8 d  k; g4 H5 e' J
    y = (x*27+15).reshape(-1)
    $ u% {, U) W) m% \* A/ y
    1 L1 i; r# w+ D! h9 _: ^' M循环次数加成10倍,就看到 b 收敛了# M' l5 V% Z' c5 U$ t0 m5 ?  x
    w , b
    ; P" n7 C5 `0 D( H, Q; t5 A& Z27.002620697021484 14.826167106628418
    / g( q4 K6 ^0 i! R( m7 F$ F5 W5 n4 ~. E* [. I& b* n; G# h
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-1 03:40 , Processed in 0.059414 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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