从零理解bp反向传播算法:图片分类任务中的梯度计算与实现

1次阅读
没有评论

共计 1540 个字符,预计需要花费 4 分钟才能阅读完成。

image.webp

一、从全连接网络到计算图

假设我们要处理 28×28 的 MNIST 手写数字图片,输入层就是 784 个像素点。一个最简单的三层神经网络结构如下:

从零理解 bp 反向传播算法:图片分类任务中的梯度计算与实现

  • 输入层:784 个神经元(对应图片展平的 784 像素)
  • 隐藏层:256 个神经元(使用 Sigmoid 激活)
  • 输出层:10 个神经元(对应 0 - 9 数字分类)

前向传播的计算图可以这样表示:

graph LR
    A[输入 X] -->|W1| B[隐藏层 Z1=W1X+b1]
    B -->|Sigmoid| C[A1=σ(Z1)]
    C -->|W2| D[输出层 Z2=W2A1+b2]
    D -->|Softmax| E[预测值 Y_hat]

二、反向传播的数学拆解

以交叉熵损失函数为例,我们需要求损失 L 对 W1、b1、W2、b2 的偏导。关键步骤如下:

  1. 输出层梯度计算:
  2. ∂L/∂Z2 = Y_hat – Y(Softmax 与交叉熵的求导特例)

  3. 隐藏层参数更新:

  4. ∂L/∂W2 = (∂L/∂Z2) · A1.T
  5. ∂L/∂b2 = sum(∂L/∂Z2, axis=0)

  6. 链式法则传递到前层:

  7. ∂L/∂A1 = W2.T · (∂L/∂Z2)
  8. ∂L/∂Z1 = ∂L/∂A1 ⊙ σ'(Z1)(⊙表示逐元素乘)

  9. 输入层参数更新:

  10. ∂L/∂W1 = (∂L/∂Z1) · X.T
  11. ∂L/∂b1 = sum(∂L/∂Z1, axis=0)

三、Python 实现关键代码

import numpy as np

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

def sigmoid_derivative(x):
    return x * (1 - x)  # 注意这里 x 应传入 sigmoid 的输出值

# 前向传播
def forward(X, W1, b1, W2, b2):
    Z1 = np.dot(W1, X.T) + b1
    A1 = sigmoid(Z1)
    Z2 = np.dot(W2, A1) + b2
    # Softmax 实现略
    return A1, Z2

# 反向传播
def backward(X, Y, A1, Z2, W2):
    m = X.shape[0]  # 样本数

    # 输出层梯度
    dZ2 = (Y_hat - Y) / m  # 注意除以 m 实现均值
    dW2 = np.dot(dZ2, A1.T)
    db2 = np.sum(dZ2, axis=1, keepdims=True)

    # 隐藏层梯度
    dA1 = np.dot(W2.T, dZ2)
    dZ1 = dA1 * sigmoid_derivative(A1)
    dW1 = np.dot(dZ1, X)
    db1 = np.sum(dZ1, axis=1, keepdims=True)

    return dW1, db1, dW2, db2

四、避坑指南

  1. 学习率与梯度消失:
  2. 当使用 Sigmoid 时,导数最大值仅 0.25,多层连乘会导致梯度指数级减小
  3. 建议初始学习率设为 0.1,配合梯度裁剪(gradient clipping)

  4. 输入归一化:

  5. MNIST 像素值原始范围 0 -255,应缩放至 0 -1
  6. 实践发现,归一化后收敛步数减少约 40%

  7. 可视化调试:

    from torch.utils.tensorboard import SummaryWriter
    writer = SummaryWriter()
    
    # 在训练循环中添加
    writer.add_scalar('Loss/train', loss.item(), epoch)
    writer.add_histogram('weights/W1', W1, epoch)

五、延伸思考

  1. ReLU 的反向传播调整:
  2. 正向传播时:max(0, x)
  3. 反向传播时:梯度为 1(当 x >0)或 0(当 x≤0)
  4. 代码实现只需修改激活函数导数部分

  5. CNN 的反向传播特点:

  6. 卷积核的梯度计算需要用到转置卷积操作
  7. 池化层的梯度上采样需要记录 max pooling 的位置
  8. 参数量大幅减少但计算复杂度更高

最后留个实践作业:尝试修改代码实现以下功能:
1. 增加动量(momentum)优化
2. 用验证集实现早停(early stopping)
3. 添加 L2 正则化项

正文完
 0
评论(没有评论)