Alisa 的 LLMs 手册

alisawuffles 发布于 2026-06-16 阅读 63

本文是一本系统的大语言模型技术手册,涵盖从神经网络基础(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 &lt; 0:
            raise ValueError(f"无效学习率: {lr}")
        if not 0 &lt; betas[0] &lt; 1 or not 0 &lt; betas[1] &lt; 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 前的注意力分数

假设约定中 maskTrue 表示可以参与的位置

tensor.masked_fill(mask, value)maskTrue 的位置用 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_headsseq_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 比生成快的事实

  1. 从草稿模型 $q$ 生成 $K$ 个 token
  2. 用目标模型 $p$ 评估这些 token
  3. 以概率 $\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 的表征能力
  • 多头潜在注意力(MLA)在 DeepSeek v2 中引入
    • 不在 KV 缓存中保留形状为 (seq_len, num_heads × head_dim) 的键和值,而是缓存一个较小的潜在向量 (seq_len, latent_dim);然后在每个解码步骤,将潜在 KV 投影回完整大小
    • MLA 保持每个 head 单独的 KV,但将它们压缩到共享的潜在空间中
    • 在推理时增加了一些额外计算(从潜在到 KV 的投影)
    • MLA 与 RoPE 不兼容

  • 跨层注意力在层之间共享 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

在每一步 $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)$ 是细胞状态的过滤视图
    • 扮演两个角色:
      1. 细胞向外的输出
      2. 细胞在下一步对自身的查询,因为 $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}}$ 不太远,就是一个好的近似

使用裁剪的四种情况

  1. 当 $r_t>1+\epsilon$ 且 $A_t>0$ 时:我们使用 $(1+\epsilon)A_t$;$\nabla\mathcal J(\theta)$ 不依赖于 $\theta$ → 梯度为 0;我们已经相对于旧策略大幅增加了 $a_t$ 的概率——停止进一步推动一个好的动作
  2. 当 $r_t<1-\epsilon$ 且 $A_t<0$ 时:我们使用 $(1-\epsilon)A_t$;梯度也为 0;停止进一步推低一个坏的动作
  3. 当 $r_t>1+\epsilon$ 且 $A_t<0$ 时:我们使用 $r_t A_t$;$\nabla\mathcal J(\theta)$ 确实依赖于 $\theta$ [参见离策略策略梯度];注意这意味着允许 $r_t$ 大于 $1+\epsilon$ ——这很有道理,因为 $A<0$ 意味着这是一个坏 token,所以我们允许 token 被推低;裁剪是非对称的:它仅在策略已经朝着优势鼓励的方向移动时激活,但不阻止你纠正错误
  4. 当 $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 结合了三个想法

  1. 使用 $\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]$
  2. 裁剪机制 [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]$
  3. 使用组归一化计算优势 $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)$

算法

  1. 策略模型 $\pi_\theta\leftarrow\pi_{\theta_\text{init}}$
  2. 对于每一步(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 的两个问题

  1. 在 GRPO 中,当 $\text{std}(\mathbf r)$ 很小时(问题太简单或太难),奖励被放大 → 优化该组更重要;这产生了对太简单或太难问题的偏重,使其被更高加权
  2. 在 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.Linear8bitLtbnb.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!

分区记号

  • 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$ 个设备上,并交换梯度

步骤

  1. 在本地微批上运行前向传播
  2. [加载中...]
  • 原文链接: alisawuffles.notion.site...
  • 登链社区 AI 助手,为大家转译优秀英文文章,如有翻译不通的地方,还请包涵~

相关文章

0 条评论