深度学习中的损失函数与优化算法基础知识

一、前言

在深度学习训练流水线中,网络结构负责特征提取,损失函数定义优化目标优化算法负责参数迭代更新,三者共同构成模型收敛的完整闭环。很多初学者仅停留在调用nn.CrossEntropyLoss()optim.Adam()的层面,不清楚公式背后的梯度流向、不同损失对异常值 / 样本不均衡的敏感度差异,也不理解 SGD、Adam、RMSprop 内在动量与自适应学习率机制,最终遇到损失不下降、震荡不收敛、过拟合严重、标签不均衡准确率虚高等问题时无从排查。

本文分为两大核心模块:

  1. 损失函数:给出数学公式、梯度推导逻辑、适用任务、缺陷与改良方案;
  2. 优化算法:从梯度下降本源出发,逐层迭代讲解 SGD→Momentum→RMSprop→Adam 原理、超参影响。

全文兼顾理论深度与工程落地,适合 AI 入门夯实底层、项目调参参考,附加完整 PyTorch 可运行代码、训练问题排查表、工程选型对照表。

二、模型训练完整闭环(底层逻辑铺垫)

2.1 完整前向 - 反向传播链路

    1. 输入样本进入网络前向传播 ,得到预测输出

    2. 损失函数 计算预测值 与真实标签 的误差Loss;

    3. 链式法则反向传播 ,对网络所有权重参数 求梯度

    4. 优化器根据梯度、学习率、历史动量自适应更新权重;

    5. 迭代多轮直至损失收敛、验证集指标稳定。

2.2 拟人化通俗理解

  • 损失函数 = 评分细则,规定 "什么样的预测算错、错误扣多少分";
  • 梯度 = 告诉参数往哪个方向调整能降低分数;
  • 优化器 = 学习策略,决定每次步子迈多大、是否保留过往经验惯性。

第一部分 损失函数:误差度量的数学本质

3.1 损失函数通用定义

损失函数是一个映射函数,输入预测值与真值,输出非负标量损失值。

  • 单个样本计算值:Loss
  • 一个 Batch 所有样本平均:Cost 代价函数 ; 训练目标:最小化整个训练集代价函数

3.2 回归任务常用损失函数(带公式 + 深度分析)

3.2.1 均方误差 MSE (L2 Loss)

数学公式

单样本:

批量平均代价:

求导梯度(关键,决定训练特性)
深度原理与优缺点
  1. 优点:损失曲线处处可导、梯度连续平滑,优化过程稳定;误差越大梯度越大,修正力度越强;适合连续值回归(温度、坐标、流量预测)。
  2. 致命缺陷:对离群异常值极度敏感。误差做平方放大,少量脏样本会主导梯度方向,导致模型为拟合噪声偏离真实分布。
  3. 梯度问题:当预测值与真值差距极大时,梯度过大引发训练震荡。

3.2.2 平均绝对误差 MAE (L1 Loss)

数学公式

单样本公式

批量平均代价公式:

梯度特性

导数为常数,误差大小不影响梯度幅值。

深度分析
  1. 优点:对异常值鲁棒性极强,不受极端值干扰,回归中位数效果更好;
  2. 缺点:在 处不可导,梯度突变易导致收敛抖动;全程梯度大小一致,远距离误差修正速度慢。

3.2.3 Huber Loss(MSE+MAE 折中方案,工程首选)

分段公式
作用

小误差用 MSE 保证平滑收敛,大误差用 MAE 抑制异常值冲击,是带噪声回归数据集最优损失。

3.3 分类任务核心损失函数(重点)

3.3.1 二元交叉熵 BCE Loss(二分类)

适用场景:0/1 二分类、逻辑回归、Sigmoid 输出。

公式

经过 Sigmoid 压缩到 (0,1) 概率区间。

梯度优势

相比 MSE 做分类,交叉熵梯度与误差正相关,预测越离谱梯度越大,快速修正;MSE 容易出现梯度消失。

3.3.2 多分类交叉熵 CrossEntropyLoss(CNN 图像分类标配)

PyTorch 中nn.CrossEntropyLoss = LogSoftmax + NLLLoss 封装一体。

公式(单样本)

C 为总类别数,内部自动完成 Softmax 概率归一化。

深度底层优势
  1. 完美适配多类别概率分布输出,梯度不会饱和;
  2. 硬标签训练收敛速度远快于 MSE;
  3. 缺陷:对类别不均衡数据集不友好,样本多的类别主导损失。

3.3.3 改良版:Focal Loss(解决正负样本不均衡)

目标检测、小目标识别常用,在交叉熵基础上加调制因子,降低易分样本权重,让模型专注学习难样本。

为衰减系数,一般取 2,极大提升不均衡数据集精度。

