一句话总结:Sparse Attention 少算连接,Linear Attention 改写相似度计算,Infini-attention 把跨 Segment 的旧 K/V 压缩进固定状态;三者解决的并不是同一个问题,也都不会免费保留完整历史。


24.1 Full Attention 的三笔账

24.1.1 连接数

长度为 NN 的 Full Self-Attention,要比较每个 Query 和每个 Key。双向 Attention 有 N2N^2 个连接;因果 Attention 允许的下三角连接是:

1+2+⋯+N=N(N+1)21+2+\cdots+N=\frac{N(N+1)}{2}

不过,“逻辑上只有下三角”不代表一个普通稠密矩阵乘法自动少算一半。是否真的跳过被遮住的区域,取决于所用 Kernel。

NN双向连接数 N2N^2因果有效连接数 N(N+1)/2N(N+1)/2
1,0241,048,576524,800
4,09616,777,2168,390,656
32,7681,073,741,824536,887,296
1,000,0001,000,000,000,000500,000,500,000

序列变长 10 倍,连接数大约变成 100 倍。这是第一个问题:算术工作仍是二次增长。

24.1.2 中间结果的存储

如果真的把一个 Head 的 N×NN\times N score 矩阵以 FP16/BF16 保存下来,它会占用:

NN一个 Head 的稠密 score 矩阵
4,09632 MiB
8,192128 MiB
32,7682 GiB
131,07232 GiB

这些数字只是一个 Head、一个矩阵,不是“整个单层只占这么多”。若真的按所有 Head 保存,还要乘 Head 数;训练还涉及反向传播所需状态。

但第 21 章的 FlashAttention 已经告诉我们:精确 Softmax Attention 不必把完整 score 矩阵写入 HBM。它可以把工作内存降到随序列近似线性增长,同时保持完全相同的数学结果。FlashAttention 解决的是 I/O 和中间存储,不会消除所有 Query–Key 配对的二次算术。

Full Attention 的连接数、显式 score 存储和 KV Cache 是三种不同成本

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 就是把“谁能看谁”画成一张稀疏图。

Sliding Window、Global 与 Random 三种稀疏 Attention 连接模式

常见的三种边是:

  1. Window / Local:每个 Query 只看附近 ww 个 Key。
  2. Global:少数 gg 个位置与全序列双向连接,像 Longformer 中任务指定的全局 token。
  3. Random:每个位置再连接 rr 个随机位置,BigBird 用它和 Window、Global 一起构图。

若每行只有固定数量的这些连接,边数的主项是:

O ⁣(N(w+g+r))O\!\left(N(w+g+r)\right)

只有当 w,g,rw,g,r 不随 NN 增长时,才可以简写成 O(N)O(N)。如果窗口跟着序列一起扩大,复杂度也会跟着变。

24.2.1 Longformer 和 BigBird 的结论边界

Longformer 将滑动窗口与任务指定的全局 Attention 组合起来,原论文主要研究长文档编码、MLM 和 LED 编码器—解码器。BigBird 再加入随机连接,原论文主要也是 BERT 风格的长序列编码任务。

BigBird 论文证明,在它给定的图结构、全局 token、层数和精度等假设下,稀疏 Transformer 仍可成为通用逼近器并保持图灵完备性。这不等于:

  • 任意随机图都保证两个 token 只隔固定 O(1)O(1) 层;
  • 一个有限深度、训练完成的 BigBird 与 Full Attention 表现必然相同;
  • 所有任务的精度损失都“很轻微”。

理论表达能力、图上的传播路径、训练后质量,是三件不同的事。

24.2.2 Mask 不等于省计算

在一个普通 N×NN\times N score 矩阵上填 -inf,可以得到正确的 Sparse Attention 结果,但矩阵已经算出来了。真正节省工作,需要 Block-sparse、FlexAttention 一类能够利用结构的 Kernel。块大小、稀疏率、形状和硬件利用率决定了实际速度;稀疏不够时,调度开销甚至可能吃掉理论收益。

在稠密矩阵上加 Sparse mask 与使用结构化 Sparse kernel 的区别

24.3 Sliding Window:最容易理解,也最容易差一格

下面我们只讲 decoder 的因果窗口。约定 window=3 表示每个位置最多看“自己 + 前两个位置”:

q0 -> k0
q1 -> k0 k1
q2 -> k0 k1 k2
q3 ->    k1 k2 k3
...

