梯度下降如何实现参数优化:从线性回归到 Sigmoid 分类

核心问题:给定输入 x 和真实标签 y,如何通过调整参数 wb,让预测函数越来越准?


1. 问题设定

机器学习中,很多任务都可以抽象为:

  • 有一组实际数据 (x, y)
  • 想构造一个预测函数 ,通过 x 预测 y
  • 用某个误差/损失函数衡量预测的好坏
  • 通过不断调整函数中的参数,使损失越来越小

这个"调整参数使损失变小"的过程,就是参数优化。梯度下降(Gradient Descent)是其中最常用的一种方法。


2. 线性预测:ŷ = wx + b

最简单的预测函数是线性函数:

text 复制代码
ŷ = wx + b

其中:

  • w 是权重(weight)
  • b 是偏置(bias)
  • ŷ 是模型对 y 的预测值

2.1 损失函数

为了衡量预测值 ŷ 与真实值 y 的差距,常用平方误差(Squared Error)

text 复制代码
e = (y - ŷ)² = (y - (wx + b))²

如果有 n 个样本,整体损失(均方误差 MSE)为:

text 复制代码
L(w, b) = (1/n) Σᵢ (yᵢ - (wxᵢ + b))²

2.1.1 为什么采用平方误差 ?

衡量预测误差的损失函数不止一种(也可以用绝对误差 |y − ŷ|、Huber 损失、交叉熵等),平方误差之所以成为最经典的起点,主要基于以下几点:

  1. 处处可导、数学性质好 。平方误差 e = (y − ŷ)² 对参数光滑可导,没有绝对值函数在 0 处的折点,因此梯度处处存在,便于用链式法则一路求到 wb。本笔记里 ∂e/∂a = −2(y − a) 这个简洁结果,正是平方误差求导直接带来的。
  2. 对大误差惩罚更重,收敛更快。误差被平方后,误差越大梯度越大,模型会优先修正那些偏差严重的样本,初期下降更迅速。
  3. 与最大似然估计等价。在假设观测噪声服从高斯(正态)分布的前提下,最小化平方误差等价于对参数做极大似然估计(MLE),这让它在统计学上有坚实的理论依据。
  4. 线性回归下为凸函数,存在闭式解 。如 2.2 节所述,L(w, b) 关于参数是凸的,可直接令偏导为 0 解析求解,无需迭代。

补充说明:在二分类(Sigmoid 输出)场景下,平方误差并不是最优损失------它会让梯度带上 a(1 − a) 因子,当预测接近 0 或 1 时梯度趋近于 0(学习缓慢)。交叉熵损失(Cross-Entropy)的梯度形态更好,是分类任务的实际首选。手写笔记使用平方误差,是为了让"链式法则求偏导"这条推导主线最直观,属于教学选择。

2.2 解析解:令偏导数为 0

在线性回归中,L(w, b) 是关于 wb 的凸函数,可以直接对 wb 求偏导并令其为 0:

text 复制代码
∂L/∂w = 0
∂L/∂b = 0

解这个方程组,就能得到最优的 wb。这是线性回归的闭式解(Normal Equation)。

但现实中,很多问题并不是简单的线性拟合。例如,人类很多判断是"离散的分类"------是或不是、0 或 1。这时需要引入非线性函数。


3. 从拟合到分类:引入 Sigmoid

3.1 为什么需要 Sigmoid

线性函数 wx + b 的输出范围是 (-∞, +∞),而分类问题(如二分类)希望输出一个在 (0, 1) 之间的概率值。

Sigmoid 函数正好能把任意实数压缩到 (0, 1) 区间:

text 复制代码
σ(z) = 1 / (1 + e^(-z))

它的图像呈"S"形:

  • z 很大时,σ(z) 接近 1
  • z 很小时,σ(z) 接近 0
  • z = 0 时,σ(z) = 0.5

3.2 新的预测函数

引入 Sigmoid 后,预测函数变为:

text 复制代码
ŷ = σ(wx + b) = 1 / (1 + e^(-(wx + b)))

损失函数仍可用平方误差:

text 复制代码
e = (y - σ(wx + b))²

现在的问题是:wb 被包在 Sigmoid 里面,直接令偏导为 0 很难解出闭式解,需要用迭代优化方法------梯度下降。


4. 复合函数求导:链式法则

Sigmoid 让 e 成为了一个关于 wb 的复合函数。为了求 ∂e/∂w∂e/∂b,需要用到链式法则(Chain Rule)

4.1 拆解复合函数

把预测过程拆成三步:

text 复制代码
z = wx + b
a = σ(z) = 1 / (1 + e^(-z))
e = (y - a)²

其中 a 就是模型的预测输出 ŷ

4.2 分别求局部导数

第一步:e 对 a 求导

text 复制代码
e = (y - a)²
∂e/∂a = -2(y - a)

第二步:a 对 z 求导(Sigmoid 的导数)

text 复制代码
σ'(z) = σ(z)(1 - σ(z)) = a(1 - a)
∂a/∂z = a(1 - a)

