一句话总结:Sparse Attention 少算连接,Linear Attention 改写相似度计算,Infini-attention 把跨 Segment 的旧 K/V 压缩进固定状态;三者解决的并不是同一个问题,也都不会免费保留完整历史。
24.1 Full Attention 的三笔账
24.1.1 连接数
长度为 的 Full Self-Attention,要比较每个 Query 和每个 Key。双向 Attention 有 个连接;因果 Attention 允许的下三角连接是:
不过,“逻辑上只有下三角”不代表一个普通稠密矩阵乘法自动少算一半。是否真的跳过被遮住的区域,取决于所用 Kernel。
| 双向连接数 | 因果有效连接数 | |
|---|---|---|
| 1,024 | 1,048,576 | 524,800 |
| 4,096 | 16,777,216 | 8,390,656 |
| 32,768 | 1,073,741,824 | 536,887,296 |
| 1,000,000 | 1,000,000,000,000 | 500,000,500,000 |
序列变长 10 倍,连接数大约变成 100 倍。这是第一个问题:算术工作仍是二次增长。
24.1.2 中间结果的存储
如果真的把一个 Head 的 score 矩阵以 FP16/BF16 保存下来,它会占用:
| 一个 Head 的稠密 score 矩阵 | |
|---|---|
| 4,096 | 32 MiB |
| 8,192 | 128 MiB |
| 32,768 | 2 GiB |
| 131,072 | 32 GiB |
这些数字只是一个 Head、一个矩阵,不是“整个单层只占这么多”。若真的按所有 Head 保存,还要乘 Head 数;训练还涉及反向传播所需状态。
但第 21 章的 FlashAttention 已经告诉我们:精确 Softmax Attention 不必把完整 score 矩阵写入 HBM。它可以把工作内存降到随序列近似线性增长,同时保持完全相同的数学结果。FlashAttention 解决的是 I/O 和中间存储,不会消除所有 Query–Key 配对的二次算术。
24.1.3 不要把 KV Cache 混进来
第 23 章的 KV Cache 随 token 数线性增长;Full Attention 的训练/Prefill 配对数随长度二次增长。这是两条不同的轴:
- FlashAttention:不保存完整 score 矩阵,数学仍是 Full Attention;
- GQA/MQA:减少 K/V Head 和 Cache 字节数;
- Sparse Attention:减少允许计算的 Query–Key 连接;
- Linear Attention:换掉 Softmax 相似度的计算形式;
- Infini-attention:用固定状态压缩跨 Segment 历史,同时保留当前 Segment 内的局部精确 Attention。
把这几件事分开,后面的复杂度才不会乱。
24.2 Sparse Attention:先画图,再谈 O(N)
想象你读到《三体》第 500 页。眼前一句话最可能先依赖附近几段;某些章节标题或人物索引可以充当全局锚点;少量跨章节连接能缩短远距离信息传播的路径。Sparse Attention 就是把“谁能看谁”画成一张稀疏图。
常见的三种边是:
- Window / Local:每个 Query 只看附近 个 Key。
- Global:少数 个位置与全序列双向连接,像 Longformer 中任务指定的全局 token。
- Random:每个位置再连接 个随机位置,BigBird 用它和 Window、Global 一起构图。
若每行只有固定数量的这些连接,边数的主项是:
只有当 不随 增长时,才可以简写成 。如果窗口跟着序列一起扩大,复杂度也会跟着变。
24.2.1 Longformer 和 BigBird 的结论边界
Longformer 将滑动窗口与任务指定的全局 Attention 组合起来,原论文主要研究长文档编码、MLM 和 LED 编码器—解码器。BigBird 再加入随机连接,原论文主要也是 BERT 风格的长序列编码任务。
BigBird 论文证明,在它给定的图结构、全局 token、层数和精度等假设下,稀疏 Transformer 仍可成为通用逼近器并保持图灵完备性。这不等于:
- 任意随机图都保证两个 token 只隔固定 层;
- 一个有限深度、训练完成的 BigBird 与 Full Attention 表现必然相同;
- 所有任务的精度损失都“很轻微”。
理论表达能力、图上的传播路径、训练后质量,是三件不同的事。
24.2.2 Mask 不等于省计算
在一个普通 score 矩阵上填 -inf,可以得到正确的 Sparse Attention 结果,但矩阵已经算出来了。真正节省工作,需要 Block-sparse、FlexAttention 一类能够利用结构的 Kernel。块大小、稀疏率、形状和硬件利用率决定了实际速度;稀疏不够时,调度开销甚至可能吃掉理论收益。
24.3 Sliding Window:最容易理解,也最容易差一格
下面我们只讲 decoder 的因果窗口。约定 window=3 表示每个位置最多看“自己 + 前两个位置”:
q0 -> k0
q1 -> k0 k1
q2 -> k0 k1 k2
q3 -> k1 k2 k3
...
长度 8 时,有效连接数不是 ,而是 ;左边界少了三条。
import math
import torch
def causal_window_mask(length, window, device=None):
if length <= 0 or window <= 0:
raise ValueError("length and window must be positive")
positions = torch.arange(length, device=device)
distance = positions[:, None] - positions[None, :]
return (distance >= 0) & (distance < window)
def masked_attention(q, k, v, mask):
"""Dense teaching reference; use a structured sparse kernel for real savings."""
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
scores = scores.masked_fill(~mask, float("-inf"))
return torch.softmax(scores, dim=-1) @ v
这段代码是用来确认 mask 语义的稠密参考,不是高性能 Sparse 实现。PyTorch 的布尔 SDPA/FlexAttention mask 语义也不完全相同,移植时要读当前 API,而不是凭 True 到底代表“保留”还是“屏蔽”来猜。
24.3.1 多层后的感受野
按上面“含当前位置共 window 个 token”的定义,没有 dilation 和全局 token 时,堆叠 层后的最大因果跨度是:
有些论文把 定义成向后看的距离,于是会写成大约 。两种说法可能只差定义,不要把窗口直径、半径和包含当前位置混在一起。
24.3.2 Mistral 7B 的历史例子
Mistral 7B v0.1 技术报告给出的配置是 32 层、Sliding Window ,并据此写出约 131K token 的理论信息传播跨度。同一张配置表的 context_len 是 8192;131K 不是说 v0.1 检查点已在 131K 上逐 token 验证过精确回忆。
报告还描述了 Rolling Buffer Cache:位置 的 K/V 写入 ,超过窗口的旧条目被覆盖。Cache 在窗口填满后不再随总生成长度增长;但被覆盖位置的精确 K/V 已不可访问,它们只能通过较新 token 的隐藏状态间接影响后续层。
24.4 Linear Attention:不是把 Softmax 挪个括号
若没有 Softmax,矩阵结合律确实给出:
左边先产生 ,右边先产生 。但标准 Attention 是:
Softmax 对每个 Query 的整行分数做归一化,不能简单穿过括号。因此常说的 Kernelized Linear Attention 定义了一个非负特征映射 ,把相似度改成:
对于 causal Attention,维护两个前缀状态:
输出为:
论文中的一个具体选择是 :
import torch
import torch.nn.functional as F
def phi(x):
return F.elu(x) + 1.0
def causal_kernel_attention(q, k, v):
"""One head: q/k [N,C], v [N,M]."""
state = q.new_zeros(q.size(1), v.size(1))
normalizer = q.new_zeros(q.size(1))
outputs = []
for q_t, k_t, v_t in zip(phi(q), phi(k), v):
state = state + torch.outer(k_t, v_t)
normalizer = normalizer + k_t
outputs.append((q_t @ state) / (q_t @ normalizer).clamp_min(1e-12))
return torch.stack(outputs)
若特征维度为 、Value 维度为 ,主项是 ;流式 Decode 的状态是 ,相对历史长度固定。它什么时候比 Full Attention 快,还取决于 、并行度和 Kernel。更重要的是,这段代码精确实现的是这个 Kernel Attention,不是标准 Softmax Attention 的无误差重排。
本章配套测试把流式前缀状态与显式构造的下三角 Kernel Attention 比较,float32/float64 的输出和梯度均一致。
24.5 Infini-attention:固定的是跨 Segment 状态
Infini-attention 把输入切成固定长度的 Segment。每一层、每个 Head 同时做两件事:
- 当前 Segment 内执行标准的 causal scaled dot-product Attention,保留局部精细信息;
- 用上一个 Segment 留下的压缩状态检索更久远的历史,并在当前 Segment 结束后更新状态。
对一个 Head,状态不是简单的 d_model × d_model,而是:
M: [d_key, d_value]
z: [d_key]
从旧状态检索:
在输出当前 Segment 后,用它的 K、V 更新:
论文还实验了 Delta 更新;这里先保留它最直接的 Linear 版本。局部输出 与记忆输出通过每 Head 一个可学习标量融合:
import torch
import torch.nn.functional as F
def phi(x):
return F.elu(x) + 1.0
def retrieve_memory(q, memory, normalizer):
q_features = phi(q)
return (q_features @ memory) / (
q_features @ normalizer
).unsqueeze(-1).clamp_min(1e-12)
def infini_segment(q, k, v, memory, normalizer, gate_logit):
"""One head and one segment; q/k [N,D_k], v [N,D_v]."""
memory_output = retrieve_memory(q, memory, normalizer)
local_output = F.scaled_dot_product_attention(
q[None, None], k[None, None], v[None, None],
is_causal=True, dropout_p=0.0,
)[0, 0]
gate = torch.sigmoid(gate_logit)
output = gate * memory_output + (1.0 - gate) * local_output
k_features = phi(k)
next_memory = memory + k_features.T @ v
next_normalizer = normalizer + k_features.sum(dim=0)
return output, next_memory, next_normalizer
注意代码顺序:当前 Segment 先从 读取,算完输出后才写入 。这样记忆分支不会偷看当前 Segment 的未来位置;当前 Segment 的因果关系由局部 Attention 负责。
24.5.1 “无限”到底是什么意思
固定的是跨 Segment 状态:每层每 Head 需要 个数,与已经处理过多少 Segment 无关。若 Segment 长度固定为 、总长度为 :
- 局部 Attention 总算术约为 ,不是总计算 ;
- 压缩记忆的状态大小相对 是 ;
- 当前 Segment 的临时工作内存仍取决于 和所用 Attention Kernel。
所以论文标题里的 infinite 指“可持续处理不设总上限的 Segment 流,并保持有界状态”,不是“所有历史 token 都能被无损、随机访问”。不同 K/V 绑定不断叠进同一矩阵,冲突和信息损失是压缩本身的一部分。
24.5.2 论文结果与 Gemini 1.5 不是一件事
Infini-attention 论文报告:经过相应训练后,1B 模型完成了 1M 长度的 passkey retrieval,8B 模型用于 500K 长度的书籍摘要;这些是论文里的特定模型、任务和训练设置。
Gemini 1.5 技术报告在 2024 年 3 月公开,Infini-attention 的首个 arXiv 版本在 2024 年 4 月公开。两份公开报告都没有说 Gemini 1.5 使用 Infini-attention。仅凭“都能处理很长上下文”把二者画等号,是推测,不是架构事实。
24.6 把方法放回正确的坐标轴
| 方法 | 改了什么 | 对总长度 的主项 | 历史是否精确保留 | 关键条件 |
|---|---|---|---|---|
| Full + FlashAttention | Kernel / I/O | 算术 | 在窗口内是 | 不保存完整 score 到 HBM |
| Sliding Window | 连接图 | 窗外只有间接影响 | 固定且 Kernel 利用稀疏结构 | |
| Longformer / BigBird | 局部、全局、随机图 | 只保留允许的边 | 原论文结论有任务与图结构边界 | |
| Kernel Linear Attention | 相似度函数与计算顺序 | 压进 | 不是 Softmax Attention 的简单等价式 | |
| Infini-attention | Segment 间的循环压缩状态 | 固定 Segment 时随 线性 | 跨 Segment 有压缩损失 | 需要相应训练;局部 Attention 仍存在 |
这些方法可以在数学上兼容某些组合,但“能组合”不等于任意框架已经提供高效融合 Kernel。FlashAttention、GQA、Block-sparse mask、Rolling Cache 和压缩记忆分别作用在不同层面;每加入一层,都要重新验证正确性、质量、峰值显存与真实吞吐。
24.7 本章总结
- Full Attention 有二次配对成本;显式 score 存储和线性 KV Cache 是另外两笔账。
- Sparse Attention 只有配合能跳过无效块的 Kernel 才真的省计算;稠密矩阵加 mask 只是得到稀疏结果。
- Longformer 与 BigBird 的线性复杂度要求窗口、全局和随机连接数不随序列增长;理论表达力不等于任意任务上的质量保证。
- Linear Attention 用特征映射定义新的 Kernel 相似度;它不是把标准 Softmax Attention 换一下括号。
- Infini-attention 的 状态相对总历史固定,但总计算仍随处理的 Segment 数增长,跨 Segment 历史也不是无损保存。
- Gemini 1.5 使用 Infini-attention 的说法没有公开证据,不能写成事实。
参考文献
- Longformer: The Long-Document Transformer(Beltagy et al., 2020)
- Big Bird: Transformers for Longer Sequences(Zaheer et al., 2020)
- Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention(Katharopoulos et al., 2020)
- Leave No Context Behind: Efficient Infinite Context Transformers with Infini-attention(Munkhdalai et al., 2024)
- Mistral 7B(Jiang et al., 2023)
- Gemini 1.5: Unlocking multimodal understanding across millions of tokens of context(Gemini Team, 2024)
- PyTorch FlexAttention 文档
下一章
这一章一直假设模型知道 token 的相对位置:窗口才能移动,因果方向才有意义,跨 Segment 也要处理位置。第 25 章回到这个基础问题,看看位置编码怎样从绝对位置走向相对距离、旋转和线性偏置。