3.4 损失函数选型对照表(工程直接套用)

表格

损失函数 适用任务 核心优点 短板
MSE 常规回归 梯度平滑、收敛稳 异常值敏感
MAE 含噪声回归 抗离群点 零点不可导、收敛慢
Huber 噪声回归通用 折中二者优点 需要调超参 δ
BCE 二分类 梯度灵敏 多分类不适用
CrossEntropy 多分类图像任务 收敛快、梯度稳定 类别不平衡效果差
Focal Loss 检测、不均衡分类 聚焦难样本 超参调参成本高

3.5 损失函数 PyTorch 完整代码示例

python 复制代码
import torch
import torch.nn as nn

# ---------------------- 1.回归类损失 ----------------------
y_pred = torch.tensor([2.1, 3.5, 4.2])
y_true = torch.tensor([2.0, 3.0, 4.0])

# MSE损失
loss_mse = nn.MSELoss()(y_pred, y_true)
# MAE损失
loss_mae = nn.L1Loss()(y_pred, y_true)
# Huber损失
loss_huber = nn.HuberLoss(delta=0.5)(y_pred, y_true)

print("MSE Loss:", loss_mse.item())
print("MAE Loss:", loss_mae.item())
print("Huber Loss:", loss_huber.item())

# ---------------------- 2.分类类损失 ----------------------
# 二分类 BCE (输入为sigmoid之前logits)
bce_logits = torch.tensor([0.8, -1.2, 0.3])
bce_label = torch.tensor([1.0, 0.0, 1.0])
loss_bce = nn.BCEWithLogitsLoss()(bce_logits, bce_label)
print("BCE Loss:", loss_bce.item())

# 多分类交叉熵 CrossEntropyLoss
# 输入:[batch, num_classes] 原始得分,无需softmax
cls_logits = torch.tensor([[2.3, 1.1, 0.2],
                           [0.5, 3.2, 1.0]])
cls_label = torch.tensor([0, 1])  # 真实类别索引
loss_ce = nn.CrossEntropyLoss()(cls_logits, cls_label)
print("CrossEntropy Loss:", loss_ce.item())

第二部分 优化算法:梯度下降的迭代进化原理

4.1 最原始:批量梯度下降 BGD

更新公式:

:学习率

:整个数据集代价函数梯度。

优缺点

  • 优点:梯度方向最准确,收敛轨迹平稳;
  • 缺点:大数据集计算全量梯度耗时爆炸,无法在线更新。

4.2 SGD 随机梯度下降(工业经典)

更新公式

每次仅用单个样本梯度更新:

Mini-Batch SGD 小批量版(现在通用):用一个 Batch 平均梯度更新。

深度特性

  1. 梯度带有噪声,收敛过程震荡,但更容易跳出局部极小值,泛化能力最强
  2. 可在线流式训练,大数据集效率极高;
  3. 缺陷:学习率固定时收敛速度慢,震荡严重,易卡在鞍点。

4.3 Momentum 动量 SGD(引入惯性思想)

公式

动量系数一般取 0.9,保留历史下降惯性。

原理通俗解释

下坡时顺着之前的速度加速向下,山谷震荡时反向抵消抖动:

  • 梯度同向:加速收敛;
  • 梯度反向:抑制来回震荡。 解决纯 SGD 收敛慢、抖动剧烈问题。

4.4 RMSprop:自适应学习率先驱

核心思路

对每个参数维护梯度平方移动均值,动态缩放学习率:

梯度频繁波动的参数,分母变大,有效学习率自动降低;梯度平稳参数学习率更大。解决不同参数更新幅度不一致问题。

4.5 Adam 自适应矩估计(目前最主流万能优化器)

本质 = Momentum + RMSprop 合体

同时维护一阶动量(惯性)、二阶动量(梯度平方自适应):

默认超参:

深度优缺点分析

优点
  1. 自带动量加速收敛 + 逐参数自适应学习率,开箱即用,新手零门槛;
  2. 对稀疏梯度、不同尺度参数适配极好,图像分类、NLP、检测都通用;
  3. 前期下降速度极快,快速看到指标提升。
严重短板(很多人忽略)
  1. 自适应学习率后期容易收敛到局部最优,泛化能力弱于纯 SGD;
  2. 小数据集极易过拟合;
  3. 长期训练二阶动量累积会导致学习率过早衰减,精度天花板低于 SGD。

工程选型经验

  • 快速实验、调模型结构、验证可行性:直接 Adam;
  • 最终上线、追求最高精度与泛化:SGD + 动量 + 学习率余弦衰减。

