Alisa 的 LLMs 手册
本文是一本系统的大语言模型技术手册,涵盖从神经网络基础(MLP、激活函数、梯度、反向传播)到现代Transformer架构(注意力机制、RMSNorm、SwiGLU、RoPE)、训练优化(缩放定律、学习率调度)、推理优化(KV缓存、投机解码、Flash Attention)、后训练(RLHF、PPO、DPO、GRPO)以及并行计算(数据/模型/流水线/张量并行)等核心知识,包含大量数学公式推导和代码实现。
神经网络基础
多层感知机
多层感知机:全连接网络,包含输入层、至少一个隐藏层和一个输出层
通常与“前馈网络”同义使用,尽管FFN在技术上是一个更广泛的类别,指信息在其中单向流动的网络。
单个神经元计算其输入的加权和,加上偏置,并将结果通过激活函数传递。
$\mathbf x\in\mathbb R^n$ 是输入向量(上一层的激活值)
$\mathbf w\in\mathbb R^n$ 是权重向量(进入神经元的边权重)
$b\in\mathbb R$ 是偏置
$f$ 是激活函数
$y=f\left(\sum_{i=1}^n w_i x_i+b\right)=f(\mathbf w^\top\mathbf x +b)$
一个具有 $n_\text{in}$ 个输入和 $n_\text{out}$ 个神经元的层可以通过矩阵乘法计算。
所以 $\mathbf x\in\mathbb R^{n_\text{in}}$(列向量)
将所有权重向量堆叠成一个权重矩阵 $W\in\mathbb R^{n_\text{out}\times n_\text{in}}$
每一行是进入单个神经元的权重
将偏置堆叠成向量 $\mathbf b\in\mathbb R^{n_\text{out}}$
输出隐藏状态的形状为 $\mathbf h\in\mathbb R^{n_\text{out}}$
$\mathbf h=f(W\mathbf x+\mathbf b)$
在实践中,我们一次处理一批 $m$ 个输入!
在这种情况下,我们将输入排列为矩阵 $X \in\mathbb R^{m\times n_\text{in}}$ 的行
通常将 $W$ 的形状改为 $\mathbb R^{n_\text{in}\times n_\text{out}}$
每一列是进入单个神经元的权重
该层变为
$H=f(XW+\mathbf b)$
其中 $\mathbf b$ 广播为形状 $m\times n_\text{out}$
用数学符号表示,线性层接收 $X\in\mathbb R^{m\times n_\text{in}}$,并应用 $W\in\mathbb R^{n_\text{in}\times n_\text{out}}$ 作为 $XW+b$。
在 PyTorch 中,权重矩阵 $W$ 实际上存储为 $n_\text{out}\times n_\text{in}$。前向传播时转置 $W$,计算 $X @ W.T$ $(m, n_\text{in})\times (n_\text{in},n_\text{out})$。转置是免费的,因为它只改变步长。这样做是为了让 $W$ 的梯度自然成为 $n_\text{out}\times n_\text{in}$,与 $W$ 的形状匹配。
让我们做 $Z=XW+b$ 的反向传播
$$\frac{\partial L}{\partial X}=\frac{\partial L}{\partial Z}W^\top\quad (m,n_\text{out})\times (n_\text{out},n_\text{in})=(m,n_\text{in})\[1em] \frac{\partial L}{\partial W}=X^\top\frac{\partial L}{\partial Z}\quad (n_\text{in}, m)\times (m,n_\text{out})=(n_\text{in},n_\text{out})$$
相同的偏置 $\mathbf b\in\mathbb R^{n_\text{out}}$ 被加到每个样本上,每个样本产生自己的 $\mathbf b$ 梯度
因此这些梯度累加
$$\frac{\partial L}{\partial b_j}=\sum_{i=1}^m\frac{\partial L}{\partial z_{ij}}\cdot\frac{\partial z_{ij}}{\partial b_j}=\sum_{i=1}^m\frac{\partial L}{\partial z_{ij}}\cdot 1=\sum_{i=1}^m\frac{\partial L}{\partial z_{ij}}$$
最直观的理解方式
我们知道 $\partial L/\partial X$(如果 $X$ 是单个样本)是 $\partial L/\partial Z \cdot W^\top (n_\text{in})$
当 $X$ 具有批处理维度时,我们知道我们寻找的输出形状为 $(m,n_\text{in})$
$Z$ 的每一行 $i$ 仅依赖于 $X$ 的第 $i$ 行(批处理样本之间不相互作用)
因此我们可以堆叠每一行的梯度
一般来说,为单个样本推导雅可比矩阵(这是简洁的,二维的)
如果张量在批处理中是共享的(如 $W$),那么批处理维度被求和掉 → 收缩(矩阵乘法,批处理维是内维度)
如果张量不共享(如 $X$,激活值),批处理维度被保留 → 堆叠(矩阵乘法,批处理维在内部)
注意,PyTorch 实现中 $Z=XW^\top$ 且 $W\in\mathbb R^{n_\text{out}\times n_\text{in}}$ 如下所示
$$\frac{\partial L}{\partial X}=\frac{\partial L}{\partial Z} W\quad (m,n_\text{out})\times (n_\text{out}, n_\text{in})=(m,n_\text{in})\[1em] \frac{\partial L}{\partial W}=\left(\frac{\partial L}{\partial Z}\right)^\top X\quad (n_\text{out}, m)\times (m,n_\text{in})=(n_\text{out},n_\text{in})$$
激活函数
sigmoid $\sigma(x)\in(0,1)$
$$\sigma(x)=\frac{1}{1+e^{-x}}$$
适合将输出解释为概率
不用于神经网络的隐藏层
梯度消失,因为导数 $\sigma(x)(1-\sigma(x)) \leq 0.25$
不是零中心的,因此单个节点的下游梯度要么全为正,要么全为负(取决于上游梯度)
tanh $\in(-1,1)$
$$\tanh(x)=\frac{e^x-e^{-x}}{e^x+e^{-x}}=2\sigma(2x)-1$$
导数在 $x=0$ 处峰值为 1.0,仍可能消失
$\tanh^\prime$ 因子只会缩小,因为 $\tanh^\prime(z)=1-\tanh^2(x)\in(0,1]$
softmax → 概率分布
$$\text{softmax}(\mathbf x)_i=\frac{e^{x_i}}{\sum_j e^{x_j}}$$
带温度
$$\text{softmax}(\mathbf x/T)_i=\frac{e^{x_i/T}}{\sum_j e^{x_j/T}}$$
ReLU $\in(0, \infty)$
$$\operatorname{ReLU}(x)=\max(x,0)$$
导数为 1($x>0$),0($x<0$)
ReLU 死亡:如果预激活永久为负(即对所有输入都为负),它永远接收零梯度
训练期间网络的一部分可能死亡
Leaky ReLU $\in(-\infty,\infty)$
$$\text{LeakyReLU}(x)=\begin{cases}x&\text{if }x>0\\alpha x&\text{if }x\leq0\end{cases}$$
修正了 ReLU 死亡问题
Swish(平滑,非单调)
$$\text{Swish}(x)=x\cdot\sigma(x)$$
GLU 使用一个线性投影产生“内容”[左],另一个产生门控[右]
$$\text{GLU}(x)=xW_1\odot\sigma(xW_2)$$
SwiGLU 将 Swish 作为 GLU 内的激活函数
$$\text{SwiGLU}(x)=(x W_1)\odot \text{Swish}(xW_2)$$
没有非线性,神经网络无法做比线性变换更多的事情。
额外的层可以归结为单个线性变换 $W_1 W_2x=Wx$。
没有非线性,增加更多层不会带来更强的表达能力。
通过包含非线性的更多层,它们可以逼近任何复杂函数!

梯度
变量上的导数告诉你整个表达式对其值的敏感度。
如果 $\partial f/\partial x=3$,那么将 $x$ 改变一个小量 $h$ 将导致 $f(x)$ 改变约 $3h$。
$$\frac{df(x)}{dx}=\frac{f(x+h)-f(x)}{h}$$
梯度 $\nabla f$ 是偏导数的向量。
给定一个具有 $m$ 个输出和 $n$ 个输入的函数,雅可比矩阵是一个 $m\times n$ 的偏导数矩阵。
$$f(\mathbf{x})=[f_1(x_1,...,x_n),...,f_m(x_1,...,x_n)]\[1em] \frac{\partial f}{\partial x}=\begin{bmatrix}\frac{\partial f_1}{\partial x_1}&\cdots &\frac{\partial f_1}{\partial x_n}\ \vdots&\ddots&\vdots\ \frac{\partial f_m}{\partial x_1}&\cdots&\frac{\partial f_m}{\partial x_n}\end{bmatrix}$$
给定一个具有 $n$ 个输入和标量输出的函数,Hessian 矩阵是一个 $n\times n$ 的二阶偏导数矩阵,其中 $H_{ij}=\frac{\partial^2 f}{\partial x_i\partial x_j}$。
损失函数的 Hessian 矩阵告诉我们损失景观的曲率。
链式法则
对于一元函数的复合,我们将导数相乘。
$$x=3y, y=x^2\[1em] \frac{dz}{dx}=\frac{dz}{dy}\frac{dy}{dx}=3\cdot 2x=6x$$
对于多变量函数,我们将雅可比矩阵相乘。
$$\mathbf h=f(\mathbf z), \mathbf z=\mathbf W\mathbf x+\mathbf b\[1em] \frac{\partial\mathbf h}{\partial\mathbf x}=\frac{\partial\mathbf h}{\partial\mathbf z}\frac{\partial\mathbf z}{\partial\mathbf x}=\cdots$$
神经网络设置
$$\begin{align*}\mathbf x&\in\mathbb{R}^d\ \mathbf h&=f(\mathbf W\mathbf x+\mathbf b)\in\mathbb{R}^k & \mathbf W\in\mathbb{R}^{k\times d},\mathbf b\in\mathbb{R}^k\ \mathbf s&=\mathbf u^\intercal\mathbf h\in\mathbb{R} & \mathbf u\in\mathbb{R}^k\end{align*}$$

