观文听傑

返回

上一篇把 RoPE 模型的名义窗口扩到了训练外,并强调“能放进 32K token”不等于“能用好 32K token”。即便位置外推可靠,标准自注意力还要为长度为 LL 的序列生成 L×LL\times L 个 Query–Key 分数:长度扩大 8 倍,分数矩阵扩大 64 倍。

稀疏注意力(Sparse Attention)不再让每个 Query 读取所有 Key,而是预先规定一张可见图。本文只讲三种紧密相连的边:滑动窗口保留邻近上下文,块稀疏让布局贴合硬件,全局 token 提供远距离中转站。核心问题不是“怎样把矩阵画得更空”,而是:哪些信息路径可以删除,仍不破坏任务真正需要的通信?

01 全注意力的瓶颈究竟在哪里?#

设 Query、Key、Value 的形状分别为:

Q [N,H,L,d] ─┐
K [N,H,L,d] ─┼─► scores [N,H,L,L] ─► softmax ─► output [N,H,L,d]
V [N,H,L,d] ─┘
text

NN 是 batch 大小,HH 是头数,LL 是序列长度,dd 是每头维度。每个头的分数为:

Sij=qikjdS_{ij}=\frac{q_i^\top k_j}{\sqrt d}

计算 QKQK^\top 约需 O(L2d)O(L^2d) 次运算;若显式保存分数或概率,则激活占用为 O(L2)O(L^2)。FlashAttention 能通过分块重算减少显存读写,却没有把任意一个 qikjq_i^\top k_j 从数学上删除;序列足够长时,平方级计算仍在。

| 方法 | 允许的 Q–K 边 | 理论边数 | 主要作用 | | ----------- | -------------: | ---------: | ------------------ | ------------ | -------- | | 全注意力 | 所有 i,ji,j | L2L^2 | 任意两点一步通信 | | 滑动窗口 | ijw | i-j | \le w | 约 L(2w+1)L(2w+1) | 局部依赖 | | 块稀疏 | 选中的块对 | 取决于布局 | 高效执行结构稀疏 | | 局部 + 全局 | 窗口边与全局边 | 约 Lw+LgLw+Lg | 局部计算与远程汇聚 |

02 稀疏注意力其实是一张有向图#

定义布尔邻接矩阵 M{0,1}L×LM\in\{0,1\}^{L\times L}。若 Query ii 允许读取 Key jj,则 Mij=1M_{ij}=1。注意力变为:

Aij=exp(Sij)Mijk:Mik=1exp(Sik),oi=j:Mij=1AijvjA_{ij}=\frac{\exp(S_{ij})M_{ij}}{\sum_{k:M_{ik}=1}\exp(S_{ik})},\qquad o_i=\sum_{j:M_{ij}=1}A_{ij}v_j

实现时通常先把不可见位置加上 -\infty,再做 softmax。每一行必须至少有一个可见 Key,否则全为 -\infty 的 softmax 会产生 NaN

token 是节点;“Query i 能读 Key j”是一条 j ─► i 的信息边。

全注意力:每对节点直接相连
局部窗口:0 ─ 1 ─ 2 ─ 3 ─ 4 ─ 5
局部+全局:0 ════════════════════╗
            └─ 1 ─ 2 ─ 3 ─ 4 ─ 5 ╝  (0 是全局中转站)
text

稀疏化改变的不只是速度,也改变归纳偏置(Inductive Bias):一层能传播到哪里、多层后信息要走几跳、哪个位置承担压缩远程信息的责任,都会变化。

03 滑动窗口怎样把平方边数降成线性?#

双向窗口半径为 ww 时,Mij=1(ijw)M_{ij}=\mathbb{1}(|i-j|\le w)。因果语言模型还必须满足 jij\le i,所以:

Mij=1(0ijw)M_{ij}=\mathbb{1}(0\le i-j\le w)

每个 Query 最多读取 w+1w+1 个 Key,总边数约为 L(w+1)L(w+1)。若 ww 固定,复杂度随 LL 线性增长。

但“一层只能看 ww 个历史 token”不等于“模型永远只能利用 ww 个”。堆叠 KK 层时,理论感受野可扩展到约 KwKw。代价是远程证据要经过多个非线性中间状态,路径更长,也可能被压缩或遗忘。

04 用 8 个 token 手算可见图#

L=8L=8、因果窗口 w=2w=2,token 0 为全局 token。普通 token 能读取合法的全局 Key;全局 Query 仍遵守因果性。