4.6 学习率 LR 对优化器的决定性影响(深度要点)

  1. LR 过大:参数跨步越过极小值,损失持续震荡不收敛,甚至发散;
  2. LR 过小:迭代极慢,困在局部最优无法跳出,训练成本极高;
  3. 最优策略:热身 warmup + 余弦退火衰减,前期大步快速收敛,后期小幅精细微调。

4.7 各类优化器 PyTorch 完整代码示例

python 复制代码
import torch
import torch.nn as nn
import torch.optim as optim

# 简单单层模拟网络
model = nn.Linear(10, 2)
loss_fn = nn.CrossEntropyLoss()

# 1. SGD 带动量
opt_sgd = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)

# 2. RMSprop
opt_rmsprop = optim.RMSprop(model.parameters(), lr=0.001, alpha=0.99)

# 3. Adam(最常用)
opt_adam = optim.Adam(model.parameters(), lr=1e-3, betas=(0.9, 0.999))

# 模拟单步训练流程
x = torch.randn(8, 10)
y = torch.tensor([0,1,0,1,0,1,0,1])

# 前向传播
pred = model(x)
loss = loss_fn(pred, y)

# 反向传播 + 参数更新
opt_adam.zero_grad()  # 清空历史梯度
loss.backward()       # 反向求梯度
opt_adam.step()       # 优化器更新权重

print("单次损失值:", loss.item())

第三部分 损失函数 + 优化器联动逻辑与实战排坑

5.1 完整联动链路(带梯度流向)

  1. 模型前向输出 logits → 损失函数计算 Loss(确定误差量化规则);
  2. 链式求导反向传播,逐层计算参数梯度;
  3. 优化器读取梯度,依靠自身算法(动量 / 自适应 LR)计算更新量;
  4. 权重迭代更新,下一轮前向传播损失降低。

5.2 高频训练问题深度排查(对应底层原理)

  1. 损失完全不下降 排查:损失函数任务不匹配、学习率过大 / 过小、梯度被冻结、数据集标签错误、激活函数导致梯度消失。

  2. 损失下降但验证集准确率不动 排查:严重过拟合,Adam 在小数据集过度拟合,改用 SGD + 正则化、Dropout、早停。

  3. 训练震荡无法收敛 排查:学习率过高、使用纯 SGD 无动量、BatchSize 过小梯度噪声太大。

  4. 类别不均衡准确率虚高 排查:CrossEntropy 对样本数多类别倾斜,替换为 Focal Loss、增加类别权重 weight 参数。

  5. 后期精度上不去 排查:Adam 局部最优陷阱,改用 SGD + 余弦学习率衰减。

5.3 工程通用固定搭配(直接照搬)

  1. 常规图像分类:CrossEntropyLoss + Adam(快速迭代),最终 SGD 微调;
  2. 目标检测 / 小目标:Focal Loss + AdamW;
  3. 数值回归带噪声:Huber Loss + SGD;
  4. 二分类任务:BCEWithLogitsLoss + Adam。

六、全文总结

  1. 损失函数本质是误差的数学度量规则,回归看 MSE/MAE/Huber,分类看交叉熵 / Focal Loss,任务错配直接导致训练失效;公式背后的梯度形态决定了收敛稳定性与对噪声的鲁棒性;
  2. 优化算法是梯度下降的层层迭代升级:SGD 保证泛化,Momentum 增加惯性提速,RMSprop 实现自适应学习率,Adam 集大成但存在泛化短板;学习率调度是决定最终精度的隐形关键;
  3. 二者必须绑定理解:损失定义 "往哪优化",优化器定义 "如何优化",搭配不合理会出现不收敛、过拟合、指标瓶颈等大量疑难问题;
  4. 实操层面记住最简落地规则:实验用 Adam,冲精度用 SGD;回归优先 Huber,不均衡分类优先 Focal Loss。
相关推荐
A15362551 小时前
五金批发电商业财一体化 ERP 推荐:打通订单、库存、财务对账
大数据·运维·人工智能·零售
happyprince1 小时前
02_OpenCodeReview 具体观:八种算法机制与背后的论文谱系
算法
变量未定义~1 小时前
图论+动态规划——魔法阵
算法
大爱编程♡1 小时前
C++基础-类和对象
c++·算法
ysa0510301 小时前
【板子】费用流
算法·深度优先·图论
把所有砖敲烂2 小时前
GLM 5.2 核心能力与效果实测全景
人工智能
神奇霸王龙2 小时前
Qwen3.7-Max屠榜:推理成本仅GPT-5.5的1/25
人工智能·python·gpt·ai·aigc·ai编程
StarkCoder2 小时前
AI 会做多、看少、不收尾:七种失效和拦住它们的办法
人工智能·架构
玖玥拾2 小时前
LeetCode 80 删除有序数组中的重复项 II
算法·leetcode