设为首页收藏本站

爱吱声

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

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

[复制链接]
  • TA的每日心情
    怒
    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 8 J  e" T0 ^3 |) k
    ( R$ z/ N% y8 p. }$ ~. J' i
    为预防老年痴呆,时不时学点新东东玩一玩。, L" f' z) }$ u0 i; W; t
    Pytorch 下面的代码做最简单的一元线性回归:
    : r6 g. z5 f) B5 }8 a----------------------------------------------
    & g2 V- C* V% N# m6 ?6 ]1 y6 Kimport torch3 }. J& x3 D5 s& i6 O
    import numpy as np# \3 z  v! S  x0 R" C
    import matplotlib.pyplot as plt
    : m  L+ }. ]& cimport random2 k; R  e2 J3 m( W: `  u* |

    ' ~3 W' ~# h- `5 fx = torch.tensor(np.arange(1,100,1))5 r1 l2 U$ g; q6 y9 m0 A/ _6 u# R
    y = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=153 l" B) m: g! S( D. k

    + J: j/ s) k' rw = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b: |$ {, ^* P/ g8 B8 O9 e
    b = torch.tensor(0.,requires_grad=True)
    ) K3 B( F- ^3 |# t" H. o
    5 n7 X( t, S, t/ N6 ]epochs = 100
    * P; m" X$ A- d8 I7 F# \! ^( t$ D* a
    losses = []
    % H' O% M7 x8 Q" a7 O9 m* q# Sfor i in range(epochs):
    % `2 y# P5 o/ H3 I; ?! d  y_pred = (x*w+b)    # 预测. L. Q9 z( V% N) S, Q
      y_pred.reshape(-1)9 T- I1 |' E! U, Q; K

    ) m6 m7 h  c' l  p  loss = torch.square(y_pred - y).mean()   #计算 loss
    8 S' T, {* U! h3 z; w  losses.append(loss)
    % a1 ^( U8 ^3 G0 j+ G  
    1 T$ u8 U+ g) i7 o& _3 C2 k  loss.backward() # autograd, R' F2 c% V7 {! C, ?
      with torch.no_grad():
    7 t' M1 |6 D" `) M" [7 F: _    w  -= w.grad*0.0001   # 回归 w
    3 f% e. L3 a; |/ {- C, c- w7 v1 e    b  -= b.grad*0.0001    # 回归 b : Q9 t2 w, d* ]3 Y& L+ ?
      w.grad.zero_()  / n  N1 d" Z! x) i& |
      b.grad.zero_()
    ! M& Q# n3 i0 V, e
    0 \" o- v. s8 J  t4 o$ o  iprint(w.item(),b.item()) #结果+ R9 i4 R. f; g& @/ H) C
    % B' t3 f$ V( B0 x; `! W9 {* v0 I* w
    Output: 27.26387596130371  0.4974517822265625
    - [+ B/ v9 G' t: h+ K2 n( O----------------------------------------------
    2 k) X2 |" S( a/ q最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    ( W0 U* W' ~. n: U( H/ V高手们帮看看是神马原因?
    4 u" {# @2 w0 \: b9 n

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 : N" ~8 R# n0 _9 E
    + N! U- g$ j; w
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?9 C7 v& F6 N; a7 r9 K6 d
    -------
    0 j% `0 P1 E! z# @! K: a: P& e, m不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。/ t4 N$ t8 s9 x& v0 x
    -------
    4 p8 Y" p9 h9 G; V算法诊断部分,建议把循环次数改为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:238 M% }; }) ?! @8 k2 ]& \
    没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?) t4 ~9 a7 S3 V3 b! L) ^% R/ \
    -------
    8 G2 r( M" W+ e; H" s* q不好意思, ...

    . Q* l/ E9 l( O" T谢谢,算法应该没问题,就是最简单的线性回归。
    $ m. f$ i: m# g2 _3 S- @9 M6 r1 z5 l/ k我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑 9 j+ k3 N4 b. ]5 D: K
    雷达 发表于 2023-2-14 21:52, C. C# Z! @& C6 T5 {
    谢谢,算法应该没问题,就是最简单的线性回归。
    4 f6 D. G; K* D+ r/ b* c& R我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    1 n8 D: f1 s7 A

    0 C2 s) _9 |) B# q& {, W/ U刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    $ N' z" U- S- A, P$ I. Y/ b
    , E/ ~; V0 r; S2 ]$ E+ i$ H: C或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情
    怒
    2025-9-22 22:19
  • 签到天数: 1183 天

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑 & x' K: `6 N$ y- W( N" P
    老福 发表于 2023-2-14 22:00- Q8 v0 {8 y1 z" ?
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。" U; V+ y  ^, f# C8 ?! c
    ) d# I8 [% E/ ^) X
    或者把b但的起点改为1试试。 ...

    4 O2 d  z0 c. ^4 m+ \+ b3 B9 y7 Q# ?
    你是对的。
    6 h2 t, |" I1 v0 `) y去掉了随机部分
    , N5 C6 i( }7 g#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    " x5 J& d2 n! t1 l, ^5 Ky = (x*27+15).reshape(-1)* O# j( U1 P; T. f

    - H3 `  a: L" o循环次数加成10倍,就看到 b 收敛了
    ' G5 z& ^5 f3 m7 Vw , b
    & C. W0 ~% ~2 J0 O, N( o* \3 G4 p27.002620697021484 14.826167106628418
      n/ q0 a0 Q: F+ g0 C- N- \0 D, Y. ^% _* x0 k& T+ D7 Q" T, o
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-9-27 18:39 , Processed in 0.103705 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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