码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • 【PyTorch】深度学习实践之反向传播 Back Propagation


    本文目录

    • 前馈计算
    • 反向传播过程
    • Tensor in PyTorch
    • 课堂练习:线性模型 Linear Model
      • 实现代码
      • 结果
    • 课后练习
    • 学习资料
    • 系列文章索引

    前馈计算

    权重维度增加,层数增加,模型变得复杂

    在这里插入图片描述

    但是化简后仍是线性,因此增加层数意义不大

    [图片]

    引入激活函数,从而增加非线性

    [图片]

    反向传播计算梯度,使用链式法则
    [图片]

    [图片]

    反向传播过程

    [图片]

    Tensor in PyTorch

    Tenso(张量):PyTorch中存储数据的基本元素。
    Tensor两个重要的成员,data和grad。(grad也是个张量)

    课堂练习:线性模型 Linear Model

    实现代码

    import torch
    
    # 已知数据:
    x_data = [1.0,2.0,3.0]
    y_data = [2.0,4.0,6.0]
    # 线性模型为y = wx, 预测x = 4时, y的值
    
    # 假设 w = 1
    w = torch.Tensor([1.0])
    w.requires_grad = True
    
    # 定义模型:
    def forward(x):
            return x*w
    
    # 定义损失函数:
    def loss(x,y):
            y_pred = forward(x)
            return (y_pred - y)**2
    
    print("Prediction before training:",4,'%.2f'%(forward(4)))
    
    for epoch in range(100):
            for x, y in zip(x_data,y_data):
                    l = loss(x,y)
                    l.backward() # 对requires_grad = True的Tensor(w)计算其梯度并进行反向传播,并且会释放计算图进行下一次计算
                    print("\tgrad:%.1f %.1f %.2f" % (x,y,w.grad.item()))
                    w.data = w.data - 0.01 * w.grad.data # 通过梯度对w进行更新
                    w.grad.data.zero_() #梯度清零
            print("Epoch:%d, w = %.2f, loss = %.2f" % (epoch,w,l.item()))
    
    print("Prediction after training:",4,'%.2f'%(forward(4)))
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21
    • 22
    • 23
    • 24
    • 25
    • 26
    • 27
    • 28
    • 29
    • 30
    • 31
    • 32
    • 本算法中反向传播主要体现在,l.backward()。调用该方法后w.grad由None更新为Tensor类型,且w.grad.data的值用于后续w.data的更新。
    • l.backward()会把计算图中所有需要梯度(grad)的地方都会求出来,然后把梯度都存在对应的待求的参数中,最终计算图被释放。
    • 取tensor中的data是不会构建计算图的。

    结果

    在这里插入图片描述

    课后练习

    1. 计算y=xw的梯度

    在这里插入图片描述

    2. 计算仿射模型y=xw+b的梯度

    在这里插入图片描述

    3. 使用计算图计算y=w1x^2+w2x+b的梯度

    在这里插入图片描述

    4. 使用Pytorch计算y=w1x^2+w2x+b的梯度

    二次模型 Quadratic Model

    在这里插入图片描述

    代码如下:

    import torch
    
    # 已知数据:
    x_data = [1.0,2.0,3.0]
    y_data = [6.0,11.0,18.0]
    # 线性模型为y = w1x²+w2x+b时, 预测x = 4时, y的值
    
    # 假设 w = 1, b = 1
    w1 = torch.Tensor([1.0])
    w1.requires_grad = True
    w2 = torch.Tensor([1.0])
    w2.requires_grad = True
    b = torch.Tensor([1.0])
    b.requires_grad = True
    
    # 定义模型:
    def forward(x):
            return x*x*w1+x*w2+b
    
    # 定义损失函数:
    def loss(x,y):
            y_pred = forward(x)
            return (y_pred - y)**2
    
    print("Prediction before training:",4,'%.2f'%(forward(4)))
    
    for epoch in range(1000):
            for x, y in zip(x_data,y_data):
                    l = loss(x,y)
                    l.backward() # 对requires_grad = True的Tensor(w)计算其梯度并进行反向传播,并且会释放计算图进行下一次计算
                    w1.data = w1.data - 0.02 * w1.grad.data # 通过梯度对w进行更新
                    w2.data = w2.data - 0.02 * w2.grad.data
                    b.data = b.data - 0.02 * b.grad.data
                    # 梯度清零
                    w1.grad.data.zero_()
                    w2.grad.data.zero_() 
                    b.grad.data.zero_()
            print("Epoch:%d, w1 = %.4f,w2 = %.4f,b = %.4f, loss = %.4f" % (epoch,w1,w2,b,l.item()))
    
    print("Prediction after training:",4,'%.4f'%(forward(4)))
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21
    • 22
    • 23
    • 24
    • 25
    • 26
    • 27
    • 28
    • 29
    • 30
    • 31
    • 32
    • 33
    • 34
    • 35
    • 36
    • 37
    • 38
    • 39
    • 40

    结果:

    在这里插入图片描述


    学习资料

    • https://blog.csdn.net/weixin_43786637/article/details/126117060
    • https://blog.csdn.net/Lilo_/article/details/113522485?utm_medium=distribute.pc_relevant.none-task-blog-2defaultbaidujs_title~default-9-113522485-blog-126117060.pc_relevant_aa&spm=1001.2101.3001.4242.6&utm_relevant_index=12
    • https://blog.csdn.net/lizhuangabby/article/details/125548170?app_version=5.7.0&code=app_1562916241&csdn_share_tail=%7B%22type%22%3A%22blog%22%2C%22rType%22%3A%22article%22%2C%22rId%22%3A%22125548170%22%2C%22source%22%3A%22qq_43800119%22%7D&ctrtid=0pZiz&uLinkId=usr1mkqgl919blen&utm_source=app

    系列文章索引

    教程指路:【《PyTorch深度学习实践》完结合集】 https://www.bilibili.com/video/BV1Y7411d7Ys?share_source=copy_web&vd_source=3d4224b4fa4af57813fe954f52f8fbe7

    1. 线性模型 Linear Model
    2. 梯度下降 Gradient Descent
    3. 反向传播 Back Propagation
    4. 用PyTorch实现线性回归 Linear Regression with Pytorch
    5. 逻辑斯蒂回归 Logistic Regression
    6. 多维度输入 Multiple Dimension Input
    7. 加载数据集Dataset and Dataloader
    8. 用Softmax和CrossEntroyLoss解决多分类问题(Minst数据集)
    9. CNN基础篇——卷积神经网络跑Minst数据集
    10. CNN高级篇——实现复杂网络
    11. RNN基础篇——实现RNN
    12. RNN高级篇—实现分类
  • 相关阅读:
    Git - 入门到熟悉_分支管理
    数字政府一网统管体系下的运维管理软件应用探讨
    人工智能的隐私保护探讨
    使用vite + vue3 + ant-design-vue + vue-router + vuex 创建一个管理应用
    iPortal如何灵活设置用户名及密码的安全规则
    下载安装Microsoft ODBC Driver for SQL Server和配置SQL Server ODBC数据源
    【正点原子STM32连载】 第六十二章 UCOSII实验2-信号量和邮箱 摘自【正点原子】MiniPro STM32H750 开发指南_V1.1
    千变万化的Promise
    echarts legend如何控制标签文字长度
    Python使用Base64编码
  • 原文地址:https://blog.csdn.net/qq_43800119/article/details/126415332
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号