长度 8 时,有效连接数不是 8×3=248\times3=24,而是 1+2+6×3=211+2+6\times3=21;左边界少了三条。

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 时,堆叠 LL 层后的最大因果跨度是:

1+L(w−1)1+L(w-1)

有些论文把 WW 定义成向后看的距离,于是会写成大约 LWLW。两种说法可能只差定义,不要把窗口直径、半径和包含当前位置混在一起。

24.3.2 Mistral 7B 的历史例子

Mistral 7B v0.1 技术报告给出的配置是 32 层、Sliding Window W=4096W=4096,并据此写出约 131K token 的理论信息传播跨度。同一张配置表的 context_len 是 8192;131K 不是说 v0.1 检查点已在 131K 上逐 token 验证过精确回忆。

报告还描述了 Rolling Buffer Cache:位置 ii 的 K/V 写入 i mod Wi\bmod W,超过窗口的旧条目被覆盖。Cache 在窗口填满后不再随总生成长度增长;但被覆盖位置的精确 K/V 已不可访问,它们只能通过较新 token 的隐藏状态间接影响后续层。

因果滑动窗口随层数扩展感受野以及 Rolling Buffer KV Cache 的覆盖过程

24.4 Linear Attention:不是把 Softmax 挪个括号

若没有 Softmax,矩阵结合律确实给出:

(QKT)V=Q(KTV)(QK^T)V=Q(K^TV)

左边先产生 N×NN\times N,右边先产生 dk×dvd_k\times d_v。但标准 Attention 是:

softmax⁡(QKT/dk)V\operatorname{softmax}(QK^T/\sqrt{d_k})V

Softmax 对每个 Query 的整行分数做归一化,不能简单穿过括号。因此常说的 Kernelized Linear Attention 定义了一个非负特征映射 ϕ\phi,把相似度改成:

sim⁡(Qi,Kj)=ϕ(Qi)Tϕ(Kj)\operatorname{sim}(Q_i,K_j)=\phi(Q_i)^T\phi(K_j)

对于 causal Attention,维护两个前缀状态:

Si=∑j≤iϕ(Kj)VjT,Zi=∑j≤iϕ(Kj)S_i=\sum_{j\le i}\phi(K_j)V_j^T,\qquad Z_i=\sum_{j\le i}\phi(K_j)

输出为:

Oi=ϕ(Qi)TSiϕ(Qi)TZiO_i=\frac{\phi(Q_i)^TS_i}{\phi(Q_i)^TZ_i}
Causal Linear Attention 用固定前缀状态 S 和 Z 代替 N 乘 N Softmax 矩阵

论文中的一个具体选择是 ϕ(x)=ELU⁡(x)+1\phi(x)=\operatorname{ELU}(x)+1:

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)

若特征维度为 CC、Value 维度为 MM,主项是 O(NCM)O(NCM);流式 Decode 的状态是 CM+CCM+C,相对历史长度固定。它什么时候比 Full Attention 快,还取决于 N,C,MN,C,M、并行度和 Kernel。更重要的是,这段代码精确实现的是这个 Kernel Attention,不是标准 Softmax Attention 的无误差重排。

本章配套测试把流式前缀状态与显式构造的下三角 Kernel Attention 比较,float32/float64 的输出和梯度均一致。


24.5 Infini-attention:固定的是跨 Segment 状态

Infini-attention 把输入切成固定长度的 Segment。每一层、每个 Head 同时做两件事:

  1. 当前 Segment 内执行标准的 causal scaled dot-product Attention,保留局部精细信息;
  2. 用上一个 Segment 留下的压缩状态检索更久远的历史,并在当前 Segment 结束后更新状态。
Infini-attention 在当前 Segment 内做局部 Attention,并通过 M 与 z 在 Segment 之间传递压缩状态

对一个 Head,状态不是简单的 d_model × d_model,而是:

M: [d_key, d_value]
z: [d_key]

从旧状态检索:

Amem=ϕ(Q)Ms−1ϕ(Q)zs−1A_{mem}=\frac{\phi(Q)M_{s-1}}{\phi(Q)z_{s-1}}

在输出当前 Segment 后,用它的 K、V 更新:

Ms=Ms−1+ϕ(K)TV,zs=zs−1+∑tϕ(Kt)M_s=M_{s-1}+\phi(K)^TV,\qquad z_s=z_{s-1}+\sum_t\phi(K_t)

论文还实验了 Delta 更新;这里先保留它最直接的 Linear 版本。局部输出 AdotA_{dot} 与记忆输出通过每 Head 一个可学习标量融合:

