Back to blog

SeerAttention: 自蒸馏注意力门控实现高效长上下文预填充

通过可学习的 AttnGate 门控机制和自蒸馏训练,直接从 LLM 学习块级注意力稀疏性,128K 序列下实现 7.3× 内核加速。

SeerAttention: 自蒸馏注意力门控实现高效长上下文预填充

一、论文概述

项目内容
标题SeerAttention: Learning Intrinsic Sparse Attention in Your LLMs
作者Yizhao Gao, Zhichen Zeng, Dayou Du, Shijie Cao, Peiyuan Zhou, Jiaxing Qi, Junjie Lai, Hayden Kwok-Hay So, Ting Cao, Fan Yang
机构港大、清华大学等
论文https://arxiv.org/abs/2410.13276
代码https://github.com/Infini-AI-Lab/SeerAttention
发布2024-10-17

二、核心思想

问题定义

注意力机制是现代大语言模型(LLM)的基石,但其二次复杂度 O(n2)O(n^2) 在长上下文场景下成为效率瓶颈。现有稀疏注意力方法主要依赖预定义模式或启发式方法在注意力头级别进行稀疏化,难以动态适应不同的上下文。

核心问题:

  • 注意力稀疏性是内在的、动态的,随输入和注意力头变化
  • 现有方法(如 MoA、MInference)使用静态或启发式模式,缺乏通用性
  • 需要一种学习驱动的方法直接从 LLM 本身学习稀疏性

解决方案概述

SeerAttention 是一种简单而有效的注意力机制,直接从 LLM 本身学习块级注意力稀疏性。

核心创新:

  1. AttnGate(注意力门控):受 MoE 门控机制启发,通过可学习的门控选择性激活注意力图中的重要块
  2. 自蒸馏训练:使用 2D-MaxPooled 注意力图作为真值,轻量级蒸馏 AttnGate
  3. 块稀疏 FlashAttention 内核:高效实现块级稀疏注意力

关键公式:

score=softmax((WqPq(Q))⋅(WkPk(K))Td)score = \text{softmax}\left(\frac{(W_q P_q(Q)) \cdot (W_k P_k(K))^T}{\sqrt{d}}\right)

其中 PqP_q 和 PkP_k 是池化操作,将 Q 和 K 沿序列维度下采样。

关键优势:

  • 学习驱动:直接从 LLM 学习稀疏性,无需预定义模式
  • 动态适应:不同输入和头自动调整稀疏模式
  • 高效训练:仅需 40 A100 小时完成蒸馏
  • 显著加速:128K 序列长度下实现 7.3× 内核加速

三、技术架构

整体框架

SeerAttention 架构

SeerAttention 工作流程:

  1. AttnGate 计算:

    • 池化 Q 和 K 沿序列维度(非重叠块)
    • 通过可学习线性层处理
    • 矩阵乘法生成门控分数
  2. 自蒸馏训练:

    • 使用 2D-MaxPooled 全注意力图作为真值
    • KL 散度损失函数蒸馏 AttnGate
  3. 推理:

    • 使用门控分数预测块级稀疏性
    • TopK 或阈值方法选择活跃块
    • 块稀疏 FlashAttention 计算

AttnGate 设计

核心组件:

组件说明输出大小
池化操作Avg/Max/Min 池化 Q 和 K[seq/B, d]
线性层WqW_q 和 WkW_k 变换[seq/B, d]
矩阵乘法生成块级分数[seq/B, seq/B]
Softmax归一化分数[seq/B, seq/B]

块大小:B=64B = 64,输出大小为原始注意力图的 14096\frac{1}{4096}

池化方法选择

池化方法对比

最优配置:

  • Q:AvgPooling
  • K:Max + Min + AvgPooling(拼接)

原因:K 张量通常包含更多异常值,Max 和 Min 池化能更好地提取特征。

Block-level RoPE

RoPE 设计对比