列是 Key j →   0 1 2 3 4 5 6 7
Query i
0               ● · · · · · · ·
1               ● ● · · · · · ·
2               ● ● ● · · · · ·
3               ● ● ● ● · · · ·
4               ● · ● ● ● · · ·
5               ● · · ● ● ● · ·
6               ● · · · ● ● ● ·
7               ● · · · · ● ● ●
text

第 7 行只计算 Key 0、5、6、7,共 4 个分数,而不是 8 个。若全局 token 位于因果序列开头,它不能读取未来,因此只能充当后续位置共享的锚点;双向编码器中的全局 token 才能同时读全序列并被全序列读取。

05 为什么还要从 token 稀疏改成块稀疏?#

逐元素掩码很灵活,却不保证更快。GPU 擅长对连续矩形做矩阵乘法;若先算完整 L×LL\times L 分数再把大部分设为 -\infty,计算量仍是全注意力。

块稀疏注意力(Block-Sparse Attention)把 Query 与 Key 轴切成大小为 Bq,BkB_q,B_k 的块。只有被选中的块对才进入内核:

Key blocks →   K0 K1 K2 K3
Query blocks
Q0              ■  ·  ·  ·
Q1              ■  ■  ·  ·
Q2              ■  ■  ■  ·
Q3              ■  ·  ■  ■

■:执行连续的小矩阵乘法;·:整块跳过
text

块边界会引入粒度误差:块中只要存在少数有效 token,内核可能仍需计算整块,再在块内应用细粒度 mask。块越大,矩阵乘效率通常越好,但多算的无效位置也可能越多。布局应根据长度、窗口和硬件实测,而不是只看理论稀疏率。

06 三种边怎样组合成可用模式?#

常见组合可写成 E=ElocalEglobalEtask\mathcal E=\mathcal E_{local}\cup\mathcal E_{global}\cup\mathcal E_{task}

  • local 保留语法、局部视觉纹理或相邻时间步;
  • global 让少量摘要、问题或特殊 token 与所有位置通信;
  • task 由文档段落、图边、检索结果或成对字段决定。

随机边也能缩短图直径,但可重复性、解释和硬件调度更复杂。设计时先问“任务中的远距离信息通过哪条路径到达”,比先照搬某篇论文的图案可靠。

07 先用稠密掩码验证语义#

下面构造 [L,L] 布尔矩阵。它适合单元测试,不是长序列性能方案。

这里特意先写稠密“真值版本”。生产稀疏内核的输出、梯度与可见边都应和它在小尺寸上对齐。

08 用 PyTorch 2.14 SDPA 做正确性基线#

PyTorch 2.14 当前官方 torch.nn.functional.scaled_dot_product_attention 接受 [N,H,L,d],其布尔 attn_maskTrue 表示允许参与;这与 nn.MultiheadAttentionkey_padding_mask 语义相反。

import torch.nn.functional as F

def dense_reference(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor):
    # q/k/v: [N,H,L,d]
    length = q.size(-2)
    mask = causal_local_global_mask(length, 128, device=q.device)
    return F.scaled_dot_product_attention(
        q, k, v,
        attn_mask=mask[None, None, :, :],  # [1,1,L,L] 广播
        dropout_p=0.0,
    )                                     # [N,H,L,d]
python

应在一个可测试的 mask 中明确合并因果与稀疏条件。模块训练时若启用 dropout,还要显式传 dropout_p=self.p if self.training else 0.0,因为 SDPA 会按传入值执行 dropout。

09 用 FlexAttention 表达真正的块布局#

PyTorch 2.14 当前官方 FlexAttention API 中,mask_mod 接收 batch、head、Query 索引和 Key/Value 索引。create_block_mask 把 token 条件压成 BlockMask,让内核跳过完整不可见块。

BlockMask 描述“哪些块可能含有效元素”,块内仍由 mask_mod 保证精确语义。固定长度与布局时应复用 mask,避免每个 step 重建;变长 batch 要把 padding 边界并入条件,或按长度分桶。

10 怎样证明稀疏实现没有算错?#

最短验证路径是让 L=8,d=4L=8,d=4

  1. 逐行打印稠密 mask,人工核对因果方向、窗口端点与全局边;
  2. 用相同 Q/K/V 比较稠密基线和稀疏内核的前向输出;
  3. 分别反传同一个标量,比较 Q/K/V 梯度;
  4. 把被屏蔽 Key 的 Value 改成极大值,确认对应 Query 输出不变;
  5. 最后才测长序列的峰值显存、tokens/s 与任务质量。
