高效架构 让模型变快不只是降低精度,还要用更聪明的架构让每个 token 做更少的活。本文件涵盖 StreamingLLM、稀疏与线性注意力、多查询/分组查询注意力、推理时的混合专家、知识蒸馏、剪枝和神经架构搜索 量化(文件 01)让每次运算更便宜,本文件则让运算本身发生得更少。两者是互补的:一个架构高效又经过量化的模型,可以比原始模型快 10-100 倍。 StreamingLLM:无限长度生成 标准 transformer 把所有历史 token 都存在 KV-cache 里,缓存随序列长度线性增长。到某个点,缓存会超过 GPU 显存,生成就崩了。StreamingLLM(Xiao 等,2023)用一个固定大小的滚动 KV-cache(rolling KV-cache)解决了这个问题。
让模型变快不只是降低精度,还要用更聪明的架构让每个 token 做更少的活。本文件涵盖 StreamingLLM、稀疏与线性注意力、多查询/分组查询注意力、推理时的混合专家、知识蒸馏、剪枝和神经架构搜索
标准 transformer 把所有历史 token 都存在 KV-cache 里,缓存随序列长度线性增长。到某个点,缓存会超过 GPU 显存,生成就崩了。StreamingLLM(Xiao 等,2023)用一个固定大小的**滚动 KV-cache(rolling KV-cache)**解决了这个问题。
关键观察:序列最前面的几个 token,无论内容如何,都会收到不成比例地高的注意力分数。它们被称为注意力沉淀(attention sinks)。如果把它们从缓存里赶走,注意力分布会塌掉,生成质量会灾难性下降。
StreamingLLM 的解法:永久保留少量沉淀 token(sink tokens)(最前面 1-4 个),再加一个最近 w 个 token 的滑动窗口(rolling window)。总缓存大小是 \text{sink} + w,无论已经生成了多少 token 都保持固定。
注意力沉淀锚定 softmax 分布,滑动窗口提供近期上下文。这就实现了常数显存的无限长度生成,代价是失去对序列中段上下文的访问。
对于天然形成注意力沉淀的模型(大多数预训练 LLM 都是),StreamingLLM 无需重训即可工作。对于没有的模型,训练时加一个可学习的沉淀 token 就能修好。
滑动窗口注意力(sliding window attention)(Mistral、Gemma):每个 token 只 attend 前面的 w 个 token(例如 w = 4096)。注意力开销变成 O(n \cdot w) 而非 O(n^2)。信息通过多层传播越过窗口:经过 L 层后,有效上下文是 L \times w。
局部 + 全局注意力(local + global attention)(Longformer、BigBird):大多数 token 用滑动窗口(局部),但少数指定 token(例如 [CLS]、每第 512 个 token)attend 到所有 token(全局)。这样既抓局部模式又抓长程依赖。
空洞注意力(dilated attention):在窗口内只 attend 每第 k 个 token,形成一种稀疏模式,用同样多的注意力分数覆盖更大范围。跨层增大空洞率,会形成类似空洞卷积(第 8 章)的层级模式。
现代 LLM 的实际赢家是滑动窗口 + 完整注意力交替:部分层用滑动窗口(便宜、处理局部上下文),部分层用完整注意力(贵、抓长程)。Mistral/Mixtral 用的就是这个模式。
我们能不能完全干掉 O(n^2) 注意力?**线性注意力(linear attention)和状态空间模型(state-space models,SSM)**通过避免显式注意力矩阵,把序列处理做到 O(n)。
线性注意力用核近似替换 softmax 注意力:
先结合 K^T V(大小是 d \times d,与序列长度无关),计算就变成 O(n \cdot d^2) 而非 O(n^2 \cdot d)。对于 n \gg d 的长序列,这是巨大的节省。
RWKV 结合了 RNN 和 transformer 的思想。它用循环形式按顺序处理 token(像 RNN),但训练时可以并行(像 transformer)。推理时每 token 是 O(1)(常数显存,KV-cache 不增长)。
Mamba(Gu & Dao,2023)是一种选择性状态空间模型。它通过学习到的状态转移处理序列:
其中 \bar{A}、\bar{B} 依赖于输入(选择性的),让 Mamba 能动态聚焦或忽略输入的某些部分。不同于固定 SSM,这种选择性让 Mamba 在语言任务上能与 transformer 抗衡,同时保持 O(n) 的扩展性。
权衡:线性注意力和 SSM 在长序列上更快,但对于需要精确长程检索的任务,一般不如完整注意力能干。混合架构(部分 transformer 层 + 部分 Mamba 层)往往能两全其美。
标准多头注意力(MHA,第 7 章)每个头都用独立的 K、V 投影。对于 h 个头,意味着 KV-cache 里有 h 组独立的 key 和 value 张量。**多查询注意力(Multi-Query Attention,MQA)和分组查询注意力(Grouped-Query Attention,GQA)**就是来减小它的。
MQA(Shazeer,2019):所有头共享同一组 K, V 投影,但每个头仍有自己的 Q 投影。KV-cache 缩小 h 倍(例如 32 个头就缩小 32 倍)。
GQA(Ainslie 等,2023):折中方案。把头分组,每组共享一组 K, V 投影。h = 32 个头、g = 8 组时,每 4 个头共享一组 K/V。KV-cache 缩小 h/g = 4 倍。
压缩向量 \mathbf{c}_t 比原始 K、V 加起来小得多。DeepSeek-V2 相比 MHA 实现了 93.3% 的 KV-cache 缩减,甚至超过 MQA,同时保持 MHA 级质量。
权衡:从潜在向量重建 K/V 给每次注意力增加一点计算开销。但因为 LLM 解码是访存带宽受限(不是计算受限),这是净赚:少搬的显存 > 每 token 多算的那点。
Flash Attention(Dao 等,2022,第 16 章文件 05 有详细讲解)不是架构改动,而是一种实现优化,但凡讨论高效注意力都绕不开它。它计算的是与标准注意力完全等价的结果,但:
Flash Attention 现在是 PyTorch(torch.nn.functional.scaled_dot_product_attention)、JAX 和所有主流推理框架的默认注意力实现。如果你在 2024 年以后跑注意力,几乎肯定用的就是 Flash Attention。
Ring Attention(Liu 等,2023)针对那些即使有 Flash Attention 也装不下单卡显存的超长序列,把注意力计算分布到多个设备上。
思路:把序列切分到 N 个设备,每个设备持有 n/N 个 token 的 Q、K、V。设备排成一个环。每一步:
通信与计算重叠:在算当前 K/V 块的注意力时,下一块正在传输。这几乎完全隐藏了通信延迟。
Ring Attention 通过把 KV-cache 分布到一圈 GPU 上,实现了百万级 token 的上下文窗口。每设备显存是 O(n/N),让任意长序列都变得可行(只受设备数量限制)。
MoE 模型(第 7 章)每个 token 只激活一小部分参数(通常是 8 个专家里激活 2 个)。推理时的独特挑战是专家缓存:所有专家都必须在显存里(因为任何 token 都可能路由到任何专家),但每个 token 只激活 2 个。
以 Mixtral 8x7B 为例:总参数 = 47B(8 × 7B 专家,但有共享组件)。每个 token 激活的参数约 13B(2 个专家 + 共享层)。它有 LLM-70B 级的质量、LLM-13B 级的推理成本,但需要 47B 参数都在显存里。
专家卸载(expert offloading):对 GPU 显存吃紧的部署,把不活跃的专家放在 CPU 或 SSD 上,按需加载。这之所以管用,是因为 token 路由足够可预测,可以预取可能用到的专家。
专家缓存:在 GPU 显存里维护一个最近用过的专家的 LRU 缓存。如果同样的专家被反复激活(领域内数据很常见),缓存命中率就很高。
其中 T 是温度(T 越大分布越软,揭示老师的不确定性),\alpha 用来平衡蒸馏损失和标准交叉熵损失。
对 LLM:蒸馏用于从大的、能力强的模型做出小的、快的模型。比如 GPT-4 → 一个 7B 学生,能捕获 GPT-4 在特定任务上大部分的行为。学生的服务成本可以便宜 10-100 倍。
任务特定蒸馏:只在与部署任务相关的数据上蒸馏。一个从 70B 老师蒸馏出来的、针对医疗问答的 7B 模型,在那个特定任务上能超过 70B 模型(因为学生有限的容量完全聚焦在目标领域)。
**剪枝(pruning)**去掉不必要的权重(置零),从而减小模型和计算量。
非结构化剪枝(unstructured pruning)(基于幅值):移除绝对值最小的单个权重,得到一个稀疏权重矩阵。压缩上简单有效,但当前硬件(GPU)无法高效加速稀疏运算,除非稀疏性遵循特定模式。
结构化剪枝(structured pruning):移除整个单元——注意力头、MLP 神经元或整层。得到一个更小的稠密模型,在标准硬件上很易加速。代价是粒度更粗(移除一整个头可能同时丢掉有用和无用的权重)。
2:4 稀疏(2:4 sparsity)(NVIDIA Ampere+):一种硬件支持的稀疏模式,每 4 个权重中有 2 个为零。GPU 的稀疏 Tensor Core 会跳过零乘法,达到约 2 倍加速。这是当今唯一有实用硬件加速的稀疏模式。
彩票假说(Lottery Ticket Hypothesis)(Frankle & Carlin,2019):在一个随机初始化的网络里,存在一个子网络("中奖彩票"),单独训练就能达到整个网络的性能。找这些子网络(先训练、再剪枝、再回退)很贵,但这个洞见推动了剪枝研究。
NAS 通过在可能的架构空间里搜索,自动完成架构设计,找到在硬件约束(延迟、显存、功耗)下精度最高的那个。
EfficientNet(第 8 章)就是 NAS 找出来的:复合缩放规则(平衡深度、宽度、分辨率)是从搜索中涌现的,不是来自人的直觉。
对于推理效率,NAS 能找到针对特定硬件目标优化的架构:"找一个在 iPhone 神经引擎上延迟 <5ms、在 ImageNet 上精度 >80% 的模型。"搜索空间包括层类型、宽度、激活函数和注意力模式。
**一次训练多部署网络(once-for-all networks)**训练一个过参数化的网络,再为不同部署目标抽取子网络。一次训练同时产出面向云端 GPU、移动 GPU 和 CPU 的模型,每个都针对自己的目标优化。
import jax import jax.numpy as jnp def full_attention(Q, K, V): """标准 O(n^2) 注意力。""" 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): """滑动窗口注意力:每个 token 只 attend 前面 window_size 个 token。""" 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")
def kv_cache_size(n_heads, n_kv_heads, d_head, seq_len, bytes=2): """KV-cache 大小,单位 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 头 gqa = kv_cache_size(n_heads, 8, d_head, seq_len) # 8 个 KV 头 mqa = kv_cache_size(n_heads, 1, d_head, seq_len) # 1 个 KV 头 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)")
import jax import jax.numpy as jnp key = jax.random.PRNGKey(0) n_heads, seq_len, d_head = 8, 64, 32 # 随机的多头注意力输出(每个头一个) head_outputs = jax.random.normal(key, (n_heads, seq_len, d_head)) # 完整输出:拼接所有头 full_output = head_outputs.reshape(seq_len, n_heads * d_head) # 重要性:用每个头的范数衡量其贡献 head_norms = jnp.linalg.norm(head_outputs, axis=(1, 2)) print("Head importance (by norm):", jnp.round(head_norms, 2)) # 剪掉最不重要的头 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) # 补零到原始大小以便对比(被剪掉的头置零) 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%}")