问题:直接使用原始 RoPE 编码的 Q 和 K 会因池化操作丢失相对位置信息。

解决方案:

  • 使用 RoPE 编码前的 Q 和 K 作为 AttnGate 输入
  • 在 AttnGate 中添加块级 RoPE
  • 使用缩减的 θ′=θ/B\theta' = \theta / B

优势:

  • 有效学习块级位置信息
  • 可外推到更长上下文长度
  • 不会过拟合训练数据长度

RoPE 设计效果

自蒸馏训练

获取真值:

  • 使用 2D-MaxPooled 全注意力图作为真值
  • 语义上:只有当块内所有注意力分数都小时,MaxPooled 结果才小
  • 定制 FlashAttention 内核直接输出 MaxPooled 结果

损失函数:

gt=MaxPool2D(softmax(QKTd))gt = \text{MaxPool2D}\left(\text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)\right)

score=AttnGate(Q,K)score = \text{AttnGate}(Q, K)

loss=DKL(gt∥score)loss = D_{KL}(gt \| score)

训练配置:

  • 数据集:RedPajama,分块为 64K
  • 学习率:1e-3,余弦衰减
  • 批大小:16
  • 训练步数:500 步
  • 硬件:A100 GPU
  • 训练时间:40 A100 小时

块稀疏 FlashAttention 内核

设计:

  • 与 FlashAttention 的 tiling 计算方案无缝集成
  • 仅计算门控分数指示的活跃块
  • 减少 I/O 和计算开销

四、核心公式

AttnGate 计算

score=softmax((WqPq(Q))⋅(WkPk(K))Td)score = \text{softmax}\left(\frac{(W_q P_q(Q)) \cdot (W_k P_k(K))^T}{\sqrt{d}}\right)

2D-MaxPool 真值

gt=MaxPool2D(softmax(QKTd))gt = \text{MaxPool2D}\left(\text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)\right)

KL 散度损失

loss=DKL(gt∥score)loss = D_{KL}(gt \| score)

块级 RoPE

θ′=θ/B\theta' = \theta / B

其中 θ\theta 是原始 RoPE 的 theta,BB 是块大小。

五、实验结果

实验设置

模型:Llama-3.1-8B-Instruct

评估基准:

  • 困惑度:PG19
  • 长上下文:LongBench、RULER
  • 短上下文:MMLU、HellaSwag、ARC-challenge、GSM8K

基线方法:

  • MoA:离线搜索静态稀疏模式
  • MInference:启发式动态稀疏索引
  • DuoAttention:区分流式头和密集头

困惑度结果

困惑度对比

PG19 测试结果:

  • SeerAttention 在不同稀疏度下提供更好的权衡
  • 单个训练的 AttnGate 可在测试时调整 TopK/阈值
  • MoA 在 128K 评估时 OOM

LongBench 结果

方法0-4k4-8k8k+平均平均稀疏度
Full Attention55.3253.9852.954.070.0
MInference55.2353.7852.1853.730.31
MoA50.7449.8451.8950.820.35
DuoAttention53.7752.1751.2752.400.5*
SeerAttention55.4354.4952.6954.200.50

*50% 流式头,实际稀疏度 <50%

关键发现:

  • SeerAttention 在 0-4k 和 4-8k 测试中甚至超过密集基线
  • 最高平均分数(54.20)和最高平均稀疏度(0.50)

RULER 结果

方法4k8k16k32k64k128k平均加速比
Full Attention95.5392.3792.0187.6384.3976.2688.011.00
MInference95.5392.6491.3785.7183.2467.0285.920.83
DuoAttention95.6492.0890.7184.7583.2475.3286.961.09
SeerAttention95.5392.7192.0288.4983.4873.3787.601.41

关键发现:

  • SeerAttention 在大多数测试(8k-64k)中达到最佳精度
  • 平均精度仅比密集基线低 0.41%
  • 最高平均端到端加速:1.41×

