设为首页收藏本站

爱吱声

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

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

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

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

    [LV.10]大乘

    跳转到指定楼层
    楼主
     楼主| 发表于 2023-2-14 13:09:28 | 只看该作者 回帖奖励 |倒序浏览 |阅读模式
    本帖最后由 雷达 于 2023-2-14 13:12 编辑 9 M: X1 K4 Z5 C: v4 f

    ) N( T: Z- G! M4 A" C, Z( K: j/ A为预防老年痴呆,时不时学点新东东玩一玩。# R3 y; r9 U7 y
    Pytorch 下面的代码做最简单的一元线性回归:
    9 r* m7 N: \4 a----------------------------------------------( T0 \( @- |; r* s: o9 `# o
    import torch
    - H( [% u; b7 ?/ H% ?1 m/ wimport numpy as np
    8 s/ J$ b$ k# w, Wimport matplotlib.pyplot as plt
    3 Y/ a9 z- d. aimport random
    * m+ E6 \& \& K0 h
    & \' ?' u. D; S0 |! j. }x = torch.tensor(np.arange(1,100,1))
    : |- q& _% ?; {1 s8 K5 Wy = (x*27+15+random.randint(-2,3)).reshape(-1)  # y=wx+b, 真实的w0 =27, b0=15$ J0 S% @9 M2 n+ ]  z

    3 P0 |8 G+ u, @w = torch.tensor(0.,requires_grad=True)  #设置随机初始 w,b& f5 ?: q- R6 G0 {+ K- T- @; ~
    b = torch.tensor(0.,requires_grad=True)
    " _9 }# t" G9 |8 Q2 W0 _$ ]" y, q# K6 d
    epochs = 100
    / H6 V5 S- [' A9 T# a) C, T" G1 |, ?$ c' s% U  V' x
    losses = []! A, z4 R0 G( ?8 j, M( V1 w
    for i in range(epochs):
    $ ^& @- i! i% L8 v1 Y- Z  _2 k1 w  y_pred = (x*w+b)    # 预测  m9 u+ w' K+ d) C) D; v# c1 `3 C
      y_pred.reshape(-1)% k; U* \5 H% {- m: Z2 o# A( D( R
    ( Q+ U% B* n$ G3 K: Z
      loss = torch.square(y_pred - y).mean()   #计算 loss
    ) V2 @6 ~/ q8 ]8 k% a  losses.append(loss)* v5 O: ~2 }" j( ]9 g0 n# b  ?; F8 G
      
    . r& E& A# Z( h/ t9 n  loss.backward() # autograd
    4 Y. K; G6 d3 f. K3 X1 E) }  with torch.no_grad():: A$ ^8 u' P9 ]# G; B: e
        w  -= w.grad*0.0001   # 回归 w! a* g, i6 K( }1 y
        b  -= b.grad*0.0001    # 回归 b : x% V: y8 l% d8 b
      w.grad.zero_()  ! ]# f; H' }( ]9 h7 o" b" D4 m
      b.grad.zero_()- p. t1 M. n) I! a7 v

    , h7 c/ C+ o. w/ \& W1 xprint(w.item(),b.item()) #结果: j7 O3 L4 H1 _+ V; ?

    $ _4 {+ b$ Q) W( NOutput: 27.26387596130371  0.4974517822265625
    2 u. V8 I  B8 u) s3 x: o----------------------------------------------9 {( i% s& a, f, J$ O: h3 [2 P. U
    最后的结果,w可以回到 w0 = 27 附近,b却回不去 b0=15。两处红字,损失函数是矢量计算后的均值,感觉 b 的回归表达有问题。- _, ^1 Y$ k$ ?0 X' {1 a" x
    高手们帮看看是神马原因?: F5 K+ O: Q& v

    评分

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

    查看全部评分

    该用户从未签到

    沙发
    发表于 2023-2-14 19:23:02 | 只看该作者
    本帖最后由 老福 于 2023-2-14 21:58 编辑 4 {  |0 _. I( C/ y) u( R

    / v% `# I$ a# U7 U, u! b9 d没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?
      w+ B# p; d3 ^5 i- d9 _3 \8 n-------4 s( m1 z& w% |% C, ?: H. U
    不好意思,再看一遍,好像你在自算回归而不是用现成的工具直接出结果,上面的评论只有一点用,就是确认是不是算法有问题。: U7 e$ |7 `8 }$ R
    -------
    4 G/ o" ]; E+ A7 A+ F$ o/ H算法诊断部分,建议把循环次数改为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
    * w* q$ P- ^) n0 I+ e5 \8 i没有用过pytorch,但你把随机噪音部分改成均值为0的正态分布再试试看是不是符合预期?0 B1 x  G9 e- J% V
    -------
    . \; ~5 ^- n5 I5 ~/ A# s' P' \" {不好意思, ...

    % y( r, {7 D2 N8 u2 {谢谢,算法应该没问题,就是最简单的线性回归。# P$ |+ t) `% l; _
    我特意没有用现成的工具,就是想从最基本的地方深入理解一下。
    回复 支持 反对

    使用道具 举报

    该用户从未签到

    地板
    发表于 2023-2-14 22:00:48 | 只看该作者
    本帖最后由 老福 于 2023-2-14 22:02 编辑
      w0 `2 ?$ \& Q9 y
    雷达 发表于 2023-2-14 21:52# Z' G. P- n" M/ s+ z" G0 j
    谢谢,算法应该没问题,就是最简单的线性回归。
    8 D- f. r2 q' e& i我特意没有用现成的工具,就是想从最基本的地方深入理解 ...
    " d6 `! W6 P- N( J& I

    + N6 [; X( k! e0 c7 r9 l& c刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。
    # w8 F) X9 O& W  D6 u1 ~  L# n8 l3 Y+ {7 G9 F3 k: g: M- J
    或者把b但的起点改为1试试。
    回复 支持 反对

    使用道具 举报

  • TA的每日心情

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

    [LV.10]大乘

    5#
     楼主| 发表于 2023-2-15 00:25:26 | 只看该作者
    本帖最后由 雷达 于 2023-2-15 00:31 编辑
    % ]! s4 c1 ?/ P
    老福 发表于 2023-2-14 22:00: q: a5 `) l, `! d2 X8 I
    刚才更新了一下,建议增加循环次数或调一下步长,查一下loss曲线。1 \( q& h; S6 K0 X: o" f

    + ?( A) u- V5 H6 [3 f或者把b但的起点改为1试试。 ...
    $ v" }3 C( c: |3 n

    3 w8 f4 Q  u( [9 w  B4 r你是对的。
    $ b* F6 O! m2 f+ O2 P. t去掉了随机部分
    ; i/ b# E. E+ j* A7 R$ E#y = (x*27+15+random.randint(-2,3)).reshape(-1)
    6 z, G, A$ t3 G1 s: Q, }9 py = (x*27+15).reshape(-1)# C# y) s* E) W4 T' z8 r; `( `9 i
    & D6 N% w! B( s% q6 w  S3 d3 ]
    循环次数加成10倍,就看到 b 收敛了) \' c- C! K3 N
    w , b
    , m% M) g! i* _  Q  T8 _+ a27.002620697021484 14.826167106628418
    0 z/ B/ L6 j; T& a! {. l+ }6 b$ t" F+ i# }, P, M
    和 b 的起始位置无关,但 labeled data 用 y = (x*27+15+random.randint(-2,3)).reshape(-1) ,收敛就很慢。
    回复 支持 反对

    使用道具 举报

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

    GMT+8, 2026-8-11 20:48 , Processed in 0.067487 second(s), 18 queries , Gzip On.

    Powered by Discuz! X3.2

    © 2001-2013 Comsenz Inc.

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