A=σ(β)Amem+(1−σ(β))AdotA=\sigma(\beta)A_{mem}+(1-\sigma(\beta))A_{dot}
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 先从 Ms−1,zs−1M_{s-1},z_{s-1} 读取,算完输出后才写入 Ms,zsM_s,z_s。这样记忆分支不会偷看当前 Segment 的未来位置;当前 Segment 的因果关系由局部 Attention 负责。

24.5.1 “无限”到底是什么意思

固定的是跨 Segment 状态:每层每 Head 需要 dkdv+dkd_kd_v+d_k 个数,与已经处理过多少 Segment 无关。若 Segment 长度固定为 SS、总长度为 TT:

  • 局部 Attention 总算术约为 O(TSd)O(TSd),不是总计算 O(1)O(1);
  • 压缩记忆的状态大小相对 TT 是 O(1)O(1);
  • 当前 Segment 的临时工作内存仍取决于 SS 和所用 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 把方法放回正确的坐标轴

方法改了什么对总长度 TT 的主项历史是否精确保留关键条件
Full + FlashAttentionKernel / I/OO(T2d)O(T^2d) 算术在窗口内是不保存完整 score 到 HBM
Sliding Window连接图O(Twd)O(Twd)窗外只有间接影响ww 固定且 Kernel 利用稀疏结构
Longformer / BigBird局部、全局、随机图O(T(w+g+r)d)O(T(w+g+r)d)只保留允许的边原论文结论有任务与图结构边界
Kernel Linear Attention相似度函数与计算顺序O(TCM)O(TCM)压进 S,ZS,Z不是 Softmax Attention 的简单等价式
Infini-attentionSegment 间的循环压缩状态固定 Segment 时随 TT 线性跨 Segment 有压缩损失需要相应训练;局部 Attention 仍存在

这些方法可以在数学上兼容某些组合,但“能组合”不等于任意框架已经提供高效融合 Kernel。FlashAttention、GQA、Block-sparse mask、Rolling Cache 和压缩记忆分别作用在不同层面;每加入一层,都要重新验证正确性、质量、峰值显存与真实吞吐。


24.7 本章总结

  1. Full Attention 有二次配对成本;显式 N2N^2 score 存储和线性 KV Cache 是另外两笔账。
  2. Sparse Attention 只有配合能跳过无效块的 Kernel 才真的省计算;稠密矩阵加 mask 只是得到稀疏结果。
  3. Longformer 与 BigBird 的线性复杂度要求窗口、全局和随机连接数不随序列增长;理论表达力不等于任意任务上的质量保证。
  4. Linear Attention 用特征映射定义新的 Kernel 相似度;它不是把标准 Softmax Attention 换一下括号。
  5. Infini-attention 的 M,zM,z 状态相对总历史固定,但总计算仍随处理的 Segment 数增长,跨 Segment 历史也不是无损保存。
  6. Gemini 1.5 使用 Infini-attention 的说法没有公开证据,不能写成事实。

参考文献

  1. Longformer: The Long-Document Transformer(Beltagy et al., 2020)
  2. Big Bird: Transformers for Longer Sequences(Zaheer et al., 2020)
  3. Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention(Katharopoulos et al., 2020)
  4. Leave No Context Behind: Efficient Infinite Context Transformers with Infini-attention(Munkhdalai et al., 2024)
  5. Mistral 7B(Jiang et al., 2023)
  6. Gemini 1.5: Unlocking multimodal understanding across millions of tokens of context(Gemini Team, 2024)
  7. PyTorch FlexAttention 文档

下一章

这一章一直假设模型知道 token 的相对位置:窗口才能移动,因果方向才有意义,跨 Segment 也要处理位置。第 25 章回到这个基础问题,看看位置编码怎样从绝对位置走向相对距离、旋转和线性偏置。

引用本文 / Cite
Zhang, Wayland (2026). 第 24 章:Sparse 与 Infini-attention. In Transformer 架构:从直觉到实现. https://waylandz.com/llm-transformer-book/%E7%AC%AC24%E7%AB%A0-Sparse%E4%B8%8EInfinite-Attention/
@incollection{zhang2026transformer_24_-Sparse_Infinite-Attention,
  author = {Zhang, Wayland},
  title = {第 24 章:Sparse 与 Infini-attention},
  booktitle = {Transformer 架构:从直觉到实现},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book/%E7%AC%AC24%E7%AB%A0-Sparse%E4%B8%8EInfinite-Attention/}
}