短上下文结果

方法MMLUHellaSwagARC-cGSM-8K
Full Attention68.180.160.775.7
SeerAttention67.979.860.275.6
平均稀疏度3.4%50.4%26%52.1%

关键发现:

  • 精度损失可忽略(如 GSM-8K 仅 0.1%)
  • 短上下文下稀疏注意力对延迟提升有限

效率评估

内核级加速

内核加速

AttnGate 开销:

  • 32K 序列、50% 稀疏度:仅增加 1% 延迟
  • 128K 序列:相对开销几乎消失

块稀疏内核加速:

  • 128K 序列、90% 稀疏度:7.3× 加速(相比 FlashAttention-2)
  • 加速与稀疏度呈线性关系

与其他方法对比

内核加速对比

SeerAttention vs MInference:

  • MInference 使用 “Vertical-slash” 模式
  • SeerAttention 将稀疏性转化为加速更有效

SeerAttention vs MoA:

  • MoA 使用 “A-shape” 块模式
  • SeerAttention 在相同稀疏度下加速更高

AttnGate 可视化

AttnGate 输出可视化

观察:

  • AttnGate 学习到有意义的稀疏模式
  • 不同注意力头显示不同的稀疏模式
  • 长上下文下稀疏性更显著

六、核心创新总结

创新点说明优势
AttnGate可学习的块级门控机制直接从 LLM 学习稀疏性
自蒸馏训练2D-MaxPool 注意力图作为真值轻量级,仅 40 A100 小时
Block-level RoPE块级相对位置编码可外推到更长上下文
池化组合Q: Avg, K: Max+Min+Avg最优困惑度表现
块稀疏内核与 FlashAttention 集成高效 GPU 实现

七、技术影响

对稀疏注意力的改进

  • 学习驱动:取代预定义模式和启发式方法
  • 动态适应:不同输入和头自动调整
  • 高效训练:仅需 500 步和 40 A100 小时
  • 显著加速:128K 下 7.3× 内核加速

与现有方法对比

方法稀疏模式训练成本精度加速
MoA静态搜索高中中
MInference启发式动态无中中
DuoAttention流式+密集头中中中
SeerAttention学习驱动低(40h)高高

实际应用价值

  • 长上下文预填充:显著减少预填充延迟
  • 灵活权衡:测试时可调整稀疏度
  • 即插即用:应用于预训练 LLM
  • 低训练成本:快速蒸馏收敛

八、局限性

  1. 仅预填充阶段:当前 AttnGate 仅应用于预填充阶段
  2. 固定块大小:块大小 B=64 固定,未探索其他值
  3. 均匀稀疏度:RULER 评估中使用均匀阈值
  4. 模型规模:主要在 8B 模型上验证
  5. 解码阶段:未应用于解码阶段(提及为未来工作)

九、相关工作

稀疏注意力

  • MoA:离线搜索静态稀疏模式
  • MInference:启发式动态稀疏索引
  • DuoAttention:流式头 + 密集头
  • FlashAttention:高效密集注意力实现

长上下文优化

  • 提示压缩:Jiang et al., 2023; Mu et al., 2024
  • KV 缓存压缩:共享、驱逐、量化
  • 稀疏解码:Yang et al., 2024; Chen et al., 2024

混合专家

  • MoE 门控:Shazeer et al., 2017; Fedus et al., 2022
  • 块级稀疏:与 FlashAttention tiling 集成

十、参考资源

论文与代码

相关工作

  • FlashAttention: Dao et al., 2022; Dao, 2023
  • MoA: Fu et al., 2024
  • MInference: Jiang et al., 2024
  • DuoAttention: Xiao et al., 2024

基准数据集

  • PG19: Rae et al., 2019
  • LongBench: Bai et al., 2023
  • RULER: Hsieh et al., 2024
  • RedPajama: Computer, 2023