Lighthouse Attention(11分钟阅读)

TLDR AI 论文

摘要

Lighthouse Attention是一种基于选择的分层注意力机制,通过在前向+反向传播中实现约17倍的速度提升(在512K上下文下),并在98K上下文中实现1.4–1.7倍的端到端加速,从而加速长上下文预训练。该机制使用Llama-3 530M模型在50B token上进行了验证。

Lighthouse Attention是一种基于选择的分层注意力机制,在大上下文下,其前向和反向传播速度比标准注意力模型快17倍。它在密集子序列上使用FlashAttention,保持了效率并与上游改进兼容。通过实现高效的长上下文训练并保留密集模型的能力,Lighthouse Attention在预训练中实现了1.4倍到1.7倍的加速,同时降低了计算成本。
查看原文
查看缓存全文

缓存时间: 2026/05/19 00:20

# Lighthouse Attention 来源:https://nousresearch.com/lighthouse-attention Lighthouse Attention:单块 B200 上按上下文的前向+反向延迟(https://arxiv.org/abs/2605.06554)2605.06554(https://arxiv.org/abs/2605.06554) > ***TL;DR。**一种基于选择的分层注意力机制,在单块 B200 上以 512K 上下文进行相同的前向+反向传播时,**比标准注意力快约 17 倍**,并在 98K 上下文中提供 **1.4–1.7 倍的端到端预训练加速**。Q、K、V 在 L 级金字塔中对称池化;每个头的 $\ell_2$ 范数选择一个小型密集子序列;FlashAttention 作用于 gather 结果——无需自定义稀疏注意力内核、直通估计器或辅助损失。在稀疏阶段之后,一个简短的 standard attention 恢复将检查点转换回密集注意力模型:在相同的 token 预算下,每次恢复的运行均达到或超过从头训练的密集模型。在 530M Llama-3、16k 优化器步骤、50B tokens 上进行了验证,并在上下文并行下使用 32 块 B200 进行 100 万 token 训练。* 长上下文预训练受到注意力二次计算成本的瓶颈。FlashAttention 削减了常数,但壁垒依然存在:你只能在能负担的上下文下进行训练。 我们引入了 ***Lighthouse Attention***,一种基于选择的分层注意力机制,它跨多分辨率金字塔*对称地* 池化 queries、keys 和 values,使用无参数函数对每个金字塔条目评分,并将选择逻辑置于*外部* 注意力内核。前向传播中昂贵的步骤是对一个小型密集子序列执行 FlashAttention。相同的内核用于训练和推理,并且我们原封不动地继承了上游 FlashAttention 的所有改进。 代码位于:github.com/ighoshsubho/lighthouse-attention (https://github.com/ighoshsubho/lighthouse-attention) ## 两个设计决策 该领域的大多数先前工作(NSA、HISA、InfLLM-v2、DSA、MoBA)做出了两个对训练悄然重要的设计决策。 **非对称性。** Queries 保持全分辨率;只有 keys 和 values 被池化。层次结构充当可压缩的可寻址内存,而不是多尺度表示。 **架构纠缠。** 选择位于注意力内核内部。现代张量核心加速的精心优化的密集注意力内核无法复用;每一种稀疏方法都自带其内核。 还有一个特定于训练的问题。*推理* -时稀疏方法至多与其密集主干一样好:稀疏替换仅针对密集前向进行评估。*训练*-时稀疏方法必须经受更严格的测试:训练完成后,**该模型是否仍是一个称职的密集注意力模型?** 如果不是,那么它只是训练了自己近似方法的专家。 我们将此问题视为核心正确性检查。 ## 方法 **对称池化。** Q、K 和 V 均在层次结构的每一层以相同因子池化。第 $\ell$ 层的池化 query 与第 $\ell$ 层的池化 key 处于相同的表示空间。这个选择将密集注意力调用从训练时的 $O(N \cdot S \cdot d)$ 变为 $O(S^2 \cdot d)$。 **无参数评分。** 每个金字塔条目获得两个标量分数:其 query 投影的 $\ell_2$ 范数及其 key 投影的 $\ell_2$ 范数。没有学习的评分头,没有辅助损失,没有 Gumbel-softmax,没有直通估计器。投影被鼓励为*在被选中时有用*,而不是*在评分时表现良好*。膨胀的 softmax 注意力评分器是更强的信号——它看到 QK 交互,而范数评分器看不到——因此我们的结果是基于选择的训练所能提供的下界。 **选择在内核之外。** 一旦决定 Top-K,我们将选中的条目收集成一个连续、因果排序的密集子序列,并对其运行 FlashAttention。训练时昂贵的步骤与密集基线使用的密集注意力内核相同;前向和反向与密集 Transformer 逐位相同。 ## 四个阶段 Lighthouse 注意力层用四个阶段替换标准缩放点积注意力,这些阶段围绕但不修改注意力内核。 Lighthouse 架构流水线 *图 1.* Lighthouse Attention。前向(黑色)将 $H_t$ 投影为 Q、K、V,应用对称金字塔池化,并根据来自分层选择器的索引 $\mathcal{I}$(评分 → Top-K)进行密集 gather、FlashAttention 和确定性 scatter-back 以产生 $O_t$。选择器分支不可微:top-K 返回整数索引,因此没有梯度流过 Score 或 Top-K。 三个小型交互面板使每个阶段具体化。 ### (i) 金字塔池化 将 Q、K、V 对称地平均池化为一个 L 级金字塔,池化因子为 $p$: $$ Q^{(\ell)} = \mathrm{Pool}_{\mu}(Q), \quad K^{(\ell)} = \mathrm{Pool}_{\mu}(K), \quad V^{(\ell)} = \mathrm{Pool}_{\mu}(V), \quad \ell = 0, 1, \ldots, L-1 $$ 第 0 级是整个序列;第 $\ell$ 级有 $N/p^{\ell}$ 个 tokens,每个 token 总结 $p^{\ell}$ 个基础位置。可视化使用 $N = 16$,$L = 3$,$p = 2$(16 个基础 tokens 向上分支到 8 + 4 个池化摘要),因此您可以准确看到粗粒度单元负责哪些基础位置。 ### (ii) Top-K 级联 计算每头 queries 和 keys 在所有层级上的 $\ell_2$ 范数,并联合选择: $$ s^{(QK)}_{\ell,i} = \|Q^{(\ell)}_i\|_2, \qquad s^{(KQ)}_{\ell,i} = \|K^{(\ell)}_i\|_2 $$ $$ \mathcal{I} = \mathrm{TopK}\!\!\left( \{ s^{(QK)}_{\ell,i},\, s^{(KQ)}_{\ell,i} : (\ell, i) \in \mathcal{P} \},\, k \right) $$ 可视化从粗到细遍历级联:在最粗层级进行 top-K,下降到幸存者的子层级,再次进行 top-K,下降,在基础层级保留所有内容。选中的单元以金色环高亮;被拒绝的单元变暗并带红色环。 第 $\ell$ 层的粗粒度条目总结了 $p^{\ell}$ 个连续的基础位置。如果我们丢弃每个被拒绝的粗条目,第 $\ell$ 层的幸存者将不再覆盖基础序列:在那些粗摘要未入选*且*更细的后代也未选中(因为选择从选中的父级继承)的位置上将出现间隙。这些间隙正是迫使后续使用*稀疏感知* 因果掩码的原因。 我们通过将拒绝的粗条目与选中的条目一起保留在缓冲区中来避免这种情况。每个层级 $\ell$ 最多贡献 $p \cdot K$ 个条目(K 来自 top-K,加上一个小的 p 因子用于因果边界记账)。在按基础序列位置对收集的三元组排序后,生成的子序列在*拓扑上* 是因果的且没有空洞:标准的 $S \times S$ 下三角因果掩码正常工作,注意力内核永远不会看到稀疏布局。 ### (iii) 作为黑盒的注意力,然后 scatter-back 将幸存的 (Q, K, V) 三元组收集成一个长度为 $$ S = N / p^{L-1} + (L - 1) \cdot p \cdot K $$ 的连续子序列,对其运行普通 FlashAttention, $$ \tilde{O} = \mathrm{Attn}(\tilde{Q}, \tilde{K}, \tilde{V};\, \tilde{M}) $$ 然后将每个输出条目散射回其代表的 $p^{\ell}$ 个基础位置,偏移量为 $p^{\ell} - 1$(因此位置 $[a, a + p^{\ell} - 1]$ 的粗摘要写入 $[a + p^{\ell} - 1, a + 2p^{\ell} - 2]$:再次是因果边界)。累积在两种内核之一中运行:默认非确定性 fp-原子积累,以及确定性整数-原子积累,后者**慢 1.2–2 倍**。确定性内核仅用于结果复现;fp-原子是默认选项。 实现的大部分是两个新文件加上在上游 torchtitan 基础上约 600 行修改:每个*可能* 需要自定义稀疏内核的步骤都被替换为 `torch.gather` 后跟 `torch.sort`,然后是普通的 FlashAttention。 ## 训练方案 训练后的模型在稀疏训练后必须仍然是称职的密集注意力模型,因此方案分两个阶段: - **阶段 1 (Lighthouse)**。在预算的大部分时间内启用 Lighthouse 选择进行训练。 - **阶段 2 (SDPA-resume)**。关闭选择恢复阶段 1 的检查点;在标准注意力下继续训练一个简短的尾部。相同的优化器状态,相同的数据加载器延续。 如果稀疏训练信号削弱了模型的密集注意力能力,阶段 2 将无法恢复。如果没有,阶段 2 将平滑收敛到一个与从头密集运行竞争的模型。 在三个分割点(总 16k 步中的 10k+6k、11k+5k、12k+4k),每个恢复的 Lighthouse 运行在相同的 16,000 步/~50B token 预算下达到或超过从头密集训练的基线。在每个恢复点,损失跳升 **1.12–1.57 nats**,因为模型首次被要求使用其未训练过的密集注意力,然后在大约 1k–1.5k SDPA 步内恢复。 这是论文的核心主张:稀疏训练不会损害模型在推理时使用完整注意力的能力,且相比于从头密集训练没有额外的 token 成本。 ## 消融实验 消融网格(530M Llama-3,16k 优化器步骤,8×B200 单节点,除非行中注明 CP): | 配置 | 评分器 | LH 步数 | 总步数 | Tokens | B200-小时 ↓ | Tok/s (k) ↑ | 最终损失 ↓ | |---|---|---|---|---|---|---|---| | **SDPA 基线 (ctx = 98K)** | n/a | n/a | 16k | 50.3B | 303.2 | 45.6 | 0.7237 | | *SDPA 可恢复性 (L=3, p=2, k=6144, ctx = 98K)* | | | | | | | | | LH → SDPA (12k+4k) | 膨胀 | 12k | 16k | 50.3B | **214.7** | 74.7 | 0.7102 | | LH → SDPA (11k+5k) | 膨胀 | 11k | 16k | 50.3B | 219.6 | **75.4** | 0.7001 | | LH → SDPA (10k+6k) | 膨胀 | 10k | 16k | 50.3B | 228.0 | 75.0 | **0.6980** | | *超参数消融 (ctx = 98K)* | | | | | | | | | L=3, p=2, k=1536 | 膨胀 | 10k | 16k | 50.3B | 203.9 | 93.9 | **0.6825** | | L=3, p=4, k=1536 | 膨胀 | 10k | 16k | 50.3B | **197.2** | **99.5** | 0.6881 | | L=3, p=8, k=1536 | 膨胀 | 10k | 16k | 50.3B | 206.2 | 92.1 | 0.6828 | | L=4, p=2, k=1536 | 膨胀 | 10k | 16k | 50.3B | 200.2 | 96.4 | 0.6978 | | L=5, p=2, k=1536 | 膨胀 | 10k | 16k | 50.3B | 201.5 | 96.3 | 0.6991 | | L=3, p=2, k=2048 | 膨胀 | 10k | 16k | 50.3B | 208.1 | 90.9 | 0.6880 | | L=3, p=2, k=4096 | 膨胀 | 10k | 16k | 50.3B | 215.7 | 83.5 | 0.6951 | | *CP 训练 (L=3, p=4)* | | | | | | | | | k=1536, ctx = 98K, CP=2, DP=4 | 范数 | 10k | 16k | 100.7B | 208.3 | 91.8 | 0.6903 | | k=2048, ctx = 98K, CP=2, DP=4 | 范数 | 10k | 16k | 100.7B | 210.9 | 89.2 | 0.6928 | | k=4096, ctx = 256K, CP=8, DP=1 | 范数 | 10k | 16k | 1.07T | 1300.3 | 48.9 | **0.6721** | *表 1.* B200-小时 = 挂钟时间 × 8 GPU(LH + SDPA 阶段合计)。Tok/s (k) 报告来自 torchtitan 的 Lighthouse 阶段吞吐量,跨 rank 聚合;SDPA 基线显示无 LH 训练时的吞吐量。 每个恢复的 Lighthouse 运行在相同 token 预算下均击败从头密集训练基线(最终损失 0.6980–0.7102 对比 0.7237),同时节省 75–106 B200-小时:1.40 倍至 1.69 倍的挂钟加速。 阶段 1 吞吐量在整个消融网格中维持在 84–126k tokens/s/GPU,相比之下密集型 SDPA 约为 46k。Lighthouse 在阶段 1 完全收回成本;SDPA 恢复尾端运行与基线相同的内核,匹配其吞吐量。 金字塔超参数宽容性良好。$L \in \{3, 4, 5\}$ 和 $p \in \{2, 4, 8\}$ 均落在彼此约 0.02 nats 内;选择主要是吞吐量/内存覆盖权衡,而非质量敏感边缘。 ## 扩展 前向和反向延迟 vs. 上下文长度 *图 2.* B200 上的单层注意力延迟(bf16, $B = 1$, $H = 8$, $d = 128$, $L = 3$, $p = 4$, 稀疏度 ≈ 1 : 64)。SDPA (cuDNN) 呈 $\Theta(N^2 \cdot d)$ 缩放;Lighthouse (cuDNN) 呈 $\Theta(S^2 \cdot d)$ 缩放,其中 $S \ll N$,因此差距随 $N$ 扩大。Lighthouse 的曲线是收集子序列的代价,包括评分/top-K/scatter 开销,而不仅仅是 FlashAttention 调用。 在短上下文中两条曲线相互接近(金字塔池化+选择的恒定开销占主导);超过交叉点后,SDPA 曲线以二次方增长,而 Lighthouse 在固定 $K$ 下接近线性于 $N$。消融表中 98K 上下文下 1.4–1.7 倍端到端训练加速正是此差距在纳入一个步骤中其他所有内容(FSDP、优化器、数据加载器、scatter-back)后产生的结果。 ## 长上下文 超过约 100K 上下文时,我们的 530M 架构在单块 B200 上无论注意力方法如何都会 OOM,因此对于长上下文场景,我们在上下文并行 (CP) 下运行 Lighthouse。金字塔池化、评分和 top-K 均在分片本地运行;收集的子序列是密集的,因此它可以参与环形注意力而无需要稀疏感知的集合通信。 这正是使 CP 路径可行的地方。Lighthouse 的选择输出是一个连续张量;那些注意力调用期望稀疏索引的方法无法在不针对稀疏布局进行特定工程的情况下表达环旋转。使用 Lighthouse,环旋转一个密集子序列,就能正常工作。 CP 引入了一个小的环形旋转开销(相对于单设备外推约 **10%** 的每 rank 吞吐量损失),并支持**在 32 块 Blackwell GPU(4 节点,CP 度 8)上进行 100 万 token 训练**,而无需更改内部注意力内核。 ## 设置 这些是 **530M Llama-3** 模型,训练了 **16,000 优化器步骤,约 50B tokens**:足够小以快速扫描消融网格,足够大以清晰确立核心正确性声明。对于长上下文检索,我们运行了一个简化的 passkey 测试:单个数字隐藏在合成字母数字填充中,对 10 个数字 tokens 进行单 token argmax 评分。四个 Lighthouse 运行中有三个在该测试中达到或超过从头密集训练基线。 ## 局限性 对称 Q/K/V 池化假定所有 queries 在同一前向传播中共同出现;自回归解码违反了这一点。我们依赖密集 SDPA 恢复将 Lighthouse 权重转换为可用于推理的模型。 收集的子序列成本为 $\Theta(S^2 \cdot d)$:在固定 $K$ 下关于 $N$ 为次二次,但并非严格线性。尚未确定 $K$ 必须随 $N$ 缩放以保持召回率的场景。 ## 开放问题 - **非对称稀疏恢复。** 使用面向推理的非对称稀疏目标(DSA、NSA、HISA、MoBA)替换密集 SDPA 恢复,使转换后的检查点原生可服务。 - **自适应选择预算。** 每层或每头的 $K$ 分配,而不是单一的固定 $K$。 - **超出文本。** 视觉、音频和视频具有天然的多尺度结构,适合金字塔。 - **服务集成。** 用于转换后的推理模型的连续批处理、推测性解码和 KV 缓存管理。 ## 代码 参考实现作为上游 torchtitan 上的单个补丁,外加两个新源文件: > github.com/ighoshsubho/lighthouse-attention (https://github.com/ighoshsubho/lighthouse-attention) 配置按消融轴组织(`topk/`、`pool/`、`levels/`、`scorer/`、`cp/`);该补丁支持三种评分器变体(`norm`、`dilated`、`gla`),可从您的 toml 文件中选择,CP 路径需要 norm 评分器。 *论文:*Long Context Pre-Training with Lighthouse Attention (https://arxiv.org/abs/2605.06554)(arXiv:2605.06554)。

相似文章

使用灯塔注意力的长上下文预训练

Hugging Face Daily Papers

灯塔注意力是一种仅用于训练的、基于层次选择的注意力算法,它降低了因果Transformer长序列训练的计算复杂度,通过恢复阶段后的竞争性最终损失实现更快的预训练。

MiniMax 稀疏注意力

Hugging Face Daily Papers

MiniMax 稀疏注意力 引入了一种分块稀疏注意力机制,针对超长上下文的大语言模型实现了显著的加速。在1M上下文长度下,每个token的注意力计算减少28.4倍,在H800 GPU上预填充阶段实际速度提升14.2倍,解码阶段提升7.6倍。该方法附带了一个开源推理内核以及一个公开发布的多模态模型。