设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑
    4 s9 S) |6 h% Y' j' ~; g) ~. b5 Q  J: f* g+ _, g+ {2 [! Y
    为预防老年痴呆,时不时学点新东东玩一玩。: l. B, n( i! O9 ^  D
    Pytorch 下面的代码做最简单的一元线性回归:! \+ \! N* @3 s& }" s0 \/ K+ ]; |! S
    ----------------------------------------------
    6 ]7 ]4 \! |  N3 n  jimport torch9 K9 m  p: z( Q% m) v
    import numpy as np* {6 M* e/ ~' C# h, S
    import matplotlib.pyplot as plt+ i8 K6 b" f& B" Z
    import random
    ! E4 y  y9 j% b, m: g3 V* T  y/ d# K: _
    x = torch.tensor(np.arange(1,100,1))
    & o& @# E" w( Y! e0 Dy = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=154 Q& _. }, V2 O# ?. u

    5 i; j+ N9 j8 G% Ww = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b% L- u8 k- b8 D! V/ S
    b = torch.tensor(0.,requires_grad=True)" V- q* o# B* G! t: s5 D2 E

    . q; n# w4 }9 O/ X! Z2 Uepochs = 100
    ' h0 G  w6 v) s3 _8 z4 Q& }4 A: o" h0 \, i) I( H
    losses = []
    ) a- W0 B) S0 z9 ?1 Y2 F8 U3 n/ ?1 wfor i in range(epochs):4 H1 i7 _0 U5 T$ C1 F' k
      y_pred = (x*w+b)    # 预测8 J: A" C! J, V8 J
      y_pred.reshape(-1)
    3 Y& i. c. I+ r8 o" ~! w; ^% @( {$ j
    # w' U# M5 `0 K! ?- ~1 X+ H  loss = torch.square(y_pred - y).mean()   #计算 loss0 Z3 M" }7 r/ y, X; h4 h
      losses.append(loss)3 E9 r6 M3 I8 x; g" \8 g, u
      
    . ~8 c% J) o& D+ n  loss.backward() # autograd
    ) d! T( `  y1 n  with torch.no_grad():& c" Y! f1 D( Q& L) f% s+ B' y7 G
        w  -= w.grad*0.0001   # 回归 w" O' S: r7 m$ A: G
        b  -= b.grad*0.0001    # 回归 b
    7 y0 \* k" D* A$ s7 ?! `7 X  w.grad.zero_()  2 n. M3 y9 r0 z$ ~
      b.grad.zero_()
    0 [. b0 _6 j/ Y( B5 V( b7 ]* F! I, {/ }! V4 H
    print(w.item(),b.item()) #结果
    * j4 w- V# ]; ~! i) w2 _# q- i5 Z8 I. }7 g# S, A
    Output: 27.26387596130371  0.4974517822265625! w. ^( @& |+ x8 U
    ----------------------------------------------: e- K4 |. {9 W
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。
    - `3 _4 m" v2 n+ b6 }  ~# O( `高手们帮看看是神马原因?8 c% O7 O% d. m9 Y0 A

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑
      a0 r4 v1 m5 G' W
    " J% F. q' F3 b& @. T9 n, v; L. R没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?$ e5 X* T- g& F, l& Y2 r+ W5 l. m
    -------* ]7 J5 D9 T8 i" N
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。
    - n2 P5 j: p# c- ~-------8 S; p; z- E! |( E1 F3 ?- G. N' {7 B
    算法诊断部分,建议把循环次数改为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
    ; H% P4 d$ A/ _6 d8 {没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
    ! t' }8 ?2 ]  ?1 G9 ?+ q3 k  w-------  X; Q. N# l7 W4 `4 @
    不好意思, ...
    7 w% n6 D9 \0 R8 p; K+ U2 q
    谢谢,算法应该没问题,就是最简单的线性回归。/ n$ h: H, D7 Y; R9 w% J
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
    5 Z# E3 w1 k2 J* [" d
    雷达 发表于 2023-2-14 21:52
    + ^. m  Q0 F  Q/ _; A谢谢,算法应该没问题,就是最简单的线性回归。
    2 f" h  i* w" ]; ]' \. A我特意没有用现成的工具,就是想从最基本的地方深入理解 ...

    / a. X2 A  G9 n# u
    : a% U# v1 z9 }% I0 V# g8 g刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。: z- b9 f* r- z+ C. ~: y

      y3 ]2 R7 O- i8 }1 V或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    $ S8 n" y) \3 }- L, B7 M
    老福 发表于 2023-2-14 22:00
    9 P0 g3 j6 G' [刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。$ r  h0 K: Z! T; D4 V) S1 n

    2 w) S  s# y% a+ D或者把b但的起点改为1试试。 ...

    1 t) H0 h# u; ^: d
    ( p8 s& h) Q( ^0 V. ?$ G8 W你是对的。2 B+ F% q8 Q# J2 |; @& k7 t; j! _: P
    去掉了随机部分2 E  F1 S  y" Y( i, O
    #y = (x*27+15+random.randint(-2,3)).reshape(-1)- }5 S% B8 W; S' Y( c* M# e- {. l
    y = (x*27+15).reshape(-1)
    6 J# u# `9 c+ H* e  h7 x$ ~' X( H4 A, I' d  c1 e
    循环次数加成10倍,就看到 b 收敛了
    ) y7 x  t. S2 G8 rw , b
    / t/ |/ [, {  L27.002620697021484 14.8261671066284187 I. B- q/ x/ }; R; \# y0 v. P

    5 Q+ D) n0 n1 _$ ^& {! e1 a9 w和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-18 08:56 , Processed in 0.059291 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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