mask = causal_local_global_mask(8, 2)
out1 = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
v2 = v.clone()
assert not mask[7, 3]
v2[..., 3, :] += 10_000
out2 = F.scaled_dot_product_attention(q, k, v2, attn_mask=mask)
torch.testing.assert_close(out1[..., 7, :], out2[..., 7, :])
python

11 训练与增量推理的数据流有什么不同?#

训练时通常一次输入完整 [N,H,L,d]。带 KV Cache 的解码步只有 Q_LEN=1,而 KV_LEN=P+1

新 Query [N,H,1,d] ─────────┐
缓存+新 Key [N,H,P+1,d] ───┼─► 最近 w 个 Key + 合法全局 Key
缓存+新 Value [N,H,P+1,d] ─┘
text

窗口注意力并不自动让 KV Cache 有界。若还允许读取开头的全局 token,需要保留“全局槽 + 最近 ww 个槽”,并维护真实全局位置供 RoPE 使用。截断缓存后把位置重新编号为 0,会破坏前文建立的位置契约。

12 常见错误与最短调试路径#

症状常见原因最短检查
loss 异常低因果不等号写反,读到未来所有允许边是否满足 jij\le i
输出 NaN某 Query 没有合法 Key检查 mask.any(-1)
显存没下降构造了完整分数或稠密 maskprofiler 中找 [L,L] 分配
稀疏内核更慢序列短、块碎或反复建 mask分离编译、建 mask 与稳态计时
长依赖骤降窗口小且没有远程路径画多层可达图,按证据距离分桶
推理训练不一致cache、位置或全局边不同同前缀逐 token 对齐 logits

性能比较必须固定 dtype、batch、头宽、序列长度和反向设置;先 warm-up,再同步设备计时。只报告“稀疏率 90%”不能说明端到端更快。

13 失败场景与相近方法#

  • 精确复制远处字符串、代码符号解析或跨文档引用时,局部路径可能太长;
  • 全局 token 太少会形成信息瓶颈,太多又把成本拉回 O(Lg)O(Lg)
  • 固定窗口与语义边界不一致,可能在段落交界处删掉关键边;
  • 不规则稀疏在通用硬件上利用率低,理论 FLOPs 减少不等于延迟下降。

还要区分:FlashAttention 精确计算全注意力,主要优化 IO;线性注意力通过核分解或状态递推改变计算顺序,不一定定义稀疏图;检索增强先从外部语料选内容;KV Cache 复用历史 K/V。这些方法可以组合,却不回答同一个问题。

14 今天真正需要记住什么?#

  1. 稀疏注意力先定义信息可达性,再谈加速;mask 是模型结构的一部分。
  2. 滑动窗口把边数从 L2L^2 降到约 LwLw,全局 token 用 LgLg 条边补充远程中转。
  3. 逐元素 mask 只验证语义;真正省计算需要能跳过整块的内核。
  4. 正确性用小尺寸稠密真值、梯度和隔离测试证明,效率用稳态端到端基准证明。

15 思考题与小练习#

  1. L=10,w=2L=10,w=2 的因果窗口,列出第 0、1、5、9 个 Query 的可见 Key,再求总边数。边界处为何少于 L(w+1)L(w+1)
  2. 两层窗口半径为 1 的注意力中,位置 5 最早能间接接收位置几的信息?加入全局 token 后路径如何改变?
  3. 实现稠密 reference 和块稀疏版本,测 L{512,2048,8192}L\in\{512,2048,8192\} 的显存与耗时,找出开始获益的交叉点。

相关工作#

  1. Child et al., Generating Long Sequences with Sparse Transformers,系统探索固定与分步稀疏模式。
  2. Beltagy et al., Longformer,组合局部窗口与任务相关全局注意力。
  3. Zaheer et al., Big Bird,结合局部、随机与全局边并分析表达能力。
  4. Dao et al., FlashAttention,展示精确全注意力的 IO 优化边界。
  5. FlexAttention 团队,FlexAttention,介绍可编程 mask 与块稀疏执行。

16 下一篇预告#

到这里,Transformer 已能在可控成本下读取长前缀,但模型为什么会学会生成下一个 token 还没有被完整展开。下一篇将从文本切片开始,追踪输入与标签怎样错开一位,推导因果语言模型的 next-token cross-entropy,并解释 padding、文档边界和 loss mask 如何决定模型究竟在学什么。

长序列为何不必让每个 token 看见全部历史?滑动窗口、块稀疏与全局 token
https://zwjcode.cn/blog/sparse-attention-sliding-window-block-global-token
作者
发布于 2026年9月9日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。