第三步:z 对 w 和 b 求导

text 复制代码
z = wx + b
∂z/∂w = x
∂z/∂b = 1

4.3 合成最终偏导数

根据链式法则:

text 复制代码
∂e/∂w = (∂e/∂a) · (∂a/∂z) · (∂z/∂w)
      = -2(y - a) · a(1 - a) · x

∂e/∂b = (∂e/∂a) · (∂a/∂z) · (∂z/∂b)
      = -2(y - a) · a(1 - a) · 1

这两个偏导数告诉了我们:wb 发生微小变化时,损失 e 会怎么变化


5. 梯度下降参数更新

5.1 核心思想

梯度下降的基本思路:

  1. 计算当前参数下损失函数对各个参数的偏导数(梯度)
  2. 沿着梯度的反方向更新参数(因为梯度指向损失增长最快的方向)
  3. 重复以上步骤,直到损失收敛

5.2 参数更新公式

text 复制代码
w := w - η · ∂e/∂w
b := b - η · ∂e/∂b

其中 η学习率(Learning Rate),控制每一步更新的步长:

  • η 太小:收敛慢
  • η 太大:可能震荡甚至发散

5.3 批量、随机与小批量梯度下降

  • 批量梯度下降(BGD):每次用全部样本计算梯度,方向稳定但计算量大
  • 随机梯度下降(SGD):每次用一个样本更新,计算快但波动大
  • 小批量梯度下降(Mini-batch):折中方案,每次用一小批(如 32、64 个)样本

实际中最常用的是 Mini-batch Gradient Descent。


6. 代码示例:用 NumPy 实现

下面是一个简化版的二分类梯度下降实现,损失采用手写笔记中的平方误差。

python 复制代码
import numpy as np


def sigmoid(z):
    return 1 / (1 + np.exp(-z))


# 构造简单的二分类数据
np.random.seed(0)
x = np.random.randn(100, 1)
w_true, b_true = 2.0, -1.0
# 真实标签:当 w*x + b > 0 时为 1,否则为 0
y = (w_true * x + b_true + 0.3 * np.random.randn(100, 1)) > 0

# 初始化参数
w, b = 0.0, 0.0
lr = 0.1
epochs = 500

for epoch in range(epochs):
    z = w * x + b
    a = sigmoid(z)

    # 复合函数求导:de/dz
    grad_z = -2 * (y - a) * a * (1 - a)

    # 链式法则展开
    dw = np.mean(grad_z * x)  # de/dw = de/dz * dz/dw
    db = np.mean(grad_z)      # de/db = de/dz * dz/db

    # 梯度下降更新
    w -= lr * dw
    b -= lr * db

    if epoch % 100 == 0:
        loss = np.mean((y - a) ** 2)
        print(f"Epoch {epoch}: loss = {loss:.4f}, w = {w:.4f}, b = {b:.4f}")

运行后会看到 wb 逐渐接近真实值 2.0-1.0,损失也逐渐下降。


7. 关键理解总结

概念 含义
预测函数 ŷ = wx + b(线性)或 ŷ = σ(wx + b)(分类)
损失函数 衡量预测与真实的差距,如平方误差 e = (y - ŷ)²
梯度 损失函数对参数的偏导数,指明参数调整方向
链式法则 用于对复合函数(如 Sigmoid 嵌套线性函数)求导
梯度下降 沿梯度反方向更新参数,使损失逐步减小
学习率 控制每次更新的步长,需要合理设置

这就是梯度下降实现参数优化的全过程。


8. 推导步骤图谱

相关推荐
霸道流氓气质13 小时前
普通办公电脑基于 Ollama 的本地 AI 能力技术文档
人工智能·电脑
点云-激光雷达-Slam-三维牙齿13 小时前
速度起飞 笔记本电脑6G显卡llama运行Qwen3.6 35BA3B MTP 大模型
人工智能·python·电脑·llama
言乐614 小时前
Python游戏水平测试辅助系统
开发语言·python·游戏·django·pygame
Shockang18 小时前
AI 智能体安全沙盒实战
人工智能
yuhulkjv33519 小时前
Claude表格复制到word不再崩溃,AI导出鸭批量导出+格式无损一键搞定
人工智能·ai·c#·word·ai导出鸭
从零开始学习人工智能20 小时前
【踩坑实录】WSL2 解决 onnxruntime\-gpu ImportError: libcudart\.so\.13 无 CUDA13 运行库问题
python
大明者省20 小时前
WSL2 Ubuntu22.04 GPU训练环境配置指南
人工智能·算法·计算机视觉
抱抱宝20 小时前
Agent-study项目教程(03):手写 Mini-ReAct Agent(不依赖框架)
javascript·人工智能·gpt·react.js·prompt·agent
抱抱宝20 小时前
大模型应用开发教程08 | 构建完整 RAG 应用(Chroma/FAISS 实战)
人工智能·gpt·prompt·agent
美团技术团队21 小时前
KDD‘26 美团学术论文精选及KDD Cup‘26 DataAgents赛道冠军思路解读
人工智能