对于逐元素激活函数 $\mathbf h=f(\mathbf z)$,其中 $\mathbf h,\mathbf z\in\mathbb{R}^n$,${\partial\mathbf h}/{\partial\mathbf z}$ 是什么?
$$\left(\frac{\partial \mathbf h}{\partial \mathbf z}\right)_{ij}=\frac{\partial h_i}{\partial z_j}=\frac{\partial}{\partial z_j}f(z_i)=\begin{cases}f'(z_i)&\text{if }i=j\0&\text{otherwise}\end{cases}$$
雅可比矩阵是一个对角矩阵
$$\frac{\partial\mathbf h}{\partial \mathbf z}=\begin{bmatrix}f'(z_1)&&\ &\ddots&\ &&f'(z_n)\end{bmatrix}=\operatorname{diag}(f'(\mathbf z))$$
有用的雅可比矩阵
$$\frac{\partial}{\partial \mathbf x}(\mathbf W\mathbf x+\mathbf b)=\begin{bmatrix}\ddots&&\ &\frac{\partial z_i}{\partial x_j}&\ &&\ddots\end{bmatrix}=\begin{bmatrix}\ddots&&\ &W_{ij}&\ &&\ddots\end{bmatrix}=\mathbf W\[1em] \frac{\partial}{\partial \mathbf b}(\mathbf W\mathbf x+\mathbf b)=\mathbf I\[1em] \frac{\partial}{\partial \mathbf z} f(\mathbf z)=\operatorname{diag}(f'(\mathbf z))\[1em] \frac{\partial}{\partial \mathbf u}(\mathbf u^\intercal\mathbf h)=\mathbf h^\intercal$$
其他有用的导数
$$\begin{align*} \frac{d}{dx} \frac 1x &=-\frac{1}{x^2}\ \frac{d}{dx} e^x &= e^x\ \frac{d}{dx} \sigma(x) &= (1-\sigma(x))\sigma(x)\quad\text{[使用商法则]}\ \frac{d}{dx}\log x &= \frac 1x\ \frac{d}{dx}\tanh(x) &= 1-\tanh^2(x) \end{align*}$$
如何得到 sigmoid 的导数
$$\begin{align*} \frac{d}{dx}\sigma(x) &= \frac{d}{dx}(1+e^{-x})^{-1} \ &=(1+e^{-x})^{-2}\cdot e^{-x}&\text{链式法则}\ &= \frac{e^{-x}}{(1+e^{-x})^2}\ &= \frac{1}{1+e^{-x}}\cdot\frac{e^{-x}}{1+e^{-x}}\ &=\sigma(x)\cdot(1-\sigma(x)) \end{align*}$$
Swish 的导数
$$\begin{align*} \frac{\partial}{\partial x}\text{Swish}(x)&=\frac{\partial}{\partial x}x\cdot\sigma(x)\ &=\sigma(x)+x\cdot\sigma'(x)&\text{乘积法则}\ &=\sigma(x)+x\cdot\sigma(x)(1-\sigma(x))&\text{sigmoid 的导数}\ &=\sigma(x)+x\sigma(x)-x\sigma(x)^2\ &=\sigma(x)+\text{Swish}(x)(1-\sigma(x)) \end{align*}$$
softmax + CE 损失的梯度
设 $\mathbf z\in\mathbb R^{\mathcal V}$ 为 logits,下标 $i\in{1,\ldots,\mathcal V}$,$\mathbf p\in\mathbb R^{\mathcal V}$ 为 softmax 后的概率
$$p_i=\frac{e^{z_i}}{\sum_je^{z_j}}$$
CE 损失梯度 $\partial L/\partial \mathbf p$
$$L=-\log p_t$$
其中 $t$ 是正确类别
梯度为
$$\frac{\partial L}{\partial p_i}=\begin{cases}-\frac{1}{p_t}&\text{if }i=t\0&\text{otherwise}\end{cases}$$
$$\frac{\partial L}{\partial \mathbf p}=\begin{bmatrix}0&\cdots&-\frac{1}{p_t}&\cdots &0\end{bmatrix}\in\mathbb R^{\mathcal V}$$
使用链式法则将 $\partial L/\partial\mathbf z$ 表示为 $\partial \mathbf p/\partial\mathbf z$
大部分项消失,因为 $\frac{\partial L}{\partial p_i}$ 仅在 $i=t$ 时非零
$$\frac{\partial L}{\partial z_i}=\sum_{j}\frac{\partial L}{\partial p_j}\frac{\partial p_j}{\partial z_i}=-\frac{1}{p_t}\frac{\partial p_t}{\partial z_i}$$
$$\frac{\partial L}{\partial\mathbf z}=\frac{\partial L}{\partial\mathbf p}\frac{\partial\mathbf p}{\partial\mathbf z}=\begin{bmatrix}0&\cdots&-\frac{1}{p_t}&\cdots &0\end{bmatrix}\begin{bmatrix}\frac{\partial p_1}{\partial z_1}&\cdots&\frac{\partial p_1}{\partial z_{\mathcal V}}\&\ddots&\\frac{\partial p_{\mathcal V}}{\partial z_1}&\cdots&\frac{\partial p_{\mathcal V}}{\partial z_{\mathcal V}}\end{bmatrix}=-\frac{1}{p_t}\frac{\partial p_t}{\partial z_i}\in\mathbb R^{\mathcal V}$$
现在计算 softmax 梯度 $\partial \mathbf p/\partial\mathbf z$
$$\frac{\partial\mathbf p}{\partial\mathbf z}=\begin{bmatrix}\frac{\partial p_1}{\partial z_1}&\cdots&\frac{\partial p_1}{\partial z_{\mathcal V}}\&\ddots&\\frac{\partial p_{\mathcal V}}{\partial z_1}&\cdots&\frac{\partial p_{\mathcal V}}{\partial z_{\mathcal V}}\end{bmatrix}\in\mathbb R^{\mathcal V\times\mathcal V}$$
$$\frac{\partial p_j}{\partial z_i}= \begin{cases}p_j(1-p_j)&\text{if }i=j\ -p_jp_i&\text{if }i\neq j \end{cases}$$
综合起来得到 $\partial L/\partial\mathbf z$
对于正确 token $(i=t)$
$$\frac{\partial L}{\partial z_t}=-\frac{1}{p_t}\frac{\partial p_t}{\partial z_t}=-\frac{1}{p_t}\cdot p_t(1-p_t)=p_t-1$$
对于所有其他 token
$$\frac{\partial L}{\partial z_i}=-\frac{1}{p_t}\cdot(-p_tp_i)=p_i$$
非常简洁的结果:$\partial L/\partial\mathbf z = \mathbf p - \operatorname{one_hot}(t)$
反向传播
NN 方程表示为计算图
函数的哪些部分被视为“门”是一个方便的问题,通常是表达式中具有简单局部梯度的部分
反向传播可以理解为门之间相互通信(通过梯度信号),告诉对方它们希望输出增加还是减少(以及增加/减少的强度),以降低损失
通过重复应用链式法则实现,这允许我们将每个梯度分解为上游梯度(已计算)和局部梯度
图中的每个节点接收一个上游梯度,并向下传递一个下游梯度
每个节点都有一个局部梯度(其输出关于其输入的梯度)
下游梯度 = 上游梯度 × 局部梯度

梯度在外部分支处求和
如果 $y$ 用于计算 $a$ 和 $b$,则 $\frac{\partial f}{\partial y}=\frac{\partial f}{\partial a}\frac{\partial a}{\partial y}+\frac{\partial f}{\partial b}\frac{\partial b}{\partial y}$
节点直觉
$+$ 将上游梯度分配给每个被加数
$\max$ 将上游梯度“路由”到多个输入参数之一
$\times$ 在下游梯度中交换前向系数
反向传播
初始化输出梯度为 1
以逆拓扑顺序访问节点:使用关于后继节点的梯度计算关于每个节点的梯度
正确完成后,前向传播和反向传播的复杂度量级相同
自动微分
梯度计算可以从前向传播的符号表达式自动推断
每个节点类型需要知道如何计算其输出,以及如何在给定关于输出的梯度时计算关于输入的梯度
局部梯度由程序员编写
手动梯度检查
对于每个参数 $x$,重新计算 $x-h$ 和 $x+h$ 处的 $f$,并检查
$$f'(x)\approx \frac{f(x+h)-f(x-h)}{2h}$$
在反向传播中,需要中间激活值
因此,神经网络通常在前向传播期间存储所有中间激活值
激活/梯度检查点通过仅存储一部分激活值(“检查点”)来以计算换内存——如果需要的激活值未保存,则从最近的检查点通过部分前向传播即时重新计算
对于具有 $N$ 层并划分为 $K$ 个检查点的模型
内存从 $O(N)$ 变为 $O(K+N/K)$
反向计算从 $O(N)$ 变为 $O(N + N*(K-1)/K)$
最优选择是 $K=\sqrt{N}$:$O(\sqrt{N})$ 内存,约 $O(2N)$ 反向计算
动手练习
为 $\partial f/\partial x$ 计算显式表达式会极其复杂,但完全没有必要!
在前向传播中构造多个中间变量,每个都是简单的表达式,我们知道其局部梯度
$$f(x,y)=\frac{x+\sigma (y)}{\sigma(x)+(x+y)^2}$$
反向传播需要从标量开始,因为它计算每个参数 $\theta$ 的 $\partial L/\partial\theta$,每个参数一个数——这仅在 $L$ 是标量时才有意义
当在标量上调用 .backward() 时,PyTorch 隐式地以 $\partial L/\partial L=1$ 播种反向传播
当我们有每个 token 的损失 $\ell_1,\ldots,\ell_n$ 并定义 $L=\frac 1N\sum_i\ell_i$(均值缩减)时,由导数的线性性质
$$\frac{\partial L}{\partial\theta}=\frac 1N\sum_{i=1}^N\frac{\partial\ell_i}{\partial\theta}$$
因此均值损失的梯度正好是每个 $\ell_i$ 单独反向传播的梯度的均值
上游梯度总是关于激活值的,关于参数的梯度用于更新并在此结束,因为参数是计算图的叶子节点
loss
│ dL/dy2 (激活梯度)
▼
Layer 2 ──→ dL/dW2 (参数梯度,存储)
│ dL/dy1 (激活梯度)
▼
Layer 1 ──→ dL/dW1 (参数梯度,存储)
│ dL/dx (激活梯度——通常丢弃)
▼
input
优化器
普通 SGD 更新
$$\theta\leftarrow \theta-\eta g_t$$
优化器决定参数更新的方向和大小
对于每个参数张量,Adam 保留:
- 参数本身 ($\theta$)
- 梯度 ($g$)
- 一阶矩(动量)
- 二阶矩(方差)
Adam 优化器
$$\theta\leftarrow\theta-\eta\frac{\hat m}{\sqrt{\hat v}+\epsilon}$$
令 $g$ 为当前步的梯度
一阶矩 $m$ 是梯度的运行均值
$$m\leftarrow\beta_1 m+(1-\beta_1)g$$
结果:如果梯度持续指向同一方向,则优化器更自信地朝该方向移动
随时间累积速度,帮助穿越噪声梯度
二阶矩 $v$ 是平方梯度的运行均值
$$v\leftarrow\beta_2 v+(1-\beta_2)g^2$$
结果:不同权重实际上获得不同的学习率——持续大的梯度 → 较小的步长
有效地将梯度归一化到同一尺度
偏差修正:修正训练早期运行均值的初始化偏差
(time 是时间步 $t$,从 1 开始;否则使用 $t+1$)
$$\hat m_t=\frac{m_t}{1-\beta_1^t},\quad\hat v_t=\frac{v_t}{1-\beta_2^t}$$
超参数 $\beta_1$ 和 $\beta_2$ 控制矩估计的更新
$m$ 和 $v$ 都初始化为 $0$
因此每个参数的内存 ≈ 4 × 参数大小
AdamW 通过向 0 添加权重衰减来修改 Adam
$$\theta\leftarrow\theta-\eta\frac{\hat m}{\sqrt{\hat v}+\epsilon}-\eta\lambda \theta$$
用模型参数初始化优化器,告诉优化器它将优化哪些值,以及 lr 参数,它决定更新的大小
如何判断某事应该是 LR 调度还是优化器?
- 仅依赖于时间步 $t$ → 可能是 LR 调度
- 需要每个参数的历史 → 优化器
代码如下:params 用于创建参数组,每个组有自己的超参数(例如,不同层有不同的学习率)
torch.optim.AdamW(model.parameters()) 创建一个单一参数组
通常我们不想在偏置和 LayerNorm 参数上应用权重衰减
torch.optim.AdamW([{'params': decay_params,'weight_decay':0.01},{'params': no_decay_params,'weight_decay':0.0},])
defaults 字典为任何未在参数组中显式指定的超参数提供回退值
实践中,最好在 Adam 更新之前应用权重衰减,因为权重衰减依赖于参数
class AdamW(torch.optim.Optimizer):
def __init__(self, params, lr, betas, eps, weight_decay):
if lr < 0:
raise ValueError(f"无效学习率: {lr}")
if not 0 < betas[0] < 1 or not 0 < betas[1] < 1:
raise ValueError(f"无效 beta 值: {betas}")
defaults = {"lr": lr, "betas": betas, "eps": eps, "weight_decay": weight_decay}
super().__init__(params, defaults)
def step(self):
for group in self.param_groups:
# 对于每一组参数
lr = group["lr"]
beta1, beta2 = group["betas"]
eps = group["eps"]
weight_decay = group["weight_decay"]
for p in group["params"]:
# 对于组中的每个参数
if p.grad is None:
continue
state = self.state[p]
# 用 0 初始化状态
t = state.get("t", 0)
m, v = state.get("m", torch.zeros_like(p.data)), state.get("v", torch.zeros_like(p.data))
# 权重衰减
p.data -= lr * weight_decay * p.data
# Adam 更新
grad = p.grad.data
m = beta1 * m + (1 - beta1) * grad
v = beta2 * v + (1 - beta2) * grad ** 2
m_hat = m / (1 - beta1 ** (t + 1))
v_hat = v / (1 - beta2 ** (t + 1))
p.data -= lr * m_hat / (v_hat.sqrt() + eps)
# 更新优化器状态
state["t"] = t + 1
state["m"] = m
state["v"] = v
梯度裁剪限制梯度范数的大小
- 计算所有梯度的全局范数
- 如果超过最大值,将所有参数按相同比例缩小,使其低于最大值
- 防止任何单个步骤灾难性地过大
学习率
预热减少了早期训练样本的首因效应
数学相关
信息论
交叉熵
$$\operatorname{CE}(p,q)=-\mathbb E_p[\log q]=-\sum_{x\in\mathcal X}p(x)\log q(x)$$
KL 散度
$$\operatorname{KL}(p\mid\mid q)=\sum_{x\in\mathcal X}p(x)(\log p(x)-\log q(x))$$
熵
$$H(p)=-\sum_{x\in\mathcal X}p(x)\log p(x)$$
另一种常见形式,给定 logits $x_i$
$$H(p)=\log\sum_i e^{x_i}-\underbrace{\frac{\sum_i e^{x_i}x_i}{\sum_i e^{x_i}}}_{\mathbb E[x]}$$
$p$ 和 $q$ 之间的交叉熵就是 $p$ 和 $q$ 之间的 KL 散度加上 $p$ 的不可约熵
$$\operatorname{CE}(p,q)=\operatorname{KL}(p\mid\mid q)+H(p)$$
证明
$$\begin{align*} \operatorname{KL}(p\mid\mid q)&=\sum_{x\in\mathcal X}p(x)(\log p(x)-\log q(x))\ &=\sum_{x\in\mathcal X}p(x)\log p(x)-\sum_{x\in\mathcal X}p(x)\log q(x)\ &=-H(p)+\operatorname{CE}(p,q) \end{align*}$$
交叉熵损失
当目标分布是 one-hot 时,交叉熵损失是下一个 token 的负对数似然
$$\mathcal L(x)=-\sum_{t=1}^T\log p(x_t\mid x_{<t})$$
也等价于 KL 散度
实现损失
如果使用 F.cross_entropy(),logits 和 labels 会在内部移位
loss = F.cross_entropy(logits.view(-1, vocab_size), targets.view(-1), ignore_index=pad_idx)
shift_logits = logits[:, :-1, :]
shift_labels = input_ids[:, 1:]
logprobs = F.log_softmax(shift_logits, dim=-1)
token_logprobs = logprobs.gather(index=shift_labels.unsqueeze(-1), dim=-1).squeeze(-1)
## 构建损失掩码
masked_logprobs = -token_logprobs * mask.float()
return masked_logprobs.sum() / mask.sum()
数值稳定性及其他技巧
一般来说,要注意
- $\exp(x)$ 对于大的 $x$ → 溢出为 $\infty$
- $\log(x)$ 对于接近 0 的 $x$ → 下溢为 $-\infty$ ($\log(0)=-\infty$)
- $\log(x)$ 对于接近 1 的 $x$ → 精度问题 ($\log(1) = 0$)
计算 softmax(x) 不稳定,因为对于大的 $x_i$,$\exp(x_i)$ 会溢出
利用 softmax 对减去常数不变性
$$\begin{align*} \text{softmax}(x)_i &= \frac{e^{x_i}}{\sum_j e^{x_j}}\ &= \frac{e^{x_i-c}\cdot e^c}{\sum_j e^{x_j-c}\cdot e^c}\ &= \frac{e^{x_i-c}\cdot e^c}{e^c\cdot \sum_j e^{x_j-c}}\ &= \frac{e^{x_i-c}}{\sum_j e^{x_j-c}}\ &=\text{softmax}(x-c)_i \end{align*}$$
从所有 $x_i$ 中减去 $x_\text{max}$ 确保 $x_i$ 不大(最大指数是 $\exp(0)=1$),因此数值稳定的实现这样做
$$\text{softmax}(x)i=\frac{e^{x_i-x\text{max}}}{\sum_j e^{x_j-x_\text{max}}}$$
计算 log(softmax(x))
直接做 log-softmax 不好,因为接近 0 的输入(低概率类别)的 log 不稳定
改用 x - logsumexp(x),避免生成微小的概率(且 logsumexp(x) 是稳定的)
它们是等价的
$$\begin{align*} \log(\text{softmax}(x))_i &= \log \frac{e^{x_i}}{\sum_j e^{x_j}}\ &= \log e^{x_i}-\log\sum_j e^{x_j}\ &= x_i-\log\sum_j e^{x_j} \end{align*}$$
计算 log(sum(exp(x)))
为什么它不稳定?
- 对于任何大的 $x_i$,$\exp(x_i)$ 会溢出为无穷
- 例如,在 float32 中,$x_i\approx 83$ 会溢出
- 如果所有 $x_i$ 都非常负,那么 $\log(0)$ 会下溢为 $-\infty$
- 如果求和接近 1,存在精度问题,因为 $\log(1)$ 不稳定
直觉:我们希望使 $x_i$ 的值变小,这可以通过减去一个常数实现!然后只需要在最后加回这个常数
logsumexp(x) 的实现如下
$$\begin{align*} \log\sum_i e^{x_i}&=\log\sum_i\left(e^{x_i-x_\text{max}} \cdot e^{x_\text{max}}\right)\ &= \log \left(e^{x_\text{max}}\sum_i e^{x_i-x_\text{max}}\right)\ &=\log e^{x_\text{max}}+\log\sum_i e^{x_i-x_\text{max}}\ &= x_\text{max}+\log\sum_i e^{x_i-x_\text{max}} \end{align*}$$
最大项是 $e^0=1$(对于 $x_i=x_\text{max}$),所以不会溢出
求和至少为 1,所以我们永远不会计算 $\log(0)$
朴素地,计算 softmax 需要两次不同的遍历来计算 $x_\text{max}$(用于稳定性)然后分母 $\sum_j e^{x_j-x_\text{max}}$
稳定的 softmax
$$\text{softmax}(x)i=\frac{e^{x_i-x\text{max}}}{\sum_j e^{x_j-x_\text{max}}}$$
在线 softmax 技巧:将 $x_\text{max}$ 和分母 $\sum_j e^{x_j-x_\text{max}}$ 的计算融合为一次遍历
思路:我们可以用当前最大值计算分母,并随着新的当前最大值不断重新缩放
维护
- 运行最大值 $m_k=\max(x_1,...,x_k)$
- 运行(移位后的)分母 $d_k=\sum_j e^{x_j-m_k}$
更新规则:当我们遇到 $x_{k+1}$ 时
$$m_{k+1}\leftarrow\max(m_k,x_{k+1})\ d_{k+1}\leftarrow d_k\cdot e^{m_k-m_{k+1}}+e^{x_{k+1}-m_{k+1}}$$
理解为什么 $d_{k+1}$ 的更新是正确的
$$\begin{align*} d_{k+1}&=\sum_{j=1}^{k+1}e^{x_j-m_{k+1}}\ &= e^{x_{k+1}-m_{k+1}}+\sum_{j=1}^k e^{x_j-m_{k+1}}&\text{分离最后一项}\ &= e^{x_{k+1}-m_{k+1}}+\sum_{j=1}^k e^{x_j-m_k+m_k-m_{k+1}}&\text{代数变换}\ &= e^{x_{k+1}-m_{k+1}}+\sum_{j=1}^k e^{x_j-m_k}e^{m_k-m_{k+1}}\ &= e^{x_{k+1}-m_{k+1}}+e^{m_k-m_{k+1}}\underbrace{\sum_{j=1}^k e^{x_j-m_k}}{d_k}&\text{提取常数}\ &= \underbrace{d_k\cdot e^{m_k-m{k+1}}}\text{重新缩放前项}+\underbrace{e^{x{k+1}-m_{k+1}}}_\text{新项}&\text{用 }d_k\text{ 表示,重排} \end{align*}$$
注意当最大值不变时,$m_k=m_{k+1}$ 且重新缩放因子 $e^{m_k-m_{k+1}}=1$
如果目标是返回 softmax,那么我们需要再遍历一次 logits 以返回 $e^{x_i-m_S}/d_S$
我们还可以用这个技巧来计算加权和,给定 logits 流 $x_i$ 和值 $v_i$
$$o=\sum_i p_i v_i=\sum_i\frac{e^{x_i}}{\sum_j e^{x_j}}v_i=\frac{\sum_i e^{x_i}v_i}{\sum_i e^{x_i}}$$
对于 FlashAttention,$o$ 对应于每个查询 $q$ 的注意力输出,其中 $x_i=q\cdot k_i$,$v_i$ 是值向量
数值稳定版本将分子和分母都乘以 $e^{-x_\text{max}}$(以避免 $e^{x_i}$ 的溢出问题)
$$o=\frac{\sum_i e^{x_i-x_\text{max}}v_i}{\sum_i e^{x_i-x_\text{max}}}$$
除了 $m_k$ 和 $d_k$,我们还维护运行分子 $o_k\in\mathbb R^H$($H$ = 每个 $v_i$ 的维度,在 FlashAttention 中为 head 维度)
$$o_k=\sum_{i=1}^k e^{x_i-m_k}v_i$$
更新使用与 $d_k$ 相同的思路(推导看起来与 $d_k$ 的相同)
$$o_k=o_k\cdot e^{m_k-m_{k+1}}+ \underbrace{e^{x_{k+1}-m_{k+1}}v_{k+1}}_\text{新项}$$
处理完所有 $N$ 个 logits 后,我们有 $(m_N, d_N, o_N)$,真实的注意力输出就是 $o_N/d_N$
当我们需要使涉及 $e^{x}$ 的表达式数值稳定时,一个标准工具是乘以 $e^{-m}$(其中 $m$ 是一个大数,通常选择 $m=x_\text{max}$)并观察幸存的结果。
基本统计
$p$ 值:在零假设为真的条件下,看到数据(至少如此极端)的概率
关键是,它是关于数据在零假设下的概率,而不是零假设在数据下的概率
这些组有差异吗?
- Kolmogorov-Smirnov 检验:给定两组连续变量的观测值,它们是否来自相同的底层分布?测量两个样本的经验 CDF 之间的最大垂直距离
- 卡方检验:给定两组分类变量的观测值,它们是否来自相同的底层分布?
- T 检验:两组连续观测值的均值是否不同?假设:数据正态分布;单样本版本:一组的均值是否等于某个值;双样本版本:两组是否具有相同均值;配对版本:同一项目上的两次测量是否系统性不同(伪装成单样本 t 检验;给定每个样本的差异(每个差异是 +1、0 或 -1),可以检验差异是否显著不为 0)
- ANOVA(F 检验):(t 检验对多于两组的推广)这 $k$ 组中是否有任何组的均值不同?
- McNemar 检验:比较两个分类器在同一数据集上的表现(这可能是标准设置中在相同测试集上评估两个模型的最佳检验,每个样本可正确或错误回答)
这些变量有关联吗?
- Pearson 相关性检验:检验两个连续变量之间是否存在线性关系;返回相关系数 $r$ 和 $r$ 是否显著不为 0 的 p 值;完全忽略非线性关系(完美的抛物线关系 → $r\approx 0$);想象在散点图上拟合一条直线
- Spearman 相关性检验:相同思路,但衡量单调关联;将两个变量转换为秩(最小值得到秩 1,等等),然后在秩上计算 Pearson 相关性;重要的是顺序,而不是具体值
- Pearson 与 Spearman:Pearson 对异常值更敏感,而 Spearman 对其鲁棒;Pearson 低估非线性单调关系;Pearson 有更清晰的解释
- 互信息:捕捉两个变量之间的任何依赖性,包括非线性的;不算是统计检验
通过采样的梯度流
Gumbel-Max 技巧
给定 logits $z_1,\ldots,z_k$,可以通过以下方式从相应的分类分布中采样
- 从 $\operatorname{Gumbel}(0,1)$ 分布中抽取独立的噪声值 $g_1,\ldots,g_k$
- 取 $(z_1+g_1,\ldots,z_k+g_k)$ 的 argmax
Gumbel-Softmax 用 softmax 替换 argmax:$\operatorname{softmax}((z_1+g_1,\ldots,z_k+g_k)/\tau)$
- softmax 处处可微!
- 高温度 → 平滑梯度,低温度 → 离散样本
- 不是为了使其可微(普通 softmax 已经是),而是使其随机化
- 普通 softmax:$y = \operatorname{softmax}(\alpha)$(确定性;总是相同的软混合)
- 真实分类采样:$y=\operatorname{one_hot}(\operatorname{sample}(\operatorname{softmax}(\alpha)))$
- gumbel-softmax:$y=\operatorname{softmax}((\alpha+G)/\tau)$(随机(Gumbel 噪声),近似离散(低温度),可微)
- 你获得了探索 AND 梯度!
直通估计器:假装 $f$ 是恒等函数
- 前向:$y=f(x)$(应用不可微函数)
- 反向:$\frac{\partial\mathcal L}{\partial x}=\frac{\partial\mathcal L}{\partial y}$
- 直接将上游梯度向下传递作为下游梯度
- 如果 $f(x)$ 是递减的(信号方向相反),则这会是错误的
- 但 STE 通常应用于单调递增的操作(取整、量化、递增阶跃函数);方向正确,即使大小可能错误
理论计算机科学
正则语言是可以用有限状态机(也称为有限自动机)识别的语言
- 确定型有限自动机(DFA)
上下文无关语言是可以用下推自动机识别的语言——基本上是一个有限状态机加一个栈
- 栈提供了无界内存,但只能通过栈顶访问
- 编程语言语法由上下文无关语言构建(匹配括号、嵌套函数调用、平衡的 HTML 标签等)
任何有限语言显然都是正则的——你可以用足够的状态列举所有有效字符串
任何 DFA 都可以由 ReLU RNN 编码
给定一个 DFA,具有
- 状态 $Q={q_1,\ldots,q_k}$
- 字母表 $\Sigma={\sigma_1,\ldots,\sigma_m}$
- 转移函数 $\delta:Q\times\Sigma\to Q$
- 起始状态 $q_1$
- 接受状态 $F\subseteq Q$
构建一个隐藏维度为 $k$ 的 RNN,它是状态的 one-hot 编码
需要构造 $W_h\in\mathbb R^{k\times k}$,$W_x\in\mathbb R^{k\times m}$,$b\in\mathbb R^k$
对于每个转移 $\delta(q_i,\sigma_k)=q_j$,设置 $(W_h){ji}=1$,$(W_x){jk}=1$,$b_j=-1$
现代 Transformer LM
架构
| 符号 | 维度 |
|---|---|
| B | 批处理中的序列数 |
| L | 层数 |
| T | 序列长度(要生成的 token 数) |
| S | 序列长度(提供的上下文) |
| V | 词表大小 |
| D | 隐藏维度 |
| H | head 维度 |
| F | MLP 隐藏维度,通常 F = 4D |
| N | 查询 head 数,N * H = D |
| K | 键/值 head 数,在 GQA 中 K < N |
| G | GQA 中的组大小 = N // K |

token 嵌入
嵌入矩阵 $\mathbf W_e\in\mathbb R^{V\times D}$,初始隐藏状态 $\mathbf X^{(0)}\in\mathbb R^{B\times S\times D}$
$$\mathbf X^{(0)}=\mathbf {W}_e[\text{tokens}]$$
层循环(对于 $\ell\in[0,\ldots,L-1]$)
RMSNorm 将 $\mathbf X^{(\ell)}$ 的每个元素除以 $\mathbf X^{(\ell)}$ 的 RMS(使得隐藏状态具有单位 RMS),然后乘以学习的重新缩放参数 $\gamma$
$$\bar{\mathbf X}^{(\ell)}=\frac{\mathbf X^{(\ell)}}{\text{RMS}(\mathbf X^{(\ell)})+\epsilon}\odot\gamma_{\text{attn}}^{(\ell)}\[1em] \operatorname{RMS}(\mathbf X)=\sqrt{\frac 1D\sum_{i=1}^D x_i^2}$$
每个 head 使用 $W^{(\ell)}_Q\in\mathbb R^{D\times D}$,$W^K\in\mathbb R^{D\times KH}$,$W^{(\ell)}_V\in\mathbb R^{D\times KH}$(其中 $H=D/N$)将 $\mathbf {X}^{(\ell)}$ 投影到 head 的低维子空间中
$$\mathbf Q=\bar {\mathbf X}\mathbf W_Q\in\mathbb R^{B\times T\times D},\quad\mathbf K=\bar {\mathbf X}\mathbf W_K\in\mathbb R^{B\times S\times D},\quad\mathbf V=\bar {\mathbf X}\mathbf W_V\in\mathbb R^{B\times S\times D}$$
[可选] QK 归一化:对查询和键向量应用 RMSNorm,以控制进入点积的向量的大小
重塑以暴露 head 维度 $D\to N\times H$,以及 $K\cdot H\to K\times H$,然后转置序列长度($S$ 或 $T$)和 head 维度($N$ 或 $K$)
$$\mathbf Q\in\mathbb R^{B\times T\times D}\to \mathbb R^{B\times N\times T\times H}\ \mathbf K\in\mathbb R^{B\times S\times (K\cdot H)}\to \mathbb R^{B\times K\times S\times H}\ \mathbf V\in\mathbb R^{B\times S\times (K\cdot H)}\to \mathbb R^{B\times K\times S\times H}$$
为 GQA 扩展 $K$、$V$
$$\mathbf K\in\mathbb R^{B\times K\times S\times H}\to \mathbb R^{B\times N\times S\times H}\ \mathbf V\in\mathbb R^{B\times K\times S\times H}\to\mathbb R^{B\times N\times S\times H}$$
通过在每个位置 $m$ 处将查询向量 $\mathbf q_m\in\mathbf R^H$(或键向量 $\mathbf k_m$)旋转 $\mathbf R_m$ 来应用 RoPE
对于维度对 $i$(对应 $\mathbf q_m$ 中的索引 $(2i, 2i+1)$),我们旋转角度 $m\theta_i$,其中 $\theta_i=\Theta^{-\frac{2i}{H}}$
超参数 $\Theta$ 控制基本旋转频率,$H$ 是 head 维度
$$\mathbf R_m=\begin{bmatrix} \ddots&&\ &\mathbf R_m^{(i)}&\ &&\ddots\ \end{bmatrix}\in\mathbb R^{H\times H}\quad\text{其中}\quad \mathbf R_m^{(i)}=\begin{bmatrix} \cos(m\theta_i)&-\sin(m\theta_i)\ \sin(m\theta_i)&\cos(m\theta_i)\ \end{bmatrix}$$
$$\mathbf q_m\leftarrow \mathbf R_m \mathbf q_m, \quad\mathbf k_m\leftarrow \mathbf R_m \mathbf k_m,$$
计算注意力分数
我们除以 head 维度 $H$,否则点积会随 $\sqrt H$ 缩放
softmax 的大输入 → 更尖的分布 → 对更新抵抗
$$\mathbf A=\frac{\mathbf Q\mathbf K^\top}{\sqrt H}\in\mathbb R^{B\times N\times T\times S}$$
应用因果掩码
$$\mathbf A_{ij}\leftarrow \begin{cases} \mathbf A_{ij}&\text{if }j\leq i\ -\infty&\text{if }j>i \end{cases}$$
应用 softmax
$$\mathbf A=\operatorname{softmax}(\mathbf A)=\frac{\exp(\mathbf A_{ij})}{\sum_{k=1}^S \exp(\mathbf A_{ik})}$$
从值的加权和得到注意力输出
$$\mathbf O=\mathbf A\mathbf V\in\mathbb R^{B\times N\times T\times H}$$
重塑
$$\mathbf O\in\mathbb R^{B\times N\times T\times H}\to\mathbb R^{B\times T\times D}$$
应用输出投影 $W^{(\ell)}_O\in\mathbb R^{D\times D}$ 以混合不同 head 的输出
$$\mathbf O_\text{proj}=\mathbf O\mathbf W_O$$
残差连接
$$\mathbf X^{(\ell)} \leftarrow\mathbf X^{(\ell)}+\mathbf O_\text{proj}$$
前馈网络
RMSNorm
$$\bar{\mathbf X}^{(\ell)}=\frac{\mathbf X^{(\ell)}}{\text{RMS}(\mathbf X^{(\ell)}) + \epsilon}\odot\gamma_{\text{ffn}}^{(\ell)}$$
门控和上投影 [扩展] 使用 $\mathbf W^{(\ell)}\text{up}\in\mathbb R^{D\times F}$,$\mathbf W^{(\ell)}\text{gate}\in\mathbb R^{D\times F}$
$$\mathbf U=\bar {\mathbf X}\mathbf W_\text{up}\ \mathbf G=\bar {\mathbf X}\mathbf W_\text{gate}$$
SwiGLU 激活
$$\operatorname{Swish}(\mathbf G)=\mathbf G\odot\sigma(\mathbf G)=\mathbf G\odot\frac{1}{1+e^{-\mathbf G}}\[1em] \mathbf H=\operatorname{Swish}(\mathbf G)\odot\mathbf U\in\mathbb R^{B\times T\times F}$$
下投影使用 $\mathbf W_\text{down}^{(\ell)}\in\mathbb R^{F\times D}$
$$\mathbf F=\mathbf H\mathbf W_\text{down}$$
残差连接
$$\mathbf X^{(\ell+1)} =\mathbf X^{(\ell)}+\mathbf F$$
最终层归一化
最终归一化
$$\mathbf X_\text{final}=\frac{\mathbf X^{(L)}}{\text{RMS}(\mathbf X^{(L)}) + \epsilon}\odot\gamma_{\text{final}}$$
去嵌入
使用 $\mathbf W_u\in\mathbb R^{D\times V}$ 投影到词表维度
$$\mathbf Z=\mathbf X_\text{final}\mathbf W_u\in\mathbb R^{B\times T\times V}$$
实现说明
scores.masked_fill(~mask, -torch.inf) 用于填充 softmax 前的注意力分数
假设约定中 mask 为 True 表示可以参与的位置
tensor.masked_fill(mask, value) 在 mask 为 True 的位置用 value 填充 tensor
RoPE
我们希望缓存每个(位置,索引)对 $(m,i)$ 的 $\cos(m\theta_i)$ 和 $\sin(m\theta_i)$
我们可以在初始化时这样做
positions = torch.arange(max_seq_len, device=device) # 形状 (max_seq_len)
thetas = self.theta ** (-torch.arange(0, d_k, 2, device=device) / d_k) # 形状 (d_k // 2)
angles = positions.unsqueeze(-1) * thetas.unsqueeze(0)
在实践中,不是做一堆 2x2 矩阵乘法,而是用点积来表达旋转
通过将最终的 head 维度 $H$ 重塑为 $(H/2,2)$ 来提取 $\mathbf Q$、$\mathbf K$ 的偶数和奇数索引
x_pairs = x.reshape(*x.shape[:-1], -1, 2)
x_even = x_pairs[..., 0]
x_odd = x_pairs[..., 1]
计算旋转矩阵中的所有偶数和奇数位置
x_out_even = x_even * cos - x_odd * sin
x_out_odd = x_even * sin + x_odd * cos
然后通过并排堆叠并展平来交错
torch.stack() 增加新维度
torch.stack([x_out_even, x_out_odd], dim=-1).flatten(start_dim=-2)
注意力看起来像这样
需要使用 .reshape() 将 d_model ($D$) 扩展为 num_heads x head_dim ($N\times H$)
需要使用 qkv.unbind() 来拆分查询、键、值向量 [可选,取决于实现]
需要使用 .transpose() 交换 num_heads 和 seq_len 维度以进行注意力计算
获得 output 后,需要再次 .transpose() 和 .reshape() 以恢复原始形状
batch_size, seq_len, _ = x.shape
x_norm = self.norm(x)
qkv = self.qkv_proj(x_norm) # (batch, seq_len, 3 * d_model)
qkv = qkv.reshape(batch, seq_len, 3, self.num_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2) # (batch, seq_len, num_heads, head_dim)
q = q.transpose(1, 2) # (batch, num_heads, seq_len, head_dim)
k = k.transpose(1, 2) # (batch, num_heads, seq_len, head_dim)
v = v.transpose(1, 2) # (batch, num_heads, seq_len, head_dim)
causal_mask = torch.tril(torch.ones(seq_len, seq_len)).bool()
output = scaled_dot_product_attention(q, k, v, mask) # (batch, num_heads, seq_len, head_dim)
output = output.transpose(1, 2) # (batch, seq_len, num_heads, head_dim)
output = output.reshape(batch, seq_len, d_model) # (batch, seq_len, d_model)
output = self.out_proj(output) # (batch, seq_len, d_model)
return x + output
注意力部分看起来像这样
def scaled_dot_product_attention(q, k, v, mask):
"""
k, q: (batch_size, ..., seq_len, d_k)
v: (batch_size, ..., seq_len, d_v)
返回 o (batch_size, ..., seq_len, d_v)
"""
d_k = q.shape[-1]
scores = (q @ k.transpose(-2, -1)) / math.sqrt(d_k)
scores = scores.masked_fill(~mask, -torch.inf)
return softmax(scores, dim=-1) @ v
核算
模型参数
嵌入:$(V,D)$
注意力是 $2D^2+2DKH\approx 4D^2$(对于标准多头注意力 $N=K$)
- $Q$ 是 $(D, D)$
- $K$ 是 $(D,KH)$
- $V$ 是 $(D,KH)$
- $O$ 是 $(D,D)$
FFN 是 $3DF$
- 上投影是 $(D,F)$
- 门控投影是 $(D,F)$
- 下投影是 $(F,D)$
层归一化是每层 $2D$(加上最终归一化)
- 注意力前和 FFN 前的层归一化各有 $D$ 个参数($D$ 中每个维度的 $\gamma$)
去嵌入:$(V,D)$
总计:$2VD+L(4D^2+2D+3DF)\approx 2VD+12LD^2$(对于 $F=8D/3$)
总模型参数为 $2VD+12LD^2$
模型激活值
注意力激活值:$6BSD+BNS^2$
- 层归一化输入是 $(B,S,D)$
- 层归一化输出是 $(B,S,D)$
- Q、K、V 输出分别是 $(B,S,D)$、$(B,S,KH)$、$(B,S,KH)$
- 注意力分数是 $(B,N,S,S)$
- 注意力输出是 $(B,S,D)$
FFN 激活值:$2BSD+2BSF\approx 8BSD$(对于 $F=8/3D$)
- 层归一化输入是 $(B,S,D)$
- 门控/上投影的输出各是 $(B,S,F)$
- 下投影的输出是 $(B,S,D)$
每层激活值为 $14BSD+BNS^2$
前向传播中的 FLOPs
假设 prefill 阶段(所以 $S=T$)
注意力每层是 $8BSD^2+4BS^2D$
- $Q$ 投影是 $(B,S,D)\times(D,D)$ → $2BSD^2$ FLOPs
- $K$ 投影是 $(B,S,D)\times (D,KH)$ → $2BSDKH\approx 2BSD^2$(对于 $K=N$)
- $V$ 投影是 $(B,S,D)\times (D,KH)$ → $2BSDKH\approx 2BSD^2$(对于 $K=N$)
- $QK^\top$ 是 $(B,N,S,H)\times(B,N,H,S)$ → $2BNS^2H=2BS^2D$(因为 $D=NH$)
- $AV$ 是 $(B,N,S,S)\times (B,N,S,H)$ → $2BS^2D$
- $O$ 投影是 $(B,S,D)\times(D,D)$ → $2BSD^2$ FLOPs
FFN 每层是 $6BSDF\approx16BSD^2$(对于 $F=8D/3$)
- 上投影是 $(B,S,D)\times (D,F)$ → $2BSDF$
- 门控投影是 $(B,S,D)\times (D,F)$ → $2BSDF$
- 下投影是 $(B,S,F)\times (F,D)$ → $2BSDF$
每层总计:$8BSD^2+4BS^2D+16BSD^2=2BSD(12D+2S)$
去嵌入层是 $2BSDV$
- 去嵌入是 $(B,S,D)\times (D,V)$ → $2BSDV$
完整前向传播是 $2LBSD(12D+2S)+2BSDV\approx 2BSD(12LD+2LS+V)$
反向传播中的 FLOPs
通常假设是前向传播 FLOPs 的 $2$ 倍
- 计算关于参数和输入的梯度,每个都是矩阵乘法
- 关于输入 $\partial L/\partial X$ 的梯度是上一层的传入梯度
推理内存使用
推理时的总内存使用:模型权重 + KV 缓存 + 峰值激活值
模型参数数量:$2VD+12LD^2$
num_params = sum(p.numel() for p in model.parameters())
KV 缓存大小:$B\cdot S\cdot (KH)\cdot L\cdot 2$
- $B$ = 批处理大小
- $S$ = 序列长度
- $K$ = KV head 数
- $H$ = head 维度
- $L$ = 层数
- 2 代表 K 和 V
激活值:prefill 中 $O(BNS^2+BSF)$,decode 中 $O(BNS+BF)$
torch.inference_mode() 立即释放内存,所以只关心单层的峰值内存
- 输入在 prefill 中是 $B\times S\times D$,在 decode 中是 $B\times T\times D$
- 使用 FlashAttention 时,$S\times S$ 矩阵从不实例化 → 注意力变为 $O(S)$
峰值激活值
- $B\cdot S\cdot D$ = 层输入
- $B\cdot S\cdot 3\cdot D$ = K、Q、V 向量
- $B\cdot N \cdot S^2$ = 注意力矩阵(不使用 FlashAttention)
- 每个批处理样本和每个查询 head 的 $S\times S$ 矩阵
- $B\cdot S\cdot F$ = FFN 中间值
激活值在没有 FA 时随序列长度 $S$ 二次缩放,有 FA 时线性缩放
- 在小批处理大小 + 短序列长度时,权重占主导
- 在 prefill 阶段:对于长序列长度(大 $S$),激活值中的 $S^2$ 注意力项占主导
- 在大批处理大小(大 $B$)时,KV 缓存和激活值都增长
训练内存使用
训练时的总内存使用:模型权重 + 优化器状态 + 梯度 + 激活值
- $P$ 模型参数
- FP32 主权重 [全精度或混合精度] → $4P$
- BF16 临时副本用于前向传播 [混合精度] → $2P$
- $2P$ 优化器状态(一阶和二阶矩)
- Adam 状态 FP32 [混合精度] → $8P$
- $P$ 梯度
- FP32 → $4P$
- 即使在混合精度中,梯度以 BF16 计算但以 FP32 累积
- 激活值
- 通常占主导部分:取决于 $B$、$S$、$D$、$L$
- 每层 $14BSD+BNS^2$ 激活值(不使用 flash attention)
- 使用 flash attention 时,第二项变为 $BNS$,激活内存随 $BS$(总 token 数)缩放
- 反向传播时需要计算梯度,可通过梯度检查点减少
注意力
标准注意力是 $O(n^2)$
KV 缓存变得独立于序列长度
滑动窗口:每个 token 只关注最后 $W$ 个 token,因此计算是 $O(nW)$ 而不是 $O(n^2)$
- $n$ 个 token 各做 $O(W)$ 工作(关注 $W$ 个键/值)
稀疏注意力
将局部注意力与全局注意力交错

RMSNorm
归一化通过防止梯度爆炸/消失来稳定训练
均匀缩放过于严格:不同特征可能需要不同量级
$\gamma$ 结合了归一化的稳定性和每个维度的量级变化
$\gamma\in\mathbb R^D$ 是学习到的每个维度重新缩放
- 归一化步骤强制隐藏状态具有单位 RMS,这破坏了任何学习的尺度信息
- $\gamma$ 将每个特征的控制权交还给网络
- $\gamma_i>1$ → 放大维度 $i$
- $\gamma_i<1$ → 抑制维度 $i$
- $\gamma_i\approx 0$ → 杀死维度 $i$
SwiGLU FFN
$\mathbf G$ 和 $\mathbf U$ 都贡献内容
$\mathbf U$ 提供一个学习的表示,$\mathbf G$ 提供另一个(由其自身置信度自门控)
RoPE
RoPE 只旋转查询和键向量,不旋转值
- 位置信息只需要影响哪些 token 相互关注,而不影响传递什么信息
对于位置 $m$ 处的查询向量 $\mathbf q$ 和位置 $n$ 处的键向量 $\mathbf k$,我们希望点积 $\mathbf q\cdot\mathbf k$ 仅依赖于相对位置 $m-n$
我们希望 $f$ 使得点积 $\langle f(\mathbf q,m),f(\mathbf k, n)\rangle$ 是一个只编码相对形式信息的函数 $g$,例如
$$\langle f(\mathbf q,m),f(\mathbf k, n)\rangle=g(\mathbf q,\mathbf k,n-m)$$
RoPE 嵌入是这样一个解
$$f(\mathbf x,m)=R_{m\theta}\mathbf x\ g(\mathbf q,\mathbf k,n-m)=\mathbf q^\top R_{(n-m)\theta}\mathbf k$$
证明
$$\begin{align*} \langle R_{m\theta}\mathbf q,R_{n\theta}\mathbf k\rangle&=(R_{m\theta}\mathbf q)^\top R_{n\theta}\mathbf k\ &=\mathbf q^\top R_{m\theta}^\top R_{n\theta}\mathbf k\ &=\mathbf q^\top R_{-m\theta}R_{n\theta}\mathbf k &\text{因为 }R_\alpha^\top=R_{-\alpha}\ &=\mathbf q^\top R_{(n-m)\theta}\mathbf k&\text{因为 } R_\alpha R_\beta=R_{\alpha+\beta} \end{align*}$$
内积随距离增加而衰减
对于二维向量 $\mathbf x=[x_1,x_2]$,绕角度 $\theta$ 旋转为
$$R_\theta=\begin{bmatrix}\cos\theta&-\sin\theta\\sin\theta&\cos\theta\end{bmatrix}$$
对于位置 $m$ 和 head 维度索引 $i$,我们旋转 $m\theta_i$
$$R_{m\theta}\mathbf x=\begin{bmatrix} x_1\cos(m\theta_i)-x_2\sin(m\theta_i)\ x_2\sin(m\theta_i)+x_2\cos(m\theta_i)\ \end{bmatrix}$$
对于 $H$ 维嵌入,我们将其划分为 $H/2$ 对,并对每对以不同频率独立旋转
$\Theta$ 是 RoPE 唯一的超参数,通常为 10,000(在 HF 中称为 rotary_base)
定义模型能够原生区分的最大距离
- 当 $m\theta_i=2\pi$ 时完成一次完整旋转 ⇒ $m=\frac{d\pi}{\theta_i}$
- 最慢(最小)的 $\theta_i$ 是 $\Theta^{-1}$ → $m=d\pi\Theta$
- 所以最慢的对在 $2\pi\Theta$ 个位置上完成一个完整圆周
对于每个维度对 $i$,$\theta_i$ 为
$$\theta_i=\Theta^{-2i/H}$$
这提供了指数间距(与原始 transformer 的正弦位置编码相同设计选择)
- 对数均匀的频率分布 → 有效覆盖尺度
- 低频对(大 $i$ → 小 $\theta_i$)旋转缓慢 → 相邻位置之间变化很小 → 编码长程信息
- 高频对(小 $i$ → 大 $\theta_i$)旋转迅速 → 相邻位置差异大 → 实现局部区分
完整的旋转矩阵(对于单个位置)是块对角的,每个 $2\times 2$ 块处理一对维度
$$\mathbf R_m=\begin{bmatrix} \ddots&&\ &\mathbf R_m^{(i)}&\ &&\ddots\ \end{bmatrix}$$
我们可以用复数重新表述
将每对 $[x_{2i},x_{2i+1}]$ 视为复数 $z=x_{2i}+ix_{2i+1}$
绕角度 $\theta$ 旋转等价于乘以 $e^{i\theta}$
Euler定理:$e^{i\theta}=\cos\theta+i\sin\theta$
$$\begin{align*} z e^{i\theta}&=(a+bi)(\cos\theta+i\sin\theta)&\text{Euler定理}\ &=(a\cos\theta-b\sin\theta)+i(a\sin\theta+b\cos\theta) \end{align*}$$
这等价于在复平面上应用旋转矩阵
$$\begin{bmatrix}\cos\theta&-\sin\theta\\sin\theta&\cos\theta\end{bmatrix}\begin{bmatrix}a\b\end{bmatrix}$$
在实践中,我们不构造完整的旋转矩阵,而是做逐元素乘法
$$\begin{bmatrix} x_1\x_2\end{bmatrix}\odot \begin{bmatrix} \cos(m\theta)\ \cos(m\theta) \end{bmatrix}+\begin{bmatrix} -x_2\x_1\end{bmatrix}\odot \begin{bmatrix}\sin(m\theta)\\sin(m\theta)\end{bmatrix}$$
推理
延迟是完成单个请求所需的时间,以秒为单位
吞吐量是所有请求中每秒能处理的 token(或请求)数量,以 token/秒为单位
批处理与打包
- 传统批处理:收集 $N$ 个请求,一起处理,等待所有完成,然后收集下一批;如果一个序列生成 500 个 token,另一个生成 10 个,短的那个会空闲等待
- 连续批处理:一旦一个序列完成,立即将新请求插入其位置,无需等待整个批处理;批处理始终保持满
- 选择性批处理:巧妙混合处于 prefill 和生成阶段的序列;思路:prefill 计算密集,而 decode 内存受限
- 序列打包:将样本连接起来,直到最大序列长度,使用注意力掩码防止交叉污染
- token 预算批处理:将样本分组成批,使批中总 token 数(填充到该批内最大序列长度后)不超过某个预算;通常在微调中这样做;这很有道理,因为 GPU 内存由
batch_size x sequence_length决定
推测解码
推测解码利用 prefill 比生成快的事实
- 从草稿模型 $q$ 生成 $K$ 个 token
- 用目标模型 $p$ 评估这些 token
- 以概率 $\min\left(1,\frac{p(x)}{q(x)}\right)$ 接受每个草稿 token $x$
- 如果 $p(x)>q(x)$,我们肯定接受它
- 如果被拒绝,从调整后的分布 $\max(0, p(x)-q(x))$ 中采样(重新归一化后),从第一个被拒绝的 token 开始
- 始终从教师模型采样一个 token,因为在评分草稿 token 时我们“免费”获得了这些下一个 token 的 logits
$P(\text{发出 token } x) = P(\text{草稿 } x) \times P(\text{接受 } x) + P(\text{采样 token 被拒绝}) \times P(\text{采样 } x)$
- 情况 1(第一项):$x$ 从草稿中被接受:$q(x)\cdot\min\left(1,\frac{p(x)}{q(x)}\right)=\min(q(x),p(x))$
- 情况 2(第二项):$x$ 在拒绝后被选择
- 草稿并拒绝草稿 token $x'$ 的概率:如果 $p(x')>q(x')$ 则为 0;否则 $q(x')\cdot\left(1-\frac{p(x')}{q(x')}\right)=q(x')-p(x')$
- 拒绝的总概率:$P(\text{拒绝})=\sum_{x'}\max(0,q(x')-p(x'))=\sum_{x'}\max(0,p(x')-q(x'))$(第二步因为 $p$ 和 $q$ 都求和为 1)
- 当我们拒绝时,我们从 $\frac{\max(0,p(x)-q(x))}{\sum_{x'}\max(0,p(x')-q(x'))}$ 采样
- 因此拒绝后得到 $x$ 的概率是 $\max(0,p(x)-q(x))$
结合两种情况:$\min(q(x),p(x))+\max(0,p(x)-q(x))=p(x)$
KV 缓存
带缓存的单 token 前向传播
缓存维度应为 (batch, num_heads, max_seq_len, head_dim);实践中,我们在自注意力模块内部使用 self.kv_cache 存储缓存
## 基于 max_seq_len 预分配缓存
kv_cache = [{'k': torch.zeros(batch, num_heads, max_seq_len, head_dim, device='cuda', dtype=torch.float16),
'v': torch.zeros(batch, num_heads, max_seq_len, head_dim, device='cuda', dtype=torch.float16),}
for _ in range(num_layers)]
## 对于注意力,使用 cache[:, :, :position+1, :] 作为键/值
def forward_with_cache(model, new_token, kv_cache, position):
"""
不处理完整序列,只处理新 token
并重用之前位置的缓存 K、V。
"""
# 只嵌入新 token
x = model.embed(new_token) # (batch, 1, d_model)
for layer_idx, layer in enumerate(model.layers):
q, k, v = layer.qkv_proj(x).chunk(3, dim=-1)
# 更新缓存
kv_cache[layer_idx]['k'][:, :, position, :] = k.squeeze(2)
kv_cache[layer_idx]['v'][:, :, position, :] = v.squeeze(2)
# 关注所有缓存的位置
k_full = kv_cache[layer_idx]['k'][:, :, :position+1, :]
v_full = kv_cache[layer_idx]['v'][:, :, :position+1, :]
x = attention(q, k_full, v_full)
x = layer.ffn(x)
return model.lm_head(x)
减小 KV 缓存大小
- 降低 KV 缓存的维度
- 在标准多头注意力(MHA)transformer 注意力中,KV 缓存按
num_layers × num_heads × seq_len × head_dim × 2(K 和 V)缩放 - 在多查询注意力(MQA)中,所有 head 共享相同的 K 和 V,但每个 head 有自己的 Q;KV 缓存缩小
num_heads倍;推理速度和内存效率高得多 - 在分组查询注意力(GQA)中,head 被分为组,每组共享 K 和 V;介于 MHA 和 MQA 之间的折中方案
- MQA 和 GQA 在 head 之间共享 KV,因此我们损失了一些每个 head 的表征能力
- 在标准多头注意力(MHA)transformer 注意力中,KV 缓存按
- 多头潜在注意力(MLA)在 DeepSeek v2 中引入
- 不在 KV 缓存中保留形状为
(seq_len, num_heads × head_dim)的键和值,而是缓存一个较小的潜在向量(seq_len, latent_dim);然后在每个解码步骤,将潜在 KV 投影回完整大小 - MLA 保持每个 head 单独的 KV,但将它们压缩到共享的潜在空间中
- 在推理时增加了一些额外计算(从潜在到 KV 的投影)
- MLA 与 RoPE 不兼容
- 不在 KV 缓存中保留形状为

- 跨层注意力在层之间共享 KV
- 局部注意力:标准注意力是 $O(n^2)$;使用局部注意力使 KV 缓存独立于序列长度;一旦 token 落在窗口外,你可以直接丢弃它
采样策略
def sample(logits, temperature=1.0, top_k=None, top_p=None):
logits = logits / temperature
if top_k is not None:
values, indices = torch.topk(logits, top_k)
logits = torch.full_like(logits, float('-inf'))
logits.scatter_(-1, indices, values)
if top_p is not None:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
# 移除累积概率超过阈值的 token
sorted_mask = cumulative_probs > top_p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
indices_to_remove = sorted_mask.scatter(-1, sorted_indices, sorted_mask)
logits = logits.masked_fill(indices_to_remove, float('-inf'))
probs = F.softmax(logits, dim=-1)
return torch.multinomial(probs, num_samples=1)
Flash Attention
标准注意力:内存问题是 attn_weights 是 (batch, num_heads, seq_len, seq_len);以 seq_len = 8192 和 32 个 head 的 FP16 为例,那就是 $8192^2 \times 32 \times 2$ 字节 ≈ 4GB
import torch
import torch.nn.functional as F
import math
def standard_attention(q, k, v):
# q, k, v 都是 (batch, num_heads, seq_len, head_dim)
scale = math.sqrt(q.size(-1))
# 实例化完整的 N×N 注意力矩阵
attn_weights = torch.matmul(q, k.transpose(-2, -1)) / scale # (batch, num_heads, seq_len, seq_len)
attn_weights = F.softmax(attn_weights, dim=-1)
output = torch.matmul(attn_weights, v) # (batch, num_heads, seq_len, head_dim)
return output
为什么标准注意力是内存受限:$N\times N$ 注意力矩阵被写入 HBM,读回(计算 softmax),再次写入,再次读取(实际使用)
flash attention 使用在线 softmax 技巧在块中计算注意力,将中间结果保留在快速 SRAM 中,而不是将完整的注意力矩阵写入 GPU 主内存
- 将注意力的激活内存从 $O(n^2)$ 减少到 $O(n)$
- 更快,因为内存带宽不再是瓶颈
- 在推理时使用
model.to_bettertransformer()或在from_pretrained()中使用attn_implementation="flash_attention_2"启用
def flash_attention(q, k, v):
# 相同输入形状: (batch, num_heads, seq_len, head_dim)
# PyTorch 自动选择最佳后端(Flash Attention、内存高效或数学)
output = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return output
FlashAttention 从不实例化完整的 $N\times N$ 矩阵,而是以足够小到适合 SRAM 的 tile 计算注意力
- 概念上:
- 加载一块 $Q$(比如 64 行)
- 加载一块 K 和 V(比如 64 列)
- 计算该 tile 的注意力分数,应用 softmax,乘以 V——全都在 SRAM 中
- 只将最终输出写入 HBM
- 对所有 tile 重复
- 由于每个(batch, head)组合完全独立,你有
batch × num_heads个并行工作者处理每个 $m \times head_dim$ 的 Q、K、V tile,其中 $m$ 是 tile 大小 - Flash Attention 是精确的,而不是近似方法
如何使用 Flash Attention:is_causal=True 标志被高效融合到内核中,而不是实例化掩码矩阵
def flash_attention(q, k, v):
# 相同输入形状: (batch, num_heads, seq_len, head_dim)
# PyTorch 自动选择最佳后端(Flash Attention、内存高效或数学)
output = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return output
缩放定律
最大更新参数化($\mu P$)
- 核心问题:在小规模找到的超参数不能推广到更大规模
- 标准参数化 → 当宽度变化时,不同层的更新幅度不一致
- $\mu P$ 逐层调整初始化和学习率,使得更新相对于权重的幅度在宽度变化时保持不变
- 主要调整宽度缩放
拟合学习率与计算量
$$\text{LR}(C)=\beta C^{-\alpha}\implies\log\text{LR}(C)=\log\beta-\alpha\log C$$
拟合损失与计算量
- 我们需要不可约损失项 $\mathcal L_\infty$(数据的熵),否则当 $C\to\infty$ 时,$\mathcal L\to 0$
$$\mathcal L(C)=\mathcal L_\infty+\beta C^{-\alpha}\implies \log(\mathcal L(C)-\mathcal L_\infty)=\log\beta-\alpha\log C$$
如何拟合方程?使用最小二乘法
- 最小二乘法是最小化残差平方和的目标
- 普通(或线性)最小二乘法和非线性最小二乘法
- 线性最小二乘法有闭式解:对于 $y=X\beta$,闭式解是 $\beta=(X^\top X)^{-1}X^\top y$
- 非线性最小二乘法通过迭代细化求解:$S=\sum_i(y_i-f(x_i))^2$
GPU
高带宽内存(HBM)是主 GPU 内存
- 从 GPU 角度看是慢速内存
- A100 上为 40GB 或 80GB
静态 RAM(SRAM)是小型快速片上内存
- 在 A100 上总共约 20MB
其他架构
RNN
普通 RNN
.png?table=block&id=33e7eb87-3605-80dd-99eb-e822a4470aac&spaceId=e06b8491-4010-4103-8a1f-b6bfce58cfba&width=1060&userId=&cache=v2&imgBuildSrc=requestProxiedImageUrl)
在每一步 $t$,处理 $x\in\mathbb R^D$ 和之前的隐藏状态 $h_{t-1}\in\mathbb R^H$,产生下一个隐藏状态 $h_t$($D$ 是输入大小,$H$ 是隐藏大小)
权重矩阵 $W_x\in\mathbb R^{H\times D}$,$W_h\in\mathbb R^{H\times H}$
$$h_t=\tanh(W_{x}x_t+W_{h}h_{t-1}+b)\ y_t=W_\text{out} h_t+b_\text{out}$$
权重在时间步之间共享
梯度消失的原因
令 $z_t=W_xx_t+W_hh_{t-1}+b$
$$\frac{\partial h_t}{\partial h_{t-1}}=\operatorname{diag}(\tanh'(z_t)) W_h$$
$\tanh^\prime$ 只会缩小,因为 $\tanh^\prime\in(0,1]$
重复应用 $W$ 要么导致梯度消失,要么导致梯度爆炸
LSTM
LSTM 使用细胞状态 $c_t$ 以及遗忘门、输入门、输出门
每一步的输入:之前的隐藏状态 $h_{t-1}$、之前的细胞状态 $c_{t-1}$、当前输入 $x_t$
遗忘门:从细胞状态中擦除什么
$$f_t=\sigma(W_f\cdot [h_{t-1}, x_t]+b_f)$$
输入门:写入什么新信息
$$i_t=\sigma(W_i\cdot[h_{t-1},x_t]+b_i)$$
细胞状态候选
$$\tilde c_t=\tanh(W_c\cdot[h_{t-1},x_t]+b_c)$$
细胞状态更新:从旧细胞状态 $c_{t-1}$ 中遗忘一些信息,并从新细胞状态 $\tilde c_t$ 中添加一些信息
$$c_t=f_t\odot c_{t-1}+i_t\odot\tilde c_t$$
输出门:暴露什么作为新的隐藏状态
$$o_t=\sigma(W_o\cdot[h_{t-1},x_t]+b_o)$$
新隐藏状态
$$h_t=o_t\odot\tanh(c_t)$$
只要 $f_t$ 保持接近 1,梯度就可以向后传播许多时间步而不消失
- 细胞状态是记忆,隐藏状态是工作输出
- $c_t$ 沿高速公路流动,只涉及逐元素乘法和加法,没有矩阵乘法或非线性
- 信息仅通过门添加或移除
- 隐藏状态 $h_t=o_t\cdot \tanh(c_t)$ 是细胞状态的过滤视图
- 扮演两个角色:
- 细胞向外的输出
- 细胞在下一步对自身的查询,因为 $t+1$ 处的门是从 $h_t$(而非 $c_t$)计算的
- 扮演两个角色:
- 细胞状态和隐藏状态的分离是 LSTM 工作的原因
- RNN 试图让单个向量同时充当长期记忆和当前输出
GRU 是一个简化版本,有两个门:重置门和更新门
与 Transformer 对比
- 远处时间步的梯度无法影响更早时间步的处理
- RNN vs Transformer
- 注意力提供任意两个 token 之间的直接 / $O(1)$ 路径,不依赖于 token 之间的距离;每个 token 直接查看所有其他 token
- 而在 RNN 中,token 1 的信息必须通过所有中间隐藏状态才能到达 token 100
- 有助于长程依赖
- 注意力同时计算所有成对关系
状态空间模型
SSM 的底层方程
$u_k$ 是一维输入信号
$x_k$ 是 $N$ 维潜在状态
$$x_k=A x_{k-1}+B u_k$$
SSM 在训练中是 $O(n)$ 且可并行,在推理中是 $O(1)$ 且顺序的
相比之下,Transformer 可并行但有 $O(n^2)$ 注意力,RNN 是 $O(n)$ 但顺序的
Mamba 使 $B$、$C$ 和 $\Delta$ 成为输入的函数
$$B_k=f_B(u_k)\ C_k=f_C(u_k)\ \Delta_k=f_\Delta(u_k)$$
- 通过 $B$ 选择性合并信息
- 通过 $C$ 选择性从状态读取
- 通过 $\Delta$ 控制时间尺度
后训练
策略梯度
符号
- 动作 $a_t\in\mathcal V$:时间步 $t$ 的下一个 token
- 状态 $s_t$:时间步 $t$ 处的文本前缀 $(s_0,a_0,\ldots,a_{t-1})$
- $a_t\sim\pi_\theta(\cdot\mid s_t)$:LM 策略
- $s_0\sim p_0$:从提示的起始分布中采样的提示 $s_0$
- $\tau$:轨迹(有限时间范围),也称为 rollout、episode
- $R(\tau)$:轨迹 $\tau$ 的奖励
目标:$\text{最大化 }J(\theta)=\mathbb E_{\tau\sim\pi_\theta} [R(\tau)]$
我们可以通过梯度上升实现:$\theta_{k+1}=\theta_k+\alpha\nabla_\theta J(\theta_k)$
普通策略梯度(也称为 REINFORCE):目标的梯度可以写成
$$\nabla_\theta J(\theta)=\mathbb E_{\tau\sim\pi_\theta}\left[\sum_t\nabla_\theta\log\pi_\theta(a_t\mid s_t) R(\tau)\right]$$
这基本上与 SFT 的更新相同,但数据 $\tau$ 是从策略中采样的,并且梯度由 $R(\tau)$ 加权
- 如果 $R(\tau)$ 为正,我们朝着增加每个 token $a_t$ 在 $\tau$ 中的 $\log\pi_\theta(a_t\mid s_t)$ 的方向移动
- 否则,我们朝相反方向移动
- $R(\tau)$ 的幅度越大,我们走的步长越大
梯度的推导
$$\begin{align*} \nabla_\theta J(\theta)&=\nabla_\theta\mathbb E_{\tau\sim\pi_\theta} [R(\tau)]\ &=\nabla_\theta\sum_\tau P(\tau\mid\theta),R(\tau)\ &=\sum_\tau\nabla_\theta P(\tau\mid\theta),R(\tau)\ &=\sum_\tau P(\tau\mid\theta)\nabla_\theta\log P(\tau\mid\theta),R(\tau)&\text{对数导数技巧: }\nabla P=P,\nabla \log P\ &=\mathbb E_{\tau\sim\pi_\theta}\nabla_\theta\log P(\tau\mid\theta),R(\tau)\ &=\mathbb E_{\tau\sim\pi_\theta}\sum_t\nabla_\theta\log\pi_\theta(a_t\mid s_t), R(\tau) \end{align*}$$
在实践中,我们通过从起始状态 $s_0^{(i)}\sim p_0$ 出发,从策略 $\pi_\theta$ 采样一批 $N$ 个 rollout $\tau^{(i)}$ 来估计 $\nabla_\theta J(\theta)$
$$\hat g=\frac 1N\sum_{i=1}^N\sum_{t=1}^T\nabla_\theta\log\pi_\theta(a_t^{(i)}\mid s_t^{(i)})(R(\tau^{(i)}))$$
所谓的策略梯度损失只是一个标量 pg_loss,使得 pg_loss.backward() 产生与近似策略梯度 $\hat g$ 等价的梯度
$$L(\theta)=\frac 1N\sum_{i=1}^N\sum_{t=1}^T\log\pi_\theta(a_t^{(i)}\mid s_t^{(i)})(R(\tau^{(i)}))$$
- 它不是典型意义上的损失:$L(\theta)$ 不告诉我们策略有多好,它只是产生正确梯度的工具
- 没有固定目标,因为它是由当前策略下采样的数据构建的
带基线的策略梯度
普通策略梯度的问题是方差非常高
- 假设我们有一个简单的提示,所有响应都获得正奖励
- 没有基线,所有响应都被加强,包括批处理中的坏响应!
- 在训练的平均水平上这是可以的,但它会导致非常嘈杂的更新
- 如果我们(例如)使用平均奖励作为基线,那么低于平均水平的响应不会被加强
带基线的策略梯度:从梯度估计中的 $R(\tau)$ 减去基线函数 $b(s_t)$
$$\nabla_\theta J(\theta)=\mathbb E_{\tau\sim\pi_\theta}\left[\sum_{t=0}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t) (R(\tau)-b(s_t))\right]$$
只要 $b(s_t)$ 仅是状态 $s_t$(而非 $a_t$)的函数,它就不会给 $\nabla J(\theta)$ 的估计引入偏差
- 我们希望 $b(s_t)$ 与 $R(\tau)$ 相关,这样 $R(\tau)-b(s_t)$ 很小 → 更少的噪声梯度
- 这不会改变我们优化的目标,只改变目标梯度的估计方式
为什么带基线的策略梯度是策略梯度的无偏估计?
我们将证明基线项 $B$ 的期望值为 0
$$B=\mathbb E_{\tau\sim\pi_\theta}\left[\sum_{t=0}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t)b(s_t)\right]$$
令 $X_t=\nabla_\theta\log\pi_\theta(a_t\mid s_t)b(s_t)$,所以我们重写
$$B=\mathbb E_{\tau\sim\pi_\theta}\sum_{t=0}^T X_t$$
首先将期望移到求和内部
$$B=\sum_{t=0}^T\mathbb E_{\tau\sim\pi_\theta} X_t$$
之前:对于每个轨迹,将 $X_t$ 在所有时间步 $t$ 上求和,然后在采样的轨迹上平均
现在:对于每个时间步 $t$,在采样的轨迹上平均 $X_t$,然后对这些平均值求和
在完整轨迹上平均 $X_t$ 与仅在 $(s_t,a_t)$ 上平均 $X_t$ 相同,因为 $X_t$ 只依赖于 $(s_t,a_t)$(轨迹的其余部分对 $X_t$ 不重要)
然后我们将 $(s_t,a_t)$ 上的联合期望分解为在 $s_t$ 的期望下对 $a_t$ 的条件期望
$$B=\sum_{t=0}^T\mathbb E_{s_t,a_t} X_t=\sum_{t=0}^T\mathbb E_{s_t}\left[\mathbb E_{a_t\mid s_t} X_t\right]$$
现在我们使用 $X_t$ 的定义
$$\begin{align*} \mathbb E_{a_t\mid s_t} X_t&=\mathbb E_{a_t\mid s_t}\nabla_\theta\log\pi_\theta (a_t\mid s_t),b(s_t)\ &= b(s_t),\mathbb E_{a_t\mid s_t}\nabla_\theta\log\pi_\theta (a_t\mid s_t)&\text{因为 }b\text{ 不依赖于 }a_t\text{!}\ &= b(s_t),\sum_{a_t}\pi_\theta(a_t\mid s_t)\nabla_\theta\log\pi_\theta(a_t\mid s_t)\ &= b(s_t),\sum_{a_t}\nabla_\theta\pi_\theta(a_t\mid s_t)&\text{再次使用对数导数技巧}\ &=b(s_t)\nabla_\theta\sum_{a_t}\pi_\theta(a_t\mid s_t)\ &= b(s_t)\nabla_\theta 1&\text{所有动作的概率和为 1}\ &= b(s_t)\cdot 0\ &= 0 \end{align*}$$
注意我们只能将 $b(s_t)$ 从 $a_t\mid s_t$ 的期望中提出来,因为 $b$ 只依赖于 $s_t$,这允许我们使用 $\sum_a\pi(a\mid s_t)=0$
将其代回,我们得到
$$B=\sum_{t=0}^T\mathbb E_{s_t}[0]=0$$
基线函数的一个非常常见的选择(PPO 使用)是 $V_\psi(s_t)$,它估计给定部分序列 $s_t$ 的期望奖励
$$\nabla_\theta J(\theta)=\mathbb E_{\tau\sim\pi_\theta}\left[\sum_{t=0}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t) (R(\tau)-V_\psi(s_t))\right]$$
这是无偏的,因为 $V_\psi$ 只依赖于状态
每个 token 的信号:最终奖励是超过还是低于生成过程中此点看起来可能的水平?
- 将看起来不好的响应(小的 $V_\psi(s_t)$)变成好响应(高的 $R(\tau)$)的 token 获得更多信用
与优势的联系
- $R(\tau)$ 是 Q 函数 $Q^\pi(s_t,a_t)$(给定此动作和状态后遵循策略 $\pi$ 的期望奖励)的单个样本
- $V_\psi(s_t)$ 是价值函数 $V^\pi(s_t)$(仅给定此状态的期望奖励,$V^\pi(s)=\sum_{a\sim\pi(\cdot\mid s)}Q^\pi(s,a)$)的估计
- 所以 $R(\tau)-V_\psi(s_t)$ 是 $Q^\pi(s,a)-V^\pi(s)=A^\pi(s,a)$ 的估计,即优势函数,动作 $a$ 比状态 $s$ 预期好多少
- 特别是,它是优势的蒙特卡洛估计,因为它使用单个样本 $\tau$ 的奖励 $R(\tau)$ 来估计 $Q^\pi$
基线函数通常都试图估计 $\mathcal V^\pi(s_t)$,即从当前状态开始的期望回报
- 学习的价值函数 [如上]
- RLOO:每个提示采样 $G$ 个响应,某个提示的基线是其他奖励的平均值(与 GRPO 非常相似,除了包含 $i$ 和除以标准差归一化)
- 批估计(REINFORCE++):使用批处理中的平均奖励
离策略策略梯度
在线策略策略梯度的问题是,对于每个梯度步骤,我们需要从策略中进行推理,即使策略可能变化不大
在离策略学习中,我们从不同于当前优化的策略采样 rollout
- 通常使用来自旧策略 $\pi_{\theta_\text{old}}$ 的 rollout 来优化当前策略 $\pi_\theta$
- 使用替代目标 $\mathcal J^\text{surrogate}$
$$\mathcal J^\text{surrogate}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_{t=1}^T\underbrace{\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_\text{old}}(a_t\mid s_t)}}_{r_t}, R(\tau)\right]$$
比例 $r_t=\pi_\theta/\pi_{\theta_\text{old}}$ 是重要性采样的重新加权项
重要性采样的背景:如果我们想要 $\mathbb E_{x\sim p} f(x)$ 但只有样本 $x\sim q$,那么我们可以重写
$$\mathbb E_{x\sim p},f(x)=\mathbb E_{x\sim q},\left[\frac{p(x)}{q(x)}f(x)\right]$$
直观地说,如果 $x$ 在 $p$ 下比在 $q$ 下更可能,那么该样本计数更多,反之亦然
注意这优化了与原始 $\mathcal J(\theta)$ 不同的目标!
真正的离策略估计只是用重要性采样重写 $\mathcal J(\theta)$,但由于比例乘积,方差很高
$$\mathcal J(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\frac{P(\tau\mid\theta)}{P(\tau\mid\theta_\text{old})},R(\tau)\right]=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\prod_{t=1}^T\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_\text{old}}(a_t\mid s_t)},R(\tau)\right]$$
$\mathcal J^\text{surrogate}$ 基本上用每个时间步项的求和替代了乘积
$\mathcal J^\text{surrogate}$ 导致以下离策略策略梯度
$$\nabla_\theta \mathcal J^\text{surrogate}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_t\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_\text{old}}(a_t\mid s_t)},\nabla_\theta\log\pi_\theta(a_t\mid s_t), R(\tau)\right]$$
要理解为什么是这样,注意只有 $r_t$ 的分子依赖于 $\theta$,然后应用对数导数技巧
在实践中我们通过以下公式估计
$$\hat g_\text{off-policy}=\frac 1N\sum_{i=1}^N\sum_{t=1}^T\frac{\pi_\theta(a_t^{(i)}\mid s_t^{(i)})}{\pi_{\theta_\text{old}}(a_t^{(i)}\mid s_t^{(i)})}\nabla_\theta\log\pi_\theta(a_t^{(i)}\mid s_t^{(i)})R(\tau^{(i)})$$
其中 $N$ = 每批 rollout 数
PPO
引入了重要性权重 $r_t=\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_\text{old}}(a_t\mid s_t)}$ 的裁剪机制
常规替代目标是
$$\mathcal J^\text{surrogate}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_t r_t A_t\right]$$
PPO 使用价值函数作为基线,因此按照惯例,我们将策略梯度中的 $R(\tau)$ 替换为每个时间步的优势 $A_t=R(\tau)-V_\psi(s_t)$
裁剪后的替代目标将 $r_t$ 裁剪在 $[1-\epsilon,1+\epsilon]$ 内
$$\mathcal J^\text{CLIP}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_t\min(r_t A_t,\text{clip}(r_t,1-\epsilon,1+\epsilon)A_t)\right]$$
裁剪在单批 rollout 上执行多个梯度步时保持稳定性
- 以无偏性换取更稳定的更新
- 裁剪后的目标只要离 $\pi_{\theta_\text{old}}$ 不太远,就是一个好的近似
使用裁剪的四种情况
- 当 $r_t>1+\epsilon$ 且 $A_t>0$ 时:我们使用 $(1+\epsilon)A_t$;$\nabla\mathcal J(\theta)$ 不依赖于 $\theta$ → 梯度为 0;我们已经相对于旧策略大幅增加了 $a_t$ 的概率——停止进一步推动一个好的动作
- 当 $r_t<1-\epsilon$ 且 $A_t<0$ 时:我们使用 $(1-\epsilon)A_t$;梯度也为 0;停止进一步推低一个坏的动作
- 当 $r_t>1+\epsilon$ 且 $A_t<0$ 时:我们使用 $r_t A_t$;$\nabla\mathcal J(\theta)$ 确实依赖于 $\theta$ [参见离策略策略梯度];注意这意味着允许 $r_t$ 大于 $1+\epsilon$ ——这很有道理,因为 $A<0$ 意味着这是一个坏 token,所以我们允许 token 被推低;裁剪是非对称的:它仅在策略已经朝着优势鼓励的方向移动时激活,但不阻止你纠正错误
- 当 $r_t<1-\epsilon$ 且 $A_t>0$ 时:我们使用 $r_t A_t$;$\nabla\mathcal J(\theta)$ 确实依赖于 $\theta$ [参见离策略策略梯度]
PPO 从当前策略收集一批轨迹,然后使用裁剪后的替代目标执行多个梯度步
RLHF
引入 KL 惩罚,防止 $\pi_\theta$ 偏离 $\pi_\text{ref}$ 太远
$$\mathcal J^\text{RLHF}(\theta)=\mathbb E_{\tau\sim\pi_\theta}\left[R(\tau)-\beta D_{\text{KL}}(\pi_\theta\mid\mid\pi_\text{ref})\right]$$
在实践中,KL 惩罚按每个 token 计算,并折叠到每个 token 的奖励中
奖励模型如何训练
Bradley-Terry 模型表示 $y_w$ 优于 $y_l$ 的概率为
$$P(y_w\succ y_l)=\frac{\exp(R(x,y_w))}{\exp(R(x,y_w))+\exp(R(x,y_l))}=\sigma(R(x,y_w)-R(x,y_l))$$
训练带有分类头的 LM 以最大化
$$\mathcal L(\varphi)=-\log P(y_w\succ y_l)$$
这正好是二元交叉熵损失 $-(y\log p+(1-y)\log(1-p))$,真实标签始终为 $y=1$,且 $p=P(y_w\succ y_l)$ 在 Bradley-Terrey 假设下
GRPO
GRPO 将奖励 $R(\tau)$ 替换为 $\tau$ 相对于一组采样轨迹的优势
注意:它并不完全算作基线选择,因为还将 $R(\tau)-b(s_t)$ 除以了归一化因子
对于问题 $q$ 的一组 rollout ${o^{(i)}}_{i=1}^G$,每个包含 token $o^{(i)}_1,\ldots,o^{(i)}_T$
使用 $R(q,o^{(i)})$ 计算每个采样输出的奖励 $\mathbf r = {r^{(i)}}_{i=1}^G$
组相对优势估计 $A^{(i)}$ 为
$$A^{(i)}=\frac{r^{(i)}-\text{mean}(\mathbf r)}{\text{std}(\mathbf r)}$$
优势 $A^{(i)}$ 对响应中的每个 token 都相同
目标为
$$\mathcal J^\text{GRPO-CLIP}(\theta)=\frac 1G\sum_{i=1}^G\frac 1T\sum_{t=1}^T\min\left(r_t A^{(i)},\operatorname{clip}(r_t,1-\epsilon,1+\epsilon)A^{(i)}\right)$$
其中 $r_t$ 是相同的每个 token 概率比例(这里的惯例是使用 $o_t$ 而不是 $a_t$)
$$r_t=\frac{\pi_\theta(o_t\mid s_t)}{\pi_{\theta_\text{old}}(o_t\mid s_t)}$$
通过相对于一组样本计算优势,简化了 PPO,移除了对评论家(价值函数)的需求
GRPO 结合了三个想法
- 使用 $\pi_{\theta_\text{old}}$ 的离策略策略梯度:$\mathcal J^\text{surrogate}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_t r_t R(\tau)\right]$
- 裁剪机制 [PPO]:$\mathcal J^\text{CLIP}(\theta)=\mathbb E_{\tau\sim\pi_{\theta_\text{old}}}\left[\sum_t\min(r_t R(\tau),\text{clip}(r_t,1-\epsilon,1+\epsilon)R(\tau))\right]$
- 使用组归一化计算优势 $A^{(i)}$ [DeepSeek R1]:$\mathcal J^\text{GRPO-CLIP}(\theta)=\frac 1G\sum_{i=1}^G\frac 1T\sum_{t=1}^T\min\left(r_t A^{(i)},\operatorname{clip}(r_t,1-\epsilon,1+\epsilon)A^{(i)}\right)$
算法
- 策略模型 $\pi_\theta\leftarrow\pi_{\theta_\text{init}}$
- 对于每一步(
n_grpo_steps):- 从 $\mathcal D$ 采样一批 $\mathcal D_b$
- 更新旧策略模型 $\pi_{\theta_\text{old}}\leftarrow\pi_\theta$
- 对每个问题 $q\in\mathcal D_b$ 从 $\pi_{\theta_\text{old}}(\cdot\mid q)$ 采样 $G$ 个输出 ${o^{(i)}}_{i=1}^G$
- 使用 $R(q,o^{(i)})$ 计算每个采样输出的奖励 ${r^{(i)}}_{i=1}^G$
- 通过组相对优势估计计算 $A^{(i)}$
- 对于每个训练步(
n_train_steps_per_rollout_batch):- 通过最大化 GRPO 目标来更新策略模型 $\pi_\theta$
- 通过使用回放机制的持续训练更新 $r_\varphi$
GRPO 有效地计算经验基线,而不是学习到的基线
- 不是学习预测期望奖励,而是通过采样来测量
- 每个提示使用更多计算($G$ 个完成)但更少内存(不需要存储价值网络)
Dr. GRPO 修复了 GRPO 的两个问题
- 在 GRPO 中,当 $\text{std}(\mathbf r)$ 很小时(问题太简单或太难),奖励被放大 → 优化该组更重要;这产生了对太简单或太难问题的偏重,使其被更高加权
- 在 GRPO 中,奖励被 rollout 长度归一化 $\frac{1}{\lvert o_i\rvert}$;在正确响应中,短长度 → 更大梯度 → 被更强地强化;在错误响应中,长长度 → 更小梯度 → 惩罚不足;模型学会如果无法得到正确答案,就生成一个非常长的答案
DPO
KL 正则化目标 $\mathcal J^\text{RLHF}$ 有一个闭式最优解
$$\pi^*(y\mid x)=\frac{1}{Z(x)}\pi_\text{ref}(y\mid x)\exp\left(\frac1\beta R(y)\right)$$
重新排列
$$R(y)=\beta\log\frac{\pi^*(y\mid x)}{\pi_\text{ref}(y\mid x)}+\beta\log Z(x)$$
DPO 损失就是观察到的偏好的负对数似然
$$\mathcal L^\text{DPO}(\theta)=-\mathbb E_{(x,y_w,y_l)}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_\text{ref}(y_w\mid x)}-\beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_\text{ref}(y_l\mid x)}\right)\right]$$
精度
混合精度训练
- 主权重在 FP32 中
- 用于前向/反向传播的 BF16 权重副本
- 激活值以 BF16 计算
- 梯度以 BF16 计算并累积到 FP32
- 这避免了将小梯度加到大权重的问题
- BF16 在接近 0 处有更多精度,因此它可以表示小梯度 0.0001,但不能表示更新后的权重 1.0001
- 我们只需要累积的梯度落在 FP32 中,单个梯度通常与权重相比非常小
- 每个参数上的
.grad与参数的数据类型相同,因此单个梯度被转换为 FP32
直觉
- 矩阵乘法对舍入噪声具有容忍性,因此前向/反向传播使用 BF16 没问题
- 主权重需要 FP32,因为单个梯度很小
- 激活值比权重更难量化
精度选项
- FP32:全精度(4 字节)
- FP16:内存减半(2 字节)
- BF16:与 FP16 相同的内存,但数值稳定性更好(更大的动态范围)
- INT8:FP32 内存的四分之一,需要量化(1 字节)
- INT4:更小
数据加载
memmap避免了一次性将整个数据加载到内存中
以 BF16 加载模型
使用 HF 的 .from_pretrained() 加载模型时
torch_dtype=torch.bfloat16用于 BF16,推荐选择;权重和激活值都在 BF16 中load_in_8bit=True用于量化,使用LLM.int8();权重在 INT8 中,激活值在 FP16 中;通过权重压缩节省内存;由于激活值使用更高精度,这不完全等同于全 INT8 计算
model.half() 或 model.to(torch.bfloat16) 将模型转换为 FP16
model = MyModel()
model.load_state_dict(torch.load('model.pt'))
model = model.half()
所有权重为 BF16,所有运算在 BF16 中运行;在实践中,所有运算在 BF16 下运行对推理来说没问题
torch.autocast() 执行自动混合精度
- 权重保持在 FP32,运算选择性地使用 BF16 或 FP32
- 管理每个运算的精度,但不管理主权重
- 内存比将整个模型加载为 BF16 多,但可能更稳定
- 矩阵乘法在 FP16 中(容忍较低精度),softmax 在 FP32 中(需要精度以保证数值稳定性),层归一化在 FP32 中(归约需要精度)
- 当模型在 FP32 中且我们无法轻松转换它,或者看到纯 BF16 有数值问题时有用
dtype指定“较低精度”类型
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
output = model(x)
要使用 bitsandbytes,将 nn.Linear 替换为 bnb.nn.Linear8bitLt 或 bnb.nn.Linear4bit
def replace_linear_with_8bit(model):
"""将所有 nn.Linear 替换为 bnb.nn.Linear8bitLt"""
for name, child in model.named_children():
if isinstance(child, nn.Linear):
# 创建量化替换
new_layer = bnb.nn.Linear8bitLt(
child.in_features,
child.out_features,
bias=child.bias is not None,
has_fp16_weights=False,
)
# 复制权重(移动到 CUDA 时它们会被量化)
new_layer.weight = bnb.nn.Int8Params(
child.weight.data,
requires_grad=False,
)
if child.bias is not None:
new_layer.bias = nn.Parameter(child.bias.data)
setattr(model, name, new_layer)
else:
# 递归进入子模块
replace_linear_with_8bit(child)
return model
## 用法
model = MyTransformer()
model.load_state_dict(torch.load('model.pt'))
model = replace_linear_with_8bit(model)
model = model.to('cuda') # 此处发生量化
并行性
数据并行性将批处理分割到设备上,而模型并行性将单个前向传播的计算分布到设备上
FSDP 将参数分配到各个 rank 以节省内存,但每个 rank 仍然通过在进行前向传播前全 gather 参数来计算完整的前向传播
数据并行性的限制
- 需要 $M<B$,这不一定好,因为我们不希望 $B$ 大于“临界批处理大小”
- 模型仍可能不适合单个设备(即使 ZeRO 阶段 3 也不减少每个设备的激活内存)
强缩放:增加训练芯片数量导致吞吐量(FLOPS/秒)按比例增加
DP 缩放吞吐量,TP/PP 缩放模型内存,SP 缩放激活内存
5D 并行性
- 数据并行性(DP):将一批数据分割到设备上
- 张量并行性(TP):将模型的不同层/阶段分割到设备上
- 流水线并行性(PP):将单个层/权重矩阵分割到设备上
- 序列并行性(SP):将输入序列长度分割到设备上
- 专家并行性(EP):将 MoE 模型中的不同专家分配到不同设备上
背景:核心集体操作
- broadcast(一对所有,相同数据):一个 GPU 拥有数据,并向所有其他 GPU 发送相同副本
- all-gather(所有到所有):每个 GPU 有一块数据,每个 GPU 获得完整集合;沿一个轴移除分片
$$\operatorname{AllGather}_Y:\mathbf A[I,J_Y]\rightarrow \mathbf A[I,J]$$
- reduce-scatter:每个 GPU 有未规约的数据,通过规约合并,结果分片到 GPU 上;与 all-gather 非常相似,但不是保留每个分片,而是将它们求和;沿一个轴添加分片
$$\operatorname{ReduceScatter}_{Y,J}:\mathbf A[I,J]{U_Y}\rightarrow \mathbf A[I,J_Y]$$
- all-reduce:每个 GPU 有未规约的数据,通过规约合并,每个 GPU 获得最终结果
$$\operatorname{AllReduce}_Y\mathbf A[I,J]{U_Y}\rightarrow A[I, J]$$

