高效推理架构实战:maths-cs-ai-compendium 中的注意力压缩、稀疏化与模型瘦身全解析
【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendium
高效推理并不仅仅是把精度从 FP16 降到 INT8——降低每个操作的"单价"只是其中一条路。真正让模型"变快"的另一个维度,是设计出每个 token 只需做更少工作的架构。本文基于 maths-cs-ai-compendium 第 17 章"AI Inference"的第二篇文档《Efficient Architectures》,系统讲解 StreamingLLM 无限长度生成、稀疏注意力、线性注意力与状态空间模型、MQA/GQA/MLA、Flash/Ring Attention、MoE 推理优化、知识蒸馏、剪枝与神经架构搜索(NAS),并给出可运行的 JAX 对比实验。读完本文,你将掌握从"KV-cache 内存墙"到"每 token 计算量"的完整推理优化工具箱,能够据此为部署任务选择最合适的架构组合。
与本文配套的上一篇 量化(quantisation) 解决的是"每个操作更便宜",本文解决的是"更少发生操作"。两者互补:一个既做了架构精简、又完成量化的模型,可以比原始模型快 10~100 倍。
StreamingLLM:常数内存下的无限长度生成
标准 Transformer 会把此前所有 token 的 Key/Value 存入KV-cache,其大小随序列长度线性增长。当 cache 超过 GPU 显存时,生成立即失败——这是长上下文推理的第一堵"内存墙"(KV-cache 的定量分析见 AI 推理 · 量化 中KV-cache size = 2 × layers × heads × d_head × seq_len × bytes的公式)。
StreamingLLM(Xiao et al., 2023)用固定大小的滚动 KV-cache解决这一问题。其核心洞察是:序列开头的少数几个 token 无论内容如何,都会获得不成比例的高注意力分数,这些 token 被称为attention sinks(注意力汇)。如果把它们从 cache 中驱逐,注意力分布会崩塌,生成质量急剧恶化。
StreamingLLM 的解法是:在 cache 中永久保留少量 sink token(序列开头的前 1~4 个 token),外加一个最近 $w$ 个 token 的滚动窗口。总 cache 大小恒定为 $\text{sink} + w$,与已生成的 token 总数无关:
$$\text{Cache} = [\text{token}_0, \text{token}1, \text{token}{t-w+1}, \ldots, \text{token}_t]$$
注意力汇为 softmax 分布提供锚点,滚动窗口提供近期上下文。这让无限长度生成在常数内存下成为可能,代价是丢失对序列中间部分上下文的访问。对大多数天然形成 attention sinks 的预训练 LLM,StreamingLLM无需任何重训练即可工作;对于不天然具备该特性的模型,训练时加入单个可学习的 sink token 即可修复。
从推理服务的角度看,这套思路与 PagedAttention 的页式 KV-cache、KV-cache 驱逐策略 H2O 同属"控制 cache 规模"这一主线,只是切入角度不同:StreamingLLM 从架构上保证 cache 有界,而 H2O 用启发式决定保留哪些 token。
稀疏注意力:用模式换掉 O(n²)
全量自注意力的代价是 $O(n^2)$(序列长度 $n$):每个 token 都要注意其他所有 token。当 $n = 128K$ 时,注意力矩阵有 $128K^2 = 160$ 亿个元素。稀疏注意力通过限制"谁能注意谁"来削减这一开销。
- 滑动窗口注意力(Mistral、Gemma 采用):每个 token 只注意前 $w$ 个 token(如 $w = 4096$),复杂度从 $O(n^2)$ 降为 $O(n \cdot w)$。信息通过多层堆叠传播到窗口之外:经过 $L$ 层后,有效上下文约为 $L \times w$。
- 局部 + 全局注意力(Longformer、BigBird):大多数 token 使用滑窗(局部),但少数指定 token(如
[CLS]、每第 512 个 token)可以注意全部 token(全局),同时捕获局部模式与长距离依赖。 - 膨胀注意力:在窗口内每隔 $k$ 个 token 采样一个参与注意,用同样数量的注意力分数覆盖更大范围;跨层递增 $k$ 可形成与膨胀卷积(见 第 8 章 · 卷积网络)类似的分层模式。
对现代 LLM 而言,实际胜出的方案是"滑动窗口 + 全注意力交错":部分层用滑窗(便宜、处理局部上下文),部分层用全注意力(昂贵、捕获长距离),Mistral/Mixtral 正是这一模式。仓库 第 7 章 · Transformers 与语言模型 的注意力机制基础(scaled dot-product attention、multi-head attention)是理解这些变体的前提。
线性注意力与状态空间模型:让 O(n²) 彻底消失
能否完全不构造显式的注意力矩阵?线性注意力与状态空间模型(SSM)用 $O(n)$ 时间处理序列,避开 $O(n^2)$ 矩阵:
$$\text{标准: } O = \text{softmax}(QK^T / \sqrt{d}) V$$ $$\text{线性: } O = \phi(Q) (\phi(K)^T V)$$
关键技巧是先做 $K^T V$ 结合——该乘积是 $d \times d$ 维,与序列长度无关。于是计算量从 $O(n^2 \cdot d)$ 降为 $O(n \cdot d^2)$。当 $n \gg d$ 时,这是数量级的节省。
- RWKV融合 RNN 与 Transformer 思想:推理时按序处理 token(类 RNN),训练时仍可并行(类 Transformer)。每 token 推理为 $O(1)$——内存恒定,KV-cache 不增长。
- Mamba(Gu & Dao, 2023)是选择性状态空间模型,通过学习到的状态转移处理序列:
$$h_t = \bar{A} h_{t-1} + \bar{B} x_t, \quad y_t = C h_t$$
其中 $\bar{A}$、$\bar{B}$ 是输入依赖的(选择性),使 Mamba 能动态聚焦或忽略输入的某部分。与固定 SSM 不同,这种选择性让 Mamba 在语言任务上可与 Transformer 竞争,同时保持 $O(n)$ 扩展性。
权衡:线性注意力和 SSM 对长序列更快,但在需要精确长距离检索的任务上通常不如全注意力。混合架构(部分 Transformer 层 + 部分 Mamba 层)常能兼得两者之长——这与 MoE 与混合专家架构、DeepSeek-V3 等前沿模型的路由稀疏设计 中的"异构层"思路一脉相承。
Multi-Query / Grouped-Query Attention:共享 KV 投影
标准多头注意力(MHA,见 第 7 章)为每个 head 保留独立的 K、V 投影,$h$ 个 head 就意味着 KV-cache 里有 $h$ 份独立的 key/value 张量。Multi-Query Attention(MQA)与Grouped-Query Attention(GQA)从"少存"入手:
- MQA(Shazeer, 2019):所有 head 共享同一组 K、V 投影,每个 head 仍保留自己的 Q 投影。KV-cache 缩小 $h$ 倍(32 个 head 即 32 倍)。
- GQA(Ainslie et al., 2023):折中方案。head 被分组,每组共享一组 K/V 投影。$h = 32$ 个 head、$g = 8$ 组时,每组 4 个 head 共享 K/V,KV-cache 缩小 $h/g = 4$ 倍。
$$\text{MHA: } h \text{ heads, } h \text{ K/V sets} \quad \to \quad \text{GQA: } h \text{ heads, } g \text{ K/V sets} \quad \to \quad \text{MQA: } h \text{ heads, } 1 \text{ K/V set}$$
多数现代 LLM(Llama 2/3、Gemma、Mistral)使用 GQA:相比 MHA,KV-cache 内存与推理延迟显著降低,而质量损失可忽略。
Multi-head Latent Attention(MLA)
MLA(DeepSeek-V2, 2024)比 GQA 更进一步:把 KV-cache 压缩进低秩隐空间。不再缓存完整的 key/value 向量,而是为每个 token 缓存一个压缩后的隐向量 $\mathbf{c}_t$,在注意力计算时按需重建 K/V:
$$\mathbf{c}t = W{\text{compress}} \cdot [\mathbf{k}_t; \mathbf{v}_t], \quad \mathbf{k}_t = W_K^{\text{up}} \cdot \mathbf{c}_t, \quad \mathbf{v}_t = W_V^{\text{up}} \cdot \mathbf{c}_t$$
压缩向量 $\mathbf{c}_t$ 远小于原始 K、V 之和。DeepSeek-V2 借此实现相对 MHA93.3% 的 KV-cache 缩减,优于 MQA,同时保持 MHA 级质量。代价是每次注意力操作多出一点重建计算——但由于 LLM decode 阶段受内存带宽约束而非计算约束(见 serving and batching 对 decode 阶段的定量分析),少加载内存的收益大于多算几笔矩阵乘的代价,这仍是净赢。
Flash Attention
Flash Attention(Dao et al., 2022)不是架构改动,而是实现级优化,但任何高效注意力讨论都绕不开它。仓库 第 16 章 · Triton、TPU 与 Pallas 将其作为自定义 kernel 案例研究:标准注意力的 $QK^T$ 矩阵在 $n = 128K$ 时占 64 GB,根本放不进显存;Flash Attention 通过分块(tiling)把计算限制在 SRAM 内,配合online softmax(维护运行最大值、发现新最大值时重缩放已算结果)增量计算 softmax,做到:
- O(n) 内存而非 O(n²)(注意力矩阵从不实例化到 HBM);
- 比标准注意力快 2~4 倍(数据驻留 SRAM,SRAM 访问速度约为 HBM 的 100 倍);
- 零质量损失——输出与标准注意力数学上完全一致。
如 第 16 章 所述,Flash Attention 同时有 Triton 与 CUDA C 实现:CUDA 版本快约 10%,Triton 版本可读性、可修改性更好,更适合研究新的注意力变体。如今它是 PyTorch(torch.nn.functional.scaled_dot_product_attention)、JAX 及所有主流推理框架的默认注意力实现。
Ring Attention
Ring Attention(Liu et al., 2023)解决"单卡显存放不下超长序列"的问题:把序列切分到 $N$ 台设备上,每台持有 $n/N$ 个 token 的 Q、K、V,设备排成环,每步:
- 每台设备计算本地注意力(自己的 Q 对本地 K/V);
- 把 K/V 块发给环上的下一台设备;
- 从上一台设备接收 K/V,对其计算注意力;
- 经过 $N$ 步,每台设备都对全部 K/V 块完成了注意力。
通信与计算重叠:计算当前 K/V 块时,下一块正在传输,通信延迟几乎被隐藏。由此 KV-cache 分布在一圈 GPU 上,每台设备内存为 $O(n/N)$,百万 token 级上下文窗口成为可能——序列长度只受设备数量限制。
推理时的 Mixture of Experts:参数驻留与专家缓存
MoE 模型(第 7 章 · MoE 基础)每个 token 只激活参数的一小部分(典型为 8 个专家中激活 2 个)。推理时的独特挑战是专家缓存:所有专家必须常驻内存(任何 token 都可能路由到任何专家),但每个 token 只激活其中 2 个。
以 Mixtral 8x7B 为例:总参数量 47B(8 × 7B 专家,含共享组件),每 token 激活参数约 13B(2 个专家 + 共享层)。它获得 LLM-70B 级质量、LLM-13B 级推理成本,但需要 47B 参数常驻内存。
- 专家卸载(expert offloading):显存受限部署时,把不活跃的专家放 CPU 或 SSD、按需加载。token 路由具备可预测性,足以预取可能命中的专家。
- 专家缓存(expert caching):在 GPU 内存维护最近使用专家的 LRU 缓存。当相同专家被反复激活(领域内数据常见),缓存命中率很高。
值得留意的是,前沿模型的稀疏路由设计(如 DeepSeek-V3 的辅助损失无关负载均衡、共享专家、Llama 4 的 top-1 + 共享专家)本质上都在调节"稀疏度与路由开销"的平衡,见 第 7 章 · 高级文本生成。
知识蒸馏:用小模型继承大模型行为
蒸馏(第 6 章 · 机器学习)训练一个小的"学生"模型去模仿大的"教师"模型。学生从教师的软预测(类别上的概率分布)中学习,软标签包含的信息远多于硬标签:
$$\mathcal{L} = \alpha \cdot \text{KL}(p_{\text{teacher}}^{T} | p_{\text{student}}^{T}) + (1 - \alpha) \cdot \mathcal{L}{\text{CE}}(y, p{\text{student}})$$
其中 $T$ 是温度($T$ 越高分布越软,越能暴露教师的置信度信息),$\alpha$ 在蒸馏损失与标准交叉熵损失之间取平衡。
- LLM 场景:用大而强的模型造小模型——例如把 GPT-4 级能力蒸馏进一个 7B 学生,使其在特定任务上捕获教师的大部分行为。学生模型的 serving 成本可低 10~100 倍。
- 任务特定蒸馏:只在部署任务相关的数据上蒸馏。一个在医疗问答上从 70B 教师蒸馏出的 7B 模型,在该任务上可能超过 70B 教师——因为学生有限的能力被完全集中在目标领域上。
蒸馏与量化的协同在仓库中同样可见:小模型 + 低精度(量化)常常是边缘端推理(edge inference)的标准组合。
剪枝:移除不需要的权重
剪枝把不必要的权重置零,同时减小模型体积与计算量:
- 非结构化剪枝(基于幅度):移除绝对值最小的单个权重,得到稀疏权重矩阵。简单有效,但当前 GPU 除非稀疏性符合特定模式,否则难以高效加速稀疏运算。
- 结构化剪枝:整单元移除——注意力头、MLP 神经元或整层。产出更小的稠密模型,标准硬件可直接加速。代价是粒度更粗(移除一个 head 可能同时带走有用与无用的权重)。
- 2:4 稀疏(NVIDIA Ampere 及之后):硬件支持的稀疏模式,每 4 个权重中 2 个为零。GPU 的稀疏 Tensor Core 跳过零乘法,获得约 2 倍加速。这是当前唯一有实用硬件加速的稀疏模式。
- 彩票假设(Frankle & Carlin, 2019):随机初始化的网络内存在一个子网络("中奖彩票"),单独训练即可匹配完整网络的表现。寻找这些子网络(训练 → 剪枝 → 回卷)成本高,但其洞察持续激励着剪枝研究。
从实现证据看,剪枝与蒸馏的产出物(更小、更稠密的模型)往往被进一步量化后推向 边缘推理 的 2~4K 上下文窗口场景,形成"架构精简 → 精度压缩 → 边缘部署"的完整链路。
神经架构搜索(NAS):让机器设计架构
NAS在候选架构空间中自动搜索,在延迟、内存、功耗等硬件约束下最大化精度。
- EfficientNet(第 8 章)即由 NAS 发现:其复合缩放规则(同时平衡深度、宽度、分辨率,约束 $\alpha \cdot \beta^2 \cdot \gamma^2 \approx 2$)出自搜索而非人类直觉,基线比例 $\alpha = 1.2$、$\beta = 1.1$、$\gamma = 1.15$ 通过网格搜索得到。
- 面向推理效率,NAS 可针对特定硬件找架构:"在 iPhone 神经引擎上延迟 <5ms、ImageNet 精度 >80% 的模型"。搜索空间包含层类型、宽度、激活函数与注意力模式。
- Once-for-all 网络:训练一个超参数化网络,再为不同部署目标抽取子网络。一次训练同时产出面向云 GPU、移动 GPU、CPU 的模型,各自针对其目标优化。
动手实验:三个可运行的 JAX 任务
以下任务源自原文档的 Coding Tasks 部分,可在 CoLab 或 notebook 中直接运行,用于验证上述理论。
任务 1:滑动窗口注意力 vs 全注意力的内存对比
import jax import jax.numpy as jnp def full_attention(Q, K, V): """Standard O(n^2) attention.""" scores = Q @ K.T / jnp.sqrt(Q.shape[-1]) weights = jax.nn.softmax(scores, axis=-1) return weights @ V def sliding_window_attention(Q, K, V, window_size=128): """Sliding window attention: each token attends to window_size previous tokens.""" n = Q.shape[0] d = Q.shape[-1] output = jnp.zeros_like(Q) for i in range(n): start = max(0, i - window_size + 1) k_window = K[start:i+1] v_window = V[start:i+1] scores = Q[i] @ k_window.T / jnp.sqrt(d) weights = jax.nn.softmax(scores) output = output.at[i].set(weights @ v_window) return output n, d = 512, 64 key = jax.random.PRNGKey(0) Q = jax.random.normal(key, (n, d)) K = jax.random.normal(jax.random.PRNGKey(1), (n, d)) V = jax.random.normal(jax.random.PRNGKey(2), (n, d)) print(f"Full attention memory: O(n^2) = {n*n} entries") print(f"Window (w=128) memory: O(n*w) = {n*128} entries") print(f"Reduction: {n*n / (n*128):.1f}x")任务 2:MHA / GQA / MQA 的 KV-cache 尺寸对比
def kv_cache_size(n_heads, n_kv_heads, d_head, seq_len, bytes=2): """KV-cache size in MB.""" return 2 * n_kv_heads * d_head * seq_len * bytes / 1e6 n_heads = 32 d_head = 128 seq_len = 32768 mha = kv_cache_size(n_heads, n_heads, d_head, seq_len) # 32 KV heads gqa = kv_cache_size(n_heads, 8, d_head, seq_len) # 8 KV heads mqa = kv_cache_size(n_heads, 1, d_head, seq_len) # 1 KV head print(f"MHA (32 KV heads): {mha:.0f} MB per layer") print(f"GQA (8 KV heads): {gqa:.0f} MB per layer ({mha/gqa:.0f}x smaller)") print(f"MQA (1 KV head): {mqa:.0f} MB per layer ({mha/mqa:.0f}x smaller)")以 32 头、$d_{head}=128$、$seq_len=32768$、FP16(2 字节)为例:MHA 每层约 268 MB,GQA(8 KV 头)约 67 MB(小 4 倍),MQA 约 8 MB(小 32 倍)。GQA 之所以是实际甜点位,正因为它用 4 倍缩减换来接近 MHA 的质量——这也是 Llama 2/3、Gemma、Mistral 的一致选择。
任务 3:模拟结构化剪枝——按范数裁掉"最不重要"的注意力头
import jax import jax.numpy as jnp key = jax.random.PRNGKey(0) n_heads, seq_len, d_head = 8, 64, 32 # Random multi-head attention output (one per head) head_outputs = jax.random.normal(key, (n_heads, seq_len, d_head)) # Full output: concatenate all heads full_output = head_outputs.reshape(seq_len, n_heads * d_head) # Importance: measure each head's contribution by its norm head_norms = jnp.linalg.norm(head_outputs, axis=(1, 2)) print("Head importance (by norm):", jnp.round(head_norms, 2)) # Prune least important heads for n_keep in [8, 6, 4, 2]: top_heads = jnp.argsort(head_norms)[-n_keep:] pruned = head_outputs[top_heads].reshape(seq_len, n_keep * d_head) # Pad to original size for comparison (zero out pruned heads) full_pruned = jnp.zeros_like(head_outputs) full_pruned = full_pruned.at[top_heads].set(head_outputs[top_heads]) full_pruned = full_pruned.reshape(seq_len, n_heads * d_head) error = jnp.linalg.norm(full_output - full_pruned) / jnp.linalg.norm(full_output) print(f"Keep {n_keep}/{n_heads} heads: relative error = {error:.4f}, " f"memory = {n_keep/n_heads:.0%}")注意:这里用输出范数作为头部重要性的粗糙代理。真实场景中,重要性度量往往来自激活统计、梯度或下游任务验证集,且常与蒸馏、量化组合成流水线。
小结:如何为你的部署组合这些技术
把本文的技术按"作用层面"归类,可以形成一张决策地图:
| 层面 | 技术 | 适用场景 | 主要收益 |
|---|---|---|---|
| 上下文长度 | StreamingLLM | 对话/流式生成,内存受限 | 常数 KV-cache,无限长度 |
| 注意力复杂度 | 滑窗 / 局部+全局 / 膨胀 / 线性注意力 / SSM | 长序列、移动端 | $O(n^2) \to O(n·w)$ 或 $O(n)$ |
| KV-cache 体积 | MQA / GQA / MLA | 所有自回归 LLM | cache 缩小 4~32 倍乃至 93% |
| 注意力实现 | Flash Attention / Ring Attention | 单卡显存不足或超长上下文 | O(n) 内存、2~4 倍加速、跨设备扩展 |
| 参数利用率 | MoE + 专家缓存/卸载 | 大容量模型部署 | 每 token 只激活部分参数 |
| 模型瘦身 | 蒸馏 / 剪枝(2:4 稀疏) / NAS | 从训练阶段设计高效模型 | 10~100 倍 serving 成本下降 |
这些技术并非互斥:GQA/MLA 压缩 KV-cache,Flash Attention 加速注意力实现,StreamingLLM 限制 cache 上界,再叠加 量化 把每操作成本压下来,最后通过 serving 与 batching 的 PagedAttention、大规模部署 的前缀缓存与序列并行拼装成完整的高效推理栈——这也是本仓库第 17 章各文件间的内在逻辑:从精度、架构、调度、边缘到规模,逐层把"快"落到工程实处。
【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendium
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考