跳转至

机器学习中常用函数的导数推导

本文将从最基础的初等函数出发,一路推导到 Transformer 核心组件 Attention 和残差网络的反向传播,力求做到每一步清晰可循,每一个结果都有工程注解

基础初等函数

幂函数

\[f(x) = x^n \quad \Rightarrow \quad f'(x) = nx^{n-1}\]

证明(导数定义):

\[ \begin{aligned} f'(x) &= \lim_{h\to 0}\frac{(x+h)^n - x^n}{h} \\ &= \lim_{h\to 0}\frac{x^n + nx^{n-1}h + \frac{n(n-1)}{2}x^{n-2}h^2 + \cdots + h^n - x^n}{h} \\ &= \lim_{h\to 0}\left(nx^{n-1} + \frac{n(n-1)}{2}x^{n-2}h + \cdots + h^{n-1}\right) \\ &= nx^{n-1} \end{aligned} \]

二项式展开中,除第一项外所有项都至少含一个因子 \(h\),因此极限后全部消失。当 \(n\) 为任意实数时,需用对数微分法证明,但结论相同。

工程意义

权重衰减(\(L_2\) 正则化)的梯度 \(\frac{\partial}{\partial w}(\frac{\lambda}{2}w^2)=\lambda w\) 就来源于此。

指数函数

\[f(x) = e^x \quad \Rightarrow \quad f'(x) = e^x\]

证明(利用自然对数的定义):

方法一:

利用 \(\lim_{h\to 0}\frac{e^h-1}{h}=1\) 这一基本极限:

\[ f'(x) = \lim_{h\to 0}\frac{e^{x+h}-e^x}{h} = e^x \lim_{h\to 0}\frac{e^h-1}{h} = e^x \]

方法二:

利用 \(e^x\) 的泰勒展开 \(e^h = 1 + h + \frac{h^2}{2!} + \cdots\)

\[ \begin{aligned} f'(x) &= \lim_{h\to 0}\frac{e^{x+h}-e^x}{h} = e^x \lim_{h\to 0}\frac{e^h-1}{h} \\ &= e^x \lim_{h\to 0}\frac{(1+h+\frac{h^2}{2!}+\cdots)-1}{h} = e^x \lim_{h\to 0}\left(1+\frac{h}{2!}+\cdots\right) = e^x \end{aligned} \]

工程意义

对于一般底数 \(a>0\)\((a^x)' = a^x \ln a\),这是机器学习中学习率衰减、早停等指数调度策略的理论基础。

对数函数

\[f(x) = \ln x \quad \Rightarrow \quad f'(x) = \frac{1}{x}\]

证明:

利用重要极限 \(\lim_{t\to 0}(1+t)^{1/t}=e\)

\[ \begin{aligned} f'(x) &= \lim_{h\to 0}\frac{\ln(x+h)-\ln x}{h} \\ &= \lim_{h\to 0}\frac{1}{h}\ln\left(1+\frac{h}{x}\right) \\ &= \lim_{h\to 0}\ln\left[\left(1+\frac{h}{x}\right)^{\frac{x}{h}}\right]^{\frac{1}{x}} \\ &= \frac{1}{x}\ln\left[\lim_{t\to 0}(1+t)^{1/t}\right] = \frac{1}{x}\ln e = \frac{1}{x} \end{aligned} \]

其中 \(t = h/x\)

工程意义

交叉熵损失的梯度中包含 \(1/\hat{y}\) 项,其根源就是对数的导数。\(\ln\) 的“乘法变加法”特性也使其成为极大似然估计中处理概率乘积的标准工具。

神经网络激活函数

激活函数是神经网络的非线性之源。这里将逐个剖析其导数,重点关注梯度流动特性

Sigmoid(Logistic 函数)

\[\sigma(x) = \frac{1}{1+e^{-x}} \quad \Rightarrow \quad \sigma'(x) = \sigma(x)(1-\sigma(x))\]

证明:(链式传播法则)

\[ \begin{aligned} \sigma'(x) &= \frac{d}{dx}(1+e^{-x})^{-1} \\ &= -(1+e^{-x})^{-2} \cdot (-e^{-x}) \\ &= \frac{e^{-x}}{(1+e^{-x})^2} \\ &= \frac{1}{1+e^{-x}} \cdot \frac{e^{-x}}{1+e^{-x}} \\ &= \sigma(x) \cdot \frac{1+e^{-x}-1}{1+e^{-x}} \\ &= \sigma(x)\big(1-\sigma(x)\big) \end{aligned} \]

梯度特性与工程技巧

\(x \to +\infty\)\(\sigma \to 1\),导数 \(\to 0\);当 \(x \to -\infty\)\(\sigma \to 0\),导数 \(\to 0\)饱和区梯度趋零是 Sigmoid 在深层网络中被弃用的主因。

前向传播时缓存 \(\sigma(x)\) 即可完成反向传播,无需重复计算指数,这也是 PyTorch 中 Sigmoid 反向 Kernel 的优化策略。

Tanh(双曲正切)

\[\tanh(x) = \frac{e^x - e^{-x}}{e^x + e^{-x}} \quad \Rightarrow \quad \tanh'(x) = 1 - \tanh^2(x)\]

证明(商的法则):

\(u=e^x-e^{-x}\)\(v=e^x+e^{-x}\),则 \(u'=e^x+e^{-x}=v\)\(v'=e^x-e^{-x}=u\)

\[ \tanh'(x) = \frac{u'v - uv'}{v^2} = \frac{v^2 - u^2}{v^2} = 1 - \left(\frac{u}{v}\right)^2 = 1 - \tanh^2(x) \]

与 Sigmoid 的关系

\(\tanh(x) = 2\sigma(2x) - 1\),其导数为 \(4\sigma(2x)(1-\sigma(2x))\)。相比 Sigmoid,Tanh 以零为中心,梯度通常更稳定。

ReLU(Rectified Linear Unit)

\[\text{ReLU}(x) = \max(0, x) \quad \Rightarrow \quad \text{ReLU}'(x) = \begin{cases} 1, & x > 0 \\ 0, & x < 0 \end{cases}\]

证明:分段线性的直接推论。

  • \(x > 0\)\(\text{ReLU}(x) = x\),导数为 \(1\)
  • \(x < 0\)\(\text{ReLU}(x) = 0\),导数为 \(0\)
  • \(x = 0\):不可导。实际中采用次梯度(subgradient),取 \([0,1]\) 中的任意值,代码中通常取 \(0\)\(1\)(PyTorch 取 \(0\))。

梯度特性

正区间导数为 1,梯度恒等传播,这是 ReLU 在深度网络中成功的核心原因。负区间置零带来稀疏性,但也导致 Dead ReLU 问题(神经元永久死亡)。

Leaky ReLU

\[\text{LeakyReLU}(x) = \begin{cases} x, & x \ge 0 \\ \alpha x, & x < 0 \end{cases} \quad \Rightarrow \quad \text{LeakyReLU}'(x) = \begin{cases} 1, & x \ge 0 \\ \alpha, & x < 0 \end{cases}\]

其中 \(\alpha\) 通常取 \(0.01\)(或作为可学习参数,即 PReLU)。

证明:同 ReLU,分段求导。\(x=0\) 处不可导,次梯度为 \([\alpha, 1]\)

设计动机

负区间保留小梯度 \(\alpha\),避免神经元完全“死亡”,同时保持了 ReLU 的主要优点。

Softmax(多元函数)

对于向量 \(\mathbf{z} \in \mathbb{R}^K\),Softmax 将 \(K\) 个实数映射为一个概率分布:

\[\sigma_i = \frac{e^{z_i}}{\sum_{j=1}^K e^{z_j}}, \quad i=1,\dots,K\]

其导数为 Jacobian 矩阵 \(J \in \mathbb{R}^{K \times K}\)

\[\frac{\partial \sigma_i}{\partial z_j} = \sigma_i(\delta_{ij} - \sigma_j)\]

其中 \(\delta_{ij}\) 是 Kronecker delta(\(i=j\) 时为 1,否则为 0)。

证明(分两种情况):

情况 1:\(i = j\)

\[ \frac{\partial \sigma_i}{\partial z_i} = \frac{e^{z_i}\sum_k e^{z_k} - e^{z_i}\cdot e^{z_i}}{(\sum_k e^{z_k})^2} = \frac{e^{z_i}}{\sum_k e^{z_k}} \cdot \frac{\sum_k e^{z_k} - e^{z_i}}{\sum_k e^{z_k}} = \sigma_i(1 - \sigma_i) \]

情况 2:\(i \neq j\)

\[ \frac{\partial \sigma_i}{\partial z_j} = \frac{0 \cdot \sum_k e^{z_k} - e^{z_i} \cdot e^{z_j}}{(\sum_k e^{z_k})^2} = -\frac{e^{z_i}}{\sum_k e^{z_k}} \cdot \frac{e^{z_j}}{\sum_k e^{z_k}} = -\sigma_i \sigma_j \]

合并两种情况即得 \(\sigma_i(\delta_{ij} - \sigma_j)\)

工程意义

该 Jacobian 具有行和为零的性质(\(\sum_j \frac{\partial \sigma_i}{\partial z_j}=0\)),这是因为 Softmax 输出之和恒为 1,任一输出的增加必然伴随其他输出的减少。

GELU(Gaussian Error Linear Unit)

GELU 是 GPT、BERT、ViT 等现代 Transformer 架构的标配激活函数

\[\text{GELU}(x) = x \cdot \Phi(x) = x \cdot \frac{1}{2}\left[1 + \operatorname{erf}\left(\frac{x}{\sqrt{2}}\right)\right]\]

其中 \(\Phi(x)\) 是标准正态分布的累积分布函数(CDF),\(\phi(x) = \frac{1}{\sqrt{2\pi}}e^{-x^2/2}\) 是其概率密度函数(PDF)。

导数:

\[\text{GELU}'(x) = \Phi(x) + x \cdot \phi(x)\]

证明:直接使用乘积法则:

\[ \frac{d}{dx}\text{GELU}(x) = \frac{d}{dx}\left[x \cdot \Phi(x)\right] = \Phi(x) + x \cdot \Phi'(x) = \Phi(x) + x \cdot \phi(x) \]

展开为完整形式:

\[ \text{GELU}'(x) = \frac{1}{2}\left[1 + \operatorname{erf}\left(\frac{x}{\sqrt{2}}\right)\right] + \frac{x}{\sqrt{2\pi}} e^{-x^2/2} \]

考虑到数值稳定性,实际实现中,\(\Phi(x)\) 通常用 erf 计算。对于极端负值,可采用近似公式避免下溢:

\[\Phi(x) \approx \frac{1}{2}\left[1 + \operatorname{erf}\left(\frac{x}{\sqrt{2}}\right)\right]\]

二阶导数

\[\text{GELU}''(x) = 2\phi(x) - x^2 \phi(x) = \phi(x)(2 - x^2)\]

这表明 GELU 在 \(x=\pm\sqrt{2}\) 附近有拐点,与 ReLU 的“硬门控”不同,GELU 提供了平滑的非线性

常见近似:由于 erf 计算较昂贵,实际中常用 tanh 近似:

\[\text{GELU}(x) \approx 0.5x\left[1 + \tanh\left(\sqrt{\frac{2}{\pi}}(x + 0.044715x^3)\right)\right]\]

该近似的导数可相应推导,但一般框架中直接对近似式自动微分即可。

SiLU / Swish

\[\operatorname{SiLU}(x) = x \cdot \sigma(x) = \frac{x}{1+e^{-x}}\]

其中 \(\sigma(x)\) 是 Sigmoid 函数。

导数:

\[f'(x) = \sigma(x) \cdot \left[1 + x(1 - \sigma(x))\right] = \sigma(x) + x \cdot \sigma(x)(1 - \sigma(x))\]

证明

\[ \begin{aligned} f'(x) &= \sigma(x) + x \cdot \sigma'(x) \\ &= \sigma(x) + x \cdot \sigma(x)(1-\sigma(x)) \\ &= \sigma(x)\left[1 + x - x\sigma(x)\right] \\ &= \sigma(x)\left[1 + x(1-\sigma(x))\right] \end{aligned} \]

Swish特性

该 Jacobian 具有行和为零的性质(\(\sum_j \frac{\partial \sigma_i}{\partial z_j}=0\)),这是因为 Softmax 输出之和恒为 1,任一输出的增加必然伴随其他输出的减少。

激活函数导数速查表

激活函数 导数 特点
\(\sigma(x)\) \(\sigma(x)(1-\sigma(x))\) 饱和区梯度消失
\(\tanh(x)\) \(1-\tanh^2(x)\) 零中心,梯度范围 \([0,1]\)
ReLU \(1_{x>0}\)(次梯度) 正区恒等传播,负区死亡
Leaky ReLU \(1_{x\ge0} + \alpha \cdot 1_{x<0}\) 缓解死亡 ReLU
GELU \(\Phi(x)+x\phi(x)\) Transformer 标配,平滑门控
Swish \(\sigma(x)[1+x(1-\sigma(x))]\) 自门控,平滑非单调

损失函数

损失函数是优化目标,其梯度决定了参数更新的方向和大小。

均方误差(MSE)

定义(批量版本):

\[L = \frac{1}{N}\sum_{i=1}^N (y_i - \hat{y}_i)^2\]

对预测值 \(\hat{y}_i\) 的导数:

\[\frac{\partial L}{\partial \hat{y}_i} = \frac{2}{N}(\hat{y}_i - y_i)\]

证明:令 \(d_i = y_i - \hat{y}_i\),则 \(\frac{\partial}{\partial \hat{y}_i} d_i^2 = 2d_i \cdot (-1) = 2(\hat{y}_i - y_i)\),再除以 \(N\)

注意

有些实现中分母为 \(2N\),目的是消去系数 2 使梯度表达更简洁,但梯度下降的收敛点不变。

对权重参数的拓展:若 \(\hat{y} = Wx + b\),则链式法则给出: $\(\frac{\partial L}{\partial W} = \frac{2}{N}(\hat{y} - y)x^T\)$

二元交叉熵(BCE,配合 Sigmoid)

定义(单样本):

\[L = -\left[y\ln\hat{y} + (1-y)\ln(1-\hat{y})\right], \quad \hat{y} = \sigma(z)\]

其中 \(y \in \{0,1\}\) 为真实标签,\(\hat{y}\) 为预测概率。

对 logit \(z\) 的导数:

\[\frac{\partial L}{\partial z} = \hat{y} - y\]

证明:链式法则分两步。

先求对 \(\hat{y}\) 的导数: $\(\frac{\partial L}{\partial \hat{y}} = -\frac{y}{\hat{y}} + \frac{1-y}{1-\hat{y}} = \frac{\hat{y} - y}{\hat{y}(1-\hat{y})}\)$

再乘上 Sigmoid 导数 \(\hat{y}(1-\hat{y})\)

\[ \frac{\partial L}{\partial z} = \frac{\hat{y} - y}{\hat{y}(1-\hat{y})} \cdot \hat{y}(1-\hat{y}) = \hat{y} - y \]

\(\frac{\partial L}{\partial \hat{y}}\)\(\hat{y}\to 0\)\(\hat{y}\to 1\) 时趋于无穷,但乘以 Sigmoid 导数后变为有限值 \(\hat{y}-y\),从而出现梯度抵消现象。这种 “分母消去” 是数学上的精确抵消,保证了数值稳定性。

工程意义

这正是为什么二元分类中,Sigmoid + BCE 组合比 MSE + Sigmoid 更优——前者梯度在饱和区不消失(当 \(y=1\)\(\hat{y}\approx 0\) 时梯度约 \(-1\)),而后者梯度会随 Sigmoid 导数趋零。

多元交叉熵(配合 Softmax)

定义:

\[L = -\sum_{i=1}^K y_i \ln \hat{y}_i, \quad \hat{y}_i = \operatorname{softmax}(z_i)\]

其中 \(y\) 为 one-hot 标签向量(或概率分布),\(\sum_i y_i = 1\)

对 logit \(z_j\) 的导数:

\[\frac{\partial L}{\partial z_j} = \hat{y}_j - y_j\]

证明:利用 Softmax 的 Jacobian \(\frac{\partial \hat{y}_i}{\partial z_j} = \hat{y}_i(\delta_{ij} - \hat{y}_j)\)

\[ \begin{aligned} \frac{\partial L}{\partial z_j} &= -\sum_{i=1}^K y_i \frac{1}{\hat{y}_i} \cdot \frac{\partial \hat{y}_i}{\partial z_j} \\ &= -\sum_{i=1}^K y_i \frac{1}{\hat{y}_i} \cdot \hat{y}_i(\delta_{ij} - \hat{y}_j) \\ &= -\sum_{i=1}^K y_i(\delta_{ij} - \hat{y}_j) \\ &= -y_j + \hat{y}_j \sum_{i=1}^K y_i \\ &= \hat{y}_j - y_j \quad (\text{因 } \sum_i y_i = 1) \end{aligned} \]

Softmax 的 Jacobian 和 Cross-Entropy 对 \(\hat{y}\) 的导数相互抵消,留下最简形式 \(\hat{y} - y\)。这就是为什么几乎所有分类框架都使用这个组合。

损失函数导数速查表

损失函数 配合激活 对 logit 的梯度
MSE 线性 \(\frac{2}{N}(\hat{y}-y)\)
BCE Sigmoid \(\hat{y}-y\)
CE Softmax \(\hat{y}-y\)

观察:三种最常用的损失函数(配合对应激活)的最终梯度形式都是 预测值 - 目标值。这并非巧合,而是指数族分布与对应连接函数的典范性质。

归一化层

归一化层是现代深度网络的“隐形冠军”。

Layer Normalization

前向传播:对输入向量 \(\mathbf{x} \in \mathbb{R}^d\)(单样本所有特征):

\[\mu = \frac{1}{d}\sum_{i=1}^d x_i, \quad \sigma^2 = \frac{1}{d}\sum_{i=1}^d (x_i - \mu)^2\]
\[\hat{x}_i = \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}}, \quad y_i = \gamma_i \hat{x}_i + \beta_i\]

其中 \(\gamma, \beta\) 为可学习的仿射参数,\(\epsilon\) 为数值稳定小量。

对输入 \(x_j\) 的偏导数:

\[\frac{\partial y_i}{\partial x_j} = \frac{\gamma_i}{\sqrt{\sigma^2+\epsilon}}\left(\delta_{ij} - \frac{1}{d} - \frac{(x_i-\mu)(x_j-\mu)}{d(\sigma^2+\epsilon)}\right)\]

证明(含中间变量推导):

\(v = \sigma^2 + \epsilon\)\(m = \mu\)。首先有:

\[\frac{\partial m}{\partial x_j} = \frac{1}{d}, \quad \frac{\partial v}{\partial x_j} = \frac{\partial \sigma^2}{\partial x_j} = \frac{2}{d}(x_j - m)\]

\(\hat{x}_i = (x_i - m)v^{-1/2}\) 求导:

\[ \frac{\partial \hat{x}_i}{\partial x_j} = \left(\delta_{ij} - \frac{1}{d}\right)v^{-1/2} + (x_i-m)\left(-\frac{1}{2}\right)v^{-3/2} \cdot \frac{2}{d}(x_j-m) \]
\[ = v^{-1/2}\left[\delta_{ij} - \frac{1}{d} - \frac{(x_i-m)(x_j-m)}{dv}\right] \]

最后乘以 \(\gamma_i\) 即得完整表达式。

LayerNorm 的求和维度是特征维度,对 batch 中每个样本独立计算。因此其导数表达式中 \(1/d\) 项来自特征维平均,而非 batch 维平均。

从计算图视角来看反向传播时,需要从 \(\frac{\partial L}{\partial y_i}\) 依次传递到 \(\hat{x}_i\)\(\sigma^2\)\(\mu\),最后到 \(x\)。直接使用上述闭式公式虽然简洁,但在实际实现中,为节省内存通常会分段计算。

Batch Normalization(训练时)

前向传播:对 batch 中 \(N\) 个样本的同一通道/特征:

\[\mu_B = \frac{1}{N}\sum_{k=1}^N x_k, \quad \sigma_B^2 = \frac{1}{N}\sum_{k=1}^N (x_k - \mu_B)^2\]
\[\hat{x}_i = \frac{x_i - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \quad y_i = \gamma \hat{x}_i + \beta\]

\(x_j\) 的偏导:

\[\frac{\partial y_i}{\partial x_j} = \frac{\gamma}{\sqrt{\sigma_B^2+\epsilon}}\left(\delta_{ij} - \frac{1}{N} - \frac{(x_i-\mu_B)(x_j-\mu_B)}{N(\sigma_B^2+\epsilon)}\right)\]

它与 Layer Normalization 形式完全一致,唯一区别是求和维度从特征维变成了 batch 维。这正是 BN 与 LN 的数学对称性。但是二者在推理时的差异:

推理时 BN 使用全局统计量 \(\mu_{\text{global}}\)\(\sigma_{\text{global}}^2\)(训练时的移动平均),此时 \(\hat{x}_i\) 与 batch 内其他样本无关,梯度形式退化为:

$\(\frac{\partial y_i}{\partial x_i} = \frac{\gamma}{\sqrt{\sigma_{\text{global}}^2+\epsilon}}\)$ 交叉项(\(i\neq j\))消失。

归一化层导数结构速览

归一化类型 求和维度 分母项 \(1/d\) 来源 梯度中的交叉项
LayerNorm 特征维 特征数 \(d\) \(\frac{(x_i-\mu)(x_j-\mu)}{d(\sigma^2+\epsilon)}\)
BatchNorm Batch 维 Batch 大小 \(N\) \(\frac{(x_i-\mu_B)(x_j-\mu_B)}{N(\sigma_B^2+\epsilon)}\)

为什么 BN 对 batch size 敏感?

从上式可见,当 \(N\) 较小时,\(\frac{1}{N}\) 项和交叉项方差增大,导致梯度估计不稳定。LN 不存在此问题。

Attention 机制

Attention 是 Transformer 的核心,其反向传播涉及矩阵微分

Scaled Dot-Product Attention

前向传播

\[S = \frac{QK^T}{\sqrt{d_k}}, \quad A = \operatorname{softmax}(S) \quad (\text{按行 softmax})\]
\[O = A V\]

其中 \(Q \in \mathbb{R}^{n \times d_k}\)\(K \in \mathbb{R}^{m \times d_k}\)\(V \in \mathbb{R}^{m \times d_v}\),输出 \(O \in \mathbb{R}^{n \times d_v}\)

\(V\) 的梯度(最直接):

\[\frac{\partial L}{\partial V} = A^T \frac{\partial L}{\partial O}\]

证明\(O_{ij} = \sum_k A_{ik}V_{kj}\),故 \(\frac{\partial O_{ij}}{\partial V_{kj}} = A_{ik}\),转置后即得。

\(Q\) 的梯度

设上游梯度 \(G = \frac{\partial L}{\partial A} \in \mathbb{R}^{n \times m}\)(即损失对注意力权重矩阵的导数,从输出 \(O\) 回传得到)。

利用行 Softmax 的 Jacobian(第 \(i\) 行内部):

\[\frac{\partial A_{ij}}{\partial S_{ik}} = A_{ij}(\delta_{jk} - A_{ik})\]

(注意 \(k\) 遍历的是同一行的列索引,不同行之间独立。)

因此:

\[ \frac{\partial L}{\partial S_{ij}} = \sum_k \frac{\partial L}{\partial A_{ik}} \cdot \frac{\partial A_{ik}}{\partial S_{ij}} = \sum_k G_{ik} \cdot A_{ik}(\delta_{jk} - A_{ij}) \]

化简(注意 \(\sum_k G_{ik}A_{ik}\delta_{jk} = G_{ij}A_{ij}\)):

\[ \frac{\partial L}{\partial S_{ij}} = A_{ij}\left(G_{ij} - \sum_k G_{ik}A_{ik}\right) \]

矩阵形式

\[\frac{\partial L}{\partial S} = A \odot \left(G - \text{rowsum}(A \odot G) \cdot \mathbf{1}^T\right)\]

其中: - \(\odot\) 是 Hadamard 积(逐元素相乘) - \(\text{rowsum}(A \odot G) \in \mathbb{R}^n\),第 \(i\) 项为 \(\sum_k A_{ik}G_{ik}\) - \(\mathbf{1}^T \in \mathbb{R}^{1 \times m}\) 是全 1 行向量(广播)

下面是简洁写法:

\(D = \text{diag}((A \odot G) \cdot \mathbf{1}_m)\),即 \(D_{ii} = \sum_k A_{ik}G_{ik}\),则:

\[\frac{\partial L}{\partial S} = A \odot (G - D \cdot \mathbf{1}^T) = A \odot G - A \odot (D \cdot \mathbf{1}^T)\]

在 FlashAttention 的论文中,这个公式被称为 "gradient of softmax with respect to logits",是 CUDA Kernel 实现的核心。

最后通过链式法则得到对 \(Q\)\(K\) 的梯度

\[S_{ij} = \frac{1}{\sqrt{d_k}} \sum_t Q_{it} K_{jt}\]

因此:

\[\frac{\partial L}{\partial Q_{it}} = \frac{1}{\sqrt{d_k}} \sum_j \frac{\partial L}{\partial S_{ij}} \cdot K_{jt} \Longrightarrow \frac{\partial L}{\partial Q} = \frac{1}{\sqrt{d_k}} \cdot \frac{\partial L}{\partial S} \cdot K\]
\[\frac{\partial L}{\partial K_{jt}} = \frac{1}{\sqrt{d_k}} \sum_i \frac{\partial L}{\partial S_{ij}} \cdot Q_{it} \Longrightarrow \frac{\partial L}{\partial K} = \frac{1}{\sqrt{d_k}} \cdot \left(\frac{\partial L}{\partial S}\right)^T \cdot Q\]

工程意义

上述矩阵公式是 FlashAttention、xFormers 等高效 Attention 实现的数学基础。FlashAttention 的核心优化之一就是在不显式保存完整 \(S\) 矩阵(\(O(n^2)\) 内存)的情况下,通过分块计算和重计算来得到 \(\frac{\partial L}{\partial S}\),从而大幅减少 HBM 访问。

对因果 Attention(Causal Masking)的修改:若施加下三角掩码(\(S_{ij}=-\infty\)\(j>i\)),则对应位置 \(A_{ij}=0\)\(\frac{\partial L}{\partial S_{ij}}\) 在掩码区域强制为 0,只需将上述公式中的 \(A\) 替换为掩码后的版本即可。

链式法则与反向传播

全连接层的反向传播

设第 \(l\) 层的前向传播为:

\[z^{[l]} = W^{[l]} a^{[l-1]} + b^{[l]}, \quad a^{[l]} = \sigma(z^{[l]})\]

定义 误差项 \(\delta^{[l]} = \frac{\partial L}{\partial z^{[l]}}\),则反向传播的三个核心公式为:

1. 误差项回传:

\[\delta^{[l]} = \left((W^{[l+1]})^T \delta^{[l+1]}\right) \odot \sigma'(z^{[l]})\]

2. 权重梯度:

\[\frac{\partial L}{\partial W^{[l]}} = \delta^{[l]} (a^{[l-1]})^T\]

3. 偏置梯度:

\[\frac{\partial L}{\partial b^{[l]}} = \delta^{[l]}\]

证明(链式法则):

首先,从损失函数对 \(a^{[l]}\) 的梯度推导对 \(z^{[l]}\) 的梯度:

\[\delta^{[l]} = \frac{\partial L}{\partial z^{[l]}} = \frac{\partial L}{\partial a^{[l]}} \odot \frac{\partial a^{[l]}}{\partial z^{[l]}} = \frac{\partial L}{\partial a^{[l]}} \odot \sigma'(z^{[l]})\]

\(\frac{\partial L}{\partial a^{[l]}}\) 可由下一层传递:

\[ \frac{\partial L}{\partial a^{[l]}} = \frac{\partial L}{\partial z^{[l+1]}} \cdot \frac{\partial z^{[l+1]}}{\partial a^{[l]}} = (W^{[l+1]})^T \delta^{[l+1]} \]

合并即得误差回传公式。

\(W^{[l]}_{ij}\) 求导:因为 \(z^{[l]}_i = \sum_j W^{[l]}_{ij} a^{[l-1]}_j + b^{[l]}_i\),所以:

\[\frac{\partial L}{\partial W^{[l]}_{ij}} = \frac{\partial L}{\partial z^{[l]}_i} \cdot \frac{\partial z^{[l]}_i}{\partial W^{[l]}_{ij}} = \delta^{[l]}_i \cdot a^{[l-1]}_j\]

矩阵形式即 \(\delta^{[l]} (a^{[l-1]})^T\)

从计算图视角来看反向传播的本质是梯度沿计算图反向流动。每一层接收上游梯度 \(\delta^{[l+1]}\),乘以本地的 Jacobian(激活函数导数和权重转置),产生下游梯度 \(\delta^{[l]}\),同时计算出对本层参数的梯度。

在工程实践过程中,会采用这样的工程优化:实际框架(如 PyTorch)不会显式构建所有 Jacobian 矩阵,而是采用 vector-Jacobian product (VJP) 的方式逐层传播梯度,内存效率远高。

深度残差网络的核心反向传播

现代深度网络(ResNet、Transformer、Diffusion U-Net)极少是纯串行的 Layer1 -> Layer2 -> ...,而是充满了跨层的跳跃连接(Skip Connection)。其典型结构为:

\[y = x + F(x)\]

其中 \(x\) 是输入(或浅层特征),\(F(x)\) 是复杂的非线性变换(如卷积块、自注意力块),\(y\) 是该模块的输出。

残差连接的梯度公式

前向传播: $\(y = x + F(x)\)$

反向传播(对输入 \(x\) 的梯度): 设上游传来损失对 \(y\) 的梯度为 \(\frac{\partial L}{\partial y}\),则:

\[ \frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot \frac{\partial y}{\partial x} = \frac{\partial L}{\partial y} \cdot \left( I + \frac{\partial F(x)}{\partial x} \right) \]

展开为矩阵形式(核心洞察)

\[ \frac{\partial L}{\partial x} = \underbrace{\frac{\partial L}{\partial y}}_{\text{直通梯度(恒等映射)}} + \underbrace{\frac{\partial L}{\partial y} \cdot \frac{\partial F(x)}{\partial x}}_{\text{非线性支路梯度}} \]

证明(雅可比矩阵视角)

\(F: \mathbb{R}^d \to \mathbb{R}^d\),其雅可比矩阵为 \(J_F = \frac{\partial F}{\partial x}\)。由于恒等映射 \(I(x)=x\) 的雅可比矩阵为单位矩阵 \(I\),加法操作的雅可比为:

\[J_y = I + J_F\]

根据链式法则,梯度反向传播(向量对向量的左乘)为:

\[\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot (I + J_F) = \frac{\partial L}{\partial y} + \frac{\partial L}{\partial y} \cdot J_F\]

ML 意义:梯度高速公路(Gradient Highway)

这是深度网络能够堆叠上千层的数学命脉

  1. 消除梯度消失:即使深层网络中 \(F(x)\) 的雅可比 \(J_F\) 变得极小(饱和)或条件数极差,加号左边依然保留着干净的 \(\frac{\partial L}{\partial y}\)。这意味着浅层总能直接接收到来自深层的无损梯度,彻底解决了 Sigmoid/Tanh 在深层中的梯度消失问题。

  2. 特征复用:浅层特征可以通过“直通管道”直接出现在深层输出中,让网络退化为恒等映射(\(F \to 0\)\(y=x\)),保证了深层网络至少不劣于浅层网络。

带投影(Projection)的残差连接

在实际代码中(如 ViT 或 ResNet 下采样),如果输入 \(x\) 和输出 \(y\) 的维度不匹配,需要加一个线性投影 \(W_s\)

\[y = W_s x + F(x)\]

此时的梯度变为:

\[ \frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot W_s + \frac{\partial L}{\partial y} \cdot \frac{\partial F(x)}{\partial x} \]

注意

\(W_s\) 不是单位矩阵时,直通路径的梯度 \(\frac{\partial L}{\partial y} \cdot W_s\) 可能会改变尺度。如果 \(W_s\) 的条件数很大,仍可能引发梯度不稳定。这就是为什么现代架构(如原始 ResNet)在维度匹配时强制使用恒等映射(\(W_s=I\)),而不使用 1x1 卷积做投影,除非必须下采样。

Post-Norm 与 Pre-Norm 结构的梯度差异(Transformer 关键)

在 Transformer 中,LayerNorm 和残差连接有两种排列方式,它们的梯度行为截然不同:

(1)Post-Norm(原始 Transformer):

\[y = \text{LN}(x + F(x))\]

(2)Pre-Norm(GPT / 现代 Transformer 标配):

\[y = x + F(\text{LN}(x))\]

梯度分析(为何 Pre-Norm 更优):

对于 Post-Norm,由于 LayerNorm 包裹了整个残差块,反向传播时:

\[\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot \frac{\partial \text{LN}}{\partial (x+F)} \cdot \left(I + J_F\right)\]

即使 \((I+J_F)\) 很好,但乘以 \(\frac{\partial \text{LN}}{\partial z}\)(LayerNorm 的雅可比,包含减去均值和除以方差的操作)会破坏恒等映射的精确性。在深层网络中,这个额外的归一化梯度会导致信号衰减。

对于 Pre-Norm

\[\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot \left(I + \frac{\partial F(\text{LN}(x))}{\partial x}\right)\]

加号左边依然是完美的单位矩阵 \(I\),不受 LayerNorm 影响。这意味着梯度可以无损地从输出直接流回输入。这正是 GPT、BERT 等模型坚持使用 Pre-Norm 架构的核心数学原因。

残差网络反向传播速查表

结构类型 前向公式 误差项递推 \(\delta^{[l]}\) 梯度特性
普通全连接 \(a^{[l+1]} = \sigma(W a^{[l]} + b)\) \((W^T\delta^{[l+1]}) \odot \sigma'\) 权重连乘,易消失/爆炸
残差直连 \(a^{[l+1]} = a^{[l]} + F(a^{[l]})\) \(\delta^{[l+1]} + \delta^{[l+1]}J_F\) 恒等路径保底,永不消失
带投影残差 \(a^{[l+1]} = W_s a^{[l]} + F(a^{[l]})\) \(\delta^{[l+1]}W_s + \delta^{[l+1]}J_F\) 依赖 \(W_s\) 尺度,不如恒等稳定
Pre-Norm \(a^{[l+1]} = a^{[l]} + F(\text{LN}(a^{[l]}))\) \(\delta^{[l+1]} + \delta^{[l+1]}J_{F\circ LN}\) LN 不影响直通路径,最稳定
Post-Norm \(a^{[l+1]} = \text{LN}(a^{[l]} + F(a^{[l]}))\) \(\delta^{[l+1]} J_{\text{LN}} (I + J_F)\) 直通路径被破坏,深层易梯度消失

设计深度网络时,务必保证存在一条或多条完全不经过激活函数和归一化层的"梯度高速公路"。这就是为什么所有现代大模型(LLM)都采用 Pre-Norm + 残差结构。

附录:完整导数速查表

函数 导数
\(x^n\) \(nx^{n-1}\)
\(e^x\) \(e^x\)
\(\ln x\) \(1/x\)
\(\sigma(x)\) \(\sigma(x)(1-\sigma(x))\)
\(\tanh(x)\) \(1-\tanh^2(x)\)
ReLU \(1_{x>0}\)(次梯度 \([0,1]\)
Leaky ReLU \(1_{x\ge0} + \alpha \cdot 1_{x<0}\)
GELU \(\Phi(x)+x\phi(x)\)
Swish \(\sigma(x)[1+x(1-\sigma(x))]\)
MSE \(\frac{2}{N}(\hat{y}-y)\)
BCE+Sigmoid \(\hat{y}-y\)
CE+Softmax \(\hat{y}-y\)
LayerNorm \(\frac{\gamma}{\sqrt{\sigma^2+\epsilon}}(\delta_{ij} - \frac{1}{d} - \frac{(x_i-\mu)(x_j-\mu)}{d(\sigma^2+\epsilon)})\)
Attention Softmax \(A \odot (G - \text{rowsum}(A\odot G)\mathbf{1}^T)\)
残差直连 \(\delta^{[l+1]} + \delta^{[l+1]}J_F\)

结语

从幂函数到 Attention,我们走完了一条从基础微积分到现代深度学习核心组件数学表达的完整路径:

基础(初等函数)→ 单点非线性(激活函数)→ 单点度量(损失函数)→ 层内依赖(归一化)→ 层间依赖(全连接 BP)→ 跨层依赖(深度残差)高阶矩阵结构(Attention)

这样从“点”到“线”到“跳跃”,完美覆盖了现代深度学习所有梯度传播场景。

回顾全文,可以发现几条贯穿始终的设计原则

  1. 梯度消失/爆炸的根源:激活函数的导数范围、权重矩阵的特征值、链式法则的累积效应——所有这些都体现在导数公式中。
  2. “优雅抵消”现象:Sigmoid+BCE、Softmax+CE 中,前向函数的复杂 Jacobian 与损失梯度精确相消,得到最简形式 \(\hat{y}-y\)。这是概率模型与指数族分布的内在数学之美。
  3. 工程与数学的映射:FlashAttention 的 Kernel 优化、BatchNorm 的 \(\epsilon\) 添加、GELU 的近似公式——每个工程决策背后都有明确的数学动机。

评论