环形 all-reduce:reduce-scatter + all-gather(每个进程与两个邻居通信)动画
- 一开始,每个 GPU 有未规约的数据
- reduce-scatter:通过规约合并数据,每个 GPU 有一个规约子集;执行所有算术,无冗余复制
- all-gather:每个 GPU 获得规约子集的完整集合;执行所有复制,无算术
- 对于 all-gather、reduce-scatter 和 all-reduce,通信时间仅取决于数组大小和带宽,不取决于数组分片所在的设备数量!
- reduce-scatter 和 all-gather 在彼此的反向传播中使用
- 前向中的 all-gather → 反向中的 reduce-scatter
- all-gather 将同一块广播到每个设备,它们参与不同的下游计算
- 上游梯度在外部分支处求和(如果 $x=a+b$,则 $\partial f/\partial x=\partial f/\partial a \cdot\partial a/\partial x + \partial f/\partial b\cdot \partial b/\partial x$)
- 这正是 reduce-scatter:将所有上游梯度求和到其来源的设备上;前向扇出 → 反向求和
- 前向中的 reduce-scatter → 反向中的 all-gather
- reduce-scatter 将多个输入求和为一个块
- 在反向中,求和节点将上游梯度复制到每个被加数
- 这正是 all-gather:将每个块的梯度广播回其贡献者;前向求和 → 反向扇出
- 这意味着 all-reduce 的反向是另一个 all-reduce!
- 前向中的 all-gather → 反向中的 reduce-scatter
分区记号
- mesh 有轴名为 $(X,Y)$,矩阵 $A$ 有轴 $(I,J)$
- $I_X$:将 $A$ 的行沿设备 mesh 的列分割;沿 $X$ mesh 轴(沿每行)分割 $A$ 的轴 $I$(行)
- $I_Y$:将 $A$ 的行沿设备 mesh 的行分割
- $J_Y$:将 $A$ 的列沿设备 mesh 的行分割
- $J_X$:将 $A$ 的列沿设备 mesh 的列分割
- $I_{XY}$:将 $A$ 的行沿平坦化 $XY$ mesh 的所有设备分割
- $I$:不分割 $A$ 的行
- 未出现的 mesh 维度表示数据沿该维度复制;例如,$Y$ 不出现:每列包含相同数据

