码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • [实践应用] 深度学习之损失函数


    文章总览:YuanDaiMa2048博客文章总览


    深度学习之损失函数

      • 1. 回归任务
        • 1.1 均方误差 (MSE)
        • 1.2 平均绝对误差 (MAE)
      • 2. 二分类任务
        • 2.1 二元交叉熵 (Binary Cross-Entropy)
      • 3. 多分类任务
        • 3.1 类别交叉熵 (Categorical Cross-Entropy)
      • 4. 序列生成任务(例如,机器翻译)
          • 4.1 序列交叉熵 (Sequence Cross-Entropy)
      • 5. 回归任务的正则化
          • 5.1 L2 正则化(权重衰减)
      • 其他介绍

    在机器学习和深度学习中,不同的任务使用不同的损失函数来衡量模型的性能。

    1. 回归任务

    任务: 预测一个连续的数值。

    1.1 均方误差 (MSE)

    原理: MSE 衡量预测值与实际值之间的平方差的平均值,适用于回归任务。它对异常值敏感。

    公式:
    MSE = 1 n ∑ i = 1 n ( 预测值 i − 实际值 i ) 2 \text{MSE} = \frac{1}{n} \sum_{i=1}^{n}(\text{预测值}_i - \text{实际值}_i)^2 MSE=n1​i=1∑n​(预测值i​−实际值i​)2

    PyTorch 代码:

    import torch
    import torch.nn as nn
    
    # 定义均方误差损失函数
    mse_loss = nn.MSELoss()
    

    1.2 平均绝对误差 (MAE)

    原理: MAE 衡量预测值与实际值之间的绝对差的平均值,对异常值不太敏感。

    公式:
    MAE = 1 n ∑ i = 1 n ∣ 预测值 i − 实际值 i ∣ \text{MAE} = \frac{1}{n} \sum_{i=1}^{n}|\text{预测值}_i - \text{实际值}_i| MAE=n1​i=1∑n​∣预测值i​−实际值i​∣

    PyTorch 代码:

    import torch
    import torch.nn as nn
    
    # 定义平均绝对误差损失函数
    mae_loss = nn.L1Loss()
    

    2. 二分类任务

    任务: 预测样本属于两个类别中的一个(例如,垃圾邮件分类)。

    2.1 二元交叉熵 (Binary Cross-Entropy)

    原理: 计算预测概率与实际标签之间的交叉熵,用于二分类任务。

    公式:
    BCE = − 1 n ∑ i = 1 n [ y i log ⁡ ( p i ) + ( 1 − y i ) log ⁡ ( 1 − p i ) ] \text{BCE} = -\frac{1}{n} \sum_{i=1}^{n} [y_i \log(p_i) + (1 - y_i) \log(1 - p_i)] BCE=−n1​i=1∑n​[yi​log(pi​)+(1−yi​)log(1−pi​)]

    PyTorch 代码:

    import torch
    import torch.nn as nn
    
    # 定义二元交叉熵损失函数
    bce_loss = nn.BCEWithLogitsLoss()  # 结合了 Sigmoid 激活和 BCE 损失
    

    3. 多分类任务

    任务: 预测样本属于多个类别中的一个(例如,手写数字分类)。

    3.1 类别交叉熵 (Categorical Cross-Entropy)

    原理: 计算预测的概率分布与实际类别之间的交叉熵,用于多分类任务。

    公式:
    CCE = − 1 n ∑ i = 1 n ∑ k = 1 K y i , k log ⁡ ( p i , k ) \text{CCE} = -\frac{1}{n} \sum_{i=1}^{n} \sum_{k=1}^{K} y_{i,k} \log(p_{i,k}) CCE=−n1​i=1∑n​k=1∑K​yi,k​log(pi,k​)

    其中 K K K 是类别数, y i , k y_{i,k} yi,k​ 是实际类别的 one-hot 编码, p i , k p_{i,k} pi,k​ 是预测的概率。

    PyTorch 代码:

    import torch
    import torch.nn as nn
    
    # 定义类别交叉熵损失函数
    cross_entropy_loss = nn.CrossEntropyLoss()  # 直接对 logits 应用 Softmax 和计算交叉熵
    

    4. 序列生成任务(例如,机器翻译)

    任务: 预测序列中每个位置的类别(例如,翻译每个单词)。

    4.1 序列交叉熵 (Sequence Cross-Entropy)

    原理: 与多分类交叉熵类似,但应用于序列数据,计算预测序列与实际序列之间的交叉熵。

    公式:
    Sequence CCE = − 1 n ∑ i = 1 n ∑ t = 1 T ∑ k = 1 K y i , t , k log ⁡ ( p i , t , k ) \text{Sequence CCE} = -\frac{1}{n} \sum_{i=1}^{n} \sum_{t=1}^{T} \sum_{k=1}^{K} y_{i,t,k} \log(p_{i,t,k}) Sequence CCE=−n1​i=1∑n​t=1∑T​k=1∑K​yi,t,k​log(pi,t,k​)

    PyTorch 代码:

    import torch
    import torch.nn as nn
    
    # 对于序列生成任务,通常使用 CrossEntropyLoss 处理每个时间步的预测
    sequence_cross_entropy_loss = nn.CrossEntropyLoss()
    

    5. 回归任务的正则化

    任务: 通过将正则化项添加到损失函数来防止过拟合。

    5.1 L2 正则化(权重衰减)

    原理: 在损失函数中添加权重的平方和,鼓励较小的权重值。

    公式:
    Regularized Loss = 原始损失 + λ ∑ j = 1 m W j 2 \text{Regularized Loss} = \text{原始损失} + \lambda \sum_{j=1}^{m} W_j^2 Regularized Loss=原始损失+λj=1∑m​Wj2​

    PyTorch 代码:

    import torch.optim as optim
    
    # 定义模型
    model = nn.Linear(10, 1)
    
    # 定义优化器,并添加 L2 正则化(weight_decay)
    optimizer = optim.SGD(model.parameters(), lr=0.01, weight_decay=0.01)
    

    其他介绍

    • 深度学习之激活函数
  • 相关阅读:
    【Docker】二、docker镜像的制作运行发布
    DayDreamInGIS 逆地理编码工具(根据经纬度获取位置描述)插件源码解析
    web应用及微信小程序版本更新检测方案实践
    【尘缘赠书活动第四期】推荐几本架构师成长和软件架构技术相关的好书,助你度过这个不太景气的寒冬!
    Elasticsearch集群连载-es集群安装
    交换机堆叠 配置(H3C)堆叠中一台故障如何替换
    matlab 矩阵逆运算的条件数
    浅谈图片展示、图片自适应解决方案
    golang validator 提示消息本地化(中英文案例)
    【无标题】
  • 原文地址:https://blog.csdn.net/2301_79288416/article/details/141935063
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    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号