码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • 【PyTorch】深度学习实践之 逻辑斯蒂回归 Logistic Regression


    本文目录

    • 回归vs分类
    • sigmoid函数
    • 损失函数
    • 例子
    • 课堂练习
      • 模型实现
      • 计算损失
      • 实现代码
      • 测试模型
    • 学习资料
    • 系列文章索引

    回归vs分类

    在这里插入图片描述

    • 回归是预测数值
    • 分类是预测类别概率

    sigmoid函数

    [图片]

    Logistic Function是最典型的sigmoid函数,因此有些书会直接说成sigmoid函数。
    实际上满足如下条件即可称为sigmoid函数:

    • 饱和函数
    • 单调递增
    • 存在极限
      [图片]

    损失函数

    使用二分类交叉熵公式:
    [图片]

    • y=1,预测值接近1,loss减小
    • y=0,预测值接近0,loss减小

    例子

    [图片]

    • 多个loss,MiniBatch求均值

    课堂练习

    模型实现

    [图片]

    可以看到init部分没有区别,因为逻辑回归没有参数增加。

    计算损失

    [图片]

    实现代码

    [图片]

    import torch
    import matplotlib.pyplot as plt
    
    #1.准备数据集
    x_data = torch.Tensor([[1.0],[2.0],[3.0]])
    y_data = torch.Tensor([[2.0],[4.0],[6.0]])
    
    #2.使用Class设计模型
    class LogisticRegressionModel(torch.nn.Module):
        def __init__(self):
            super(LinearModel,self).__init__()
            self.linear = torch.nn.Linear(1,1)  
        def forward(self,x):
            y_pred = F.sigmoid(self.linear(x))
            return y_pred
     
    model = LogisticRegressionModel()  #创建类LinearModel的实例
    
    #3.构建损失函数和优化器的选择
    criterion = torch.nn.BCELoss(size_average=False)
    optimizer = torch.optim.SGD(model.parameters(),lr=0.01)
    
    #4.进行训练迭代
    epoch_list =[]
    loss_list=[]
    for epoch in range(1000):
        y_pred = model(x_data)
        loss = criterion(y_pred,y_data)
        print(epoch,loss.item()) 
       
        optimizer.zero_grad()
        loss.backward()    
        optimizer.step()  
        epoch_list.append(epoch+1)
        loss_list.append(loss.item())
    
    # 画图
    plt.plot(epoch_list,loss_list)
    plt.xlabel('epoch')
    plt.ylabel('loss')
    plt.show()
    
    • 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
    • 41

    [图片]

    测试模型

    [图片]

    import numpy as np
    import matplotlib.pyplot as plt
    
    x=np.linspace(0,10,200)
    x_t=torch.Tensor(x).view((200,1))
    # 使用训练好的模型
    y_t=model(x_t)
    y=y_t.data.numpy()
    
    plt.plot(x,y)
    plt.plot([0,10],[0.5,0.5],c='r')
    plt.xlabel('Hours')
    plt.ylabel('Probability of Pass')
    plt.grid()
    plt.show()
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15

    在这里插入图片描述


    学习资料

    • https://blog.csdn.net/qq_42585108/article/details/108148210

    系列文章索引

    教程指路:【《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高级篇—实现分类
  • 相关阅读:
    曲线艺术编程 coding curves 第九章 旋轮曲线(ROULETTE CURVES)
    【LeetCode】字节面试-行、列递增的二维数组数字查找
    【文件后缀名批量修改,python,webp】
    MVC 框架安全
    85、Redis连接相关的命令, key相关命令
    什么样的程序员在 35 岁以后依然被公司抢着要?
    【超实用】教你生成GUID
    文件编码格式
    k8s教程(02)-入门及案例
    一、Lua 教程的学习
  • 原文地址:https://blog.csdn.net/qq_43800119/article/details/126415539
  • 最新文章
  • 攻防演习之三天拿下官网站群
    数据安全治理学习——前期安全规划和安全管理体系建设
    企业安全 | 企业内一次钓鱼演练准备过程
    内网渗透测试 | Kerberos协议及其部分攻击手法
    0day的产生 | 不懂代码的"代码审计"
    安装scrcpy-client模块av模块异常,环境问题解决方案
    leetcode hot100【LeetCode 279. 完全平方数】java实现
    OpenWrt下安装Mosquitto
    AnatoMask论文汇总
    【AI日记】24.11.01 LangChain、openai api和github copilot
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1
正则表达式工具 cron表达式工具 密码生成工具

京公网安备 11010502049817号