矩阵乘法的好性质:当矩阵乘数用块表示时,乘积可以用块矩阵乘法表示
四种矩阵乘法情况
- 情况 1:两个矩阵都没有分片的收缩维度
$$\mathbf A[I_X,J]\cdot \mathbf B[J, K_Y]\rightarrow \mathbf C[I_X, K_Y]$$
不需要通信;执行局部块矩阵乘法;输出自然地以预期方式分片

- 情况 2:$A$ 或 $B$ 有一个分片的收缩维度
$$\mathbf A[I,J_X]\cdot \mathbf B[J, K]\rightarrow \mathbf C[I,K]$$
- all-gather 收集 $A$ 的分片,使每个设备都有完整副本,然后与 $B$ 相乘
$$\operatorname{AllGather}_X[I, J_X]\rightarrow \mathbf A[I, J]\ \mathbf A[I,J]\cdot\mathbf B[J, K]\rightarrow\mathbf C[I,K]$$
- 情况 3:$A$ 和 $B$ 都有分片的收缩维度
$$\mathbf A[I,J_X]\cdot \mathbf B[J_X, K]\rightarrow \mathbf C[I,K]$$
-
矩阵乘法是可能的,但每个设备只表示所需乘积的部分和
-
沿 $X$ 维度的每个设备有不同的部分和
-
使用记号 $C[I,K]{U_X}$ 表示沿 $X$ mesh 轴未规约
-
使用跨 $X$ 轴的 all-reduce 完成最终求和:$\operatorname{AllReduce}_XC[I,K]{U_X}\rightarrow C[I,K]$
-
结果每个设备具有相同的完全求和值
-
情况 4:$A$ 和 $B$ 都有沿同一轴分片的非收缩维度
数据并行性
朴素数据并行性(DDP)
- 将 $B$ 大小的批中的样本分割到 $M$ 个设备上,并交换梯度
步骤
- 在本地微批上运行前向传播
- [加载中...]
- 原文链接: alisawuffles.notion.site...
- 登链社区 AI 助手,为大家转译优秀英文文章,如有翻译不通的地方,还请包涵~