Back to blog

HGA: Hierarchical Global Attention - Drop-In Exact-Token Routing for Pretrained Long-Context Transformers

作为预训练长上下文 Transformer 的即插即用分层路由算法,保留原始注意力投影,仅通过查询-键内容路由选择历史 token 进行精确 token 级注意力计算

HGA: Hierarchical Global Attention — Drop-In Exact-Token Routing for Pretrained Long-Context Transformers

一、论文概述

项目内容
标题Hierarchical Global Attention: Drop-In Exact-Token Routing for Pretrained Long-Context Transformers
作者Frank Woernle, Vladimir Fedosov, Artemiy Grinenko
机构—
论文arXiv:2606.30709
代码github.com/vfedosov77/HierarchicalGlobalAttention
发布2026-06-29

核心贡献:

  1. 提出 HGA(Hierarchical Global Attention)——一种无需重新训练即可替代预训练长上下文 Transformer 中密集因果注意力的即插即用分层路由算法
  2. 设计两级层次化路由(chunk-level → group-level),使用预训练的 key 空间本身作为路由空间,无需额外可训练参数
  3. 实现精确 token 级注意力输出:路由后打开真实 token 的 K/V,在 softmax 上精确计算,无 summary output tokens、summary gates 或标度校准参数
  4. 在 Qwen3-30B-A3B-Instruct-2507-FP8 上实现 32K-token 上下文零样本推理(单张 RTX 5090, 32GB VRAM),无需微调
  5. Needle-in-a-Haystack 检索测试在 64K-token 上下文下达到 100% 通过率(3/3 depth),仅 1.9% 稀疏度
  6. 在 40M SmallLM 上,直接权重复制仅产生 +0.018 nat loss gap,Triton-fused 实现训练速度提升 2.72×,prefill 速度提升 2.43×

二、核心思想

问题定义

预训练长上下文 LLM 的实际限制不仅在于 O(n2)O(n^2) 注意力计算,更在于必须驻留在加速器内存中的密集 K/V cache。对于量化大模型尤其明显:Qwen3-30B-A3B-Instruct-2507-FP8 的 FP8 模型权重已占满 RTX 5090 的 32GB VRAM,存储所有历史 token 的 K/V 没有剩余空间。

解决方案概述

HGA 作为系统级补丁,不修改也不替换 checkpoint 的注意力投影。它保持原始的 Q/K/V/O 投影和注意力 norm,仅改变每个处理块获取哪些历史 keys 和 values。GPU 仅持有:预训练模型、紧凑的 chunk-summary table、当前 chunk、少量 always-hot chunks、以及路由工作集。完整的历史 token K/V 存储在主机 RAM 中,每步仅将选定的 chunks 拉取到计算设备。

与现有方法的对比

方法路由策略输出类型可训练参数即插即用RAM-KV Cache
Sparse Transformer固定模式精确 token×××
Longformer固定窗口+global精确 token×××
ReformerLSH hashingSummary tokens×××
Routing Transformer在线聚类Summary tokens×××
MInference固定稀疏模式精确 token×××
vLLMPagedAttention精确 token××GPU-only
HGA层次化 chunk→group精确 token×✓✓ (RAM)

三、技术架构

整体框架

HGA 将序列划分为固定长度 C=64C=64 的 chunks,每个 chunk 进一步细分为大小为 gsg_s 的 groups。关闭的 chunks 维护三类数据:

  1. Chunk key summary:第一级路由的候选集
  2. Group key summaries:第二级路由(当启用 group routing 时)
  3. 原始 token-level K/V:路由后的精确注意力输出

只有最后一项用于最终的注意力输出。

核心公式

标准因果注意力(Background)

Attn(Q,K,V)i=∑j≤isoftmaxj ⁣(qikj⊤dh)vj(1)\mathrm{Attn}(Q,K,V)_i = \sum_{j \leq i} \mathrm{softmax}_j\!\left(\frac{q_i k_j^\top}{\sqrt{d_h}}\right) v_j \tag{1}

精确 token 级感受野随 ii 增长,总查询-键比较次数为 O(n2)O(n^2),K/V cache 与所有先前 token 成正比。

Token 级计算成本

Ttoken-attn=O ⁣(n⋅(C+(F+L)C+Broute))=O(n)(2)T_{\text{token-attn}} = O\!\left(n \cdot \bigl(C + (F+L)C + B_{\text{route}}\bigr)\right) = O(n) \tag{2}

其中 FF 和 LL 分别是 first 和 recent chunks 数量,BrouteB_{\text{route}} 是固定路由 token 预算(如 KggsK_g g_s 或 KcCK_c C)。

Mixed-RoPE Routing Summaries

Chunk summary 必须与 RoPE-rotated queries 可比。Averaging already-rotated keys 适合快速变化的高频对,但会放大 token-level phase noise;Averaging raw keys 然后在代表位置旋转更适合缓慢变化的低频对。

HGA 因此采用混合策略。对每个 key dimension pair:

  • 高频 RoPE pairs:逐 token 旋转后平均
  • 低频 RoPE pairs:在 raw-key space 中平均,然后在 chunk/group 中间位置旋转
summaryk={1∣G∣∑j∈GRp(high)(kj),high-freq pairsRpmid(low) ⁣(1∣G∣∑j∈Gkj),low-freq pairs\text{summary}_k = \begin{cases} \frac{1}{|G|}\sum_{j \in G} R_{p}^{(\text{high})}(k_j), & \text{high-freq pairs} \\ R_{p_{\text{mid}}}^{(\text{low})}\!\left(\frac{1}{|G|}\sum_{j \in G} k_j\right), & \text{low-freq pairs} \end{cases}

这些 summaries 是已投影 keys 的和或均值,没有单独的 summary linear layers。它们仅是路由 keys(routing keys only),在当前输出 softmax 中永不用作 value vectors。

路由策略(Routing Policy)

确定性可见性(Deterministic Visibility):

  • 当前 chunk 始终可见(因果掩码)
  • 可配置数量的首个 chunks 始终可见(attention sinks,如前 2 chunks = 128 sink tokens)
  • 可配置数量的最近 closed chunks 始终可见(local context,如后 8 chunks = 512 recent tokens)
  • 如果紧邻的前一个 chunk 不在 recent window 中,强制包含

基于内容的 middle routing(Content-based Middle Routing):

  • 中间上下文(排除 always-visible first/last windows)竞争路由预算
  • 查询对 chunk key summaries 打分:scorem=q⊤cm\text{score}_m = q^\top c_m
  • Chunk-group 变体:对选定 chunks 内的 group summaries 打分,打开 top-KgK_g groups 到精确 token-level K/V
  • Qwen3 exact-token 变体:选定的 chunks 直接按 KV head 获取 token K/V
Mselected=topKc({q⊤cm}m),Gselected=topKg({q⊤gm,g}m∈Mselected,g)\mathcal{M}_{\text{selected}} = \mathrm{topK}_c(\{q^\top c_m\}_m), \quad \mathcal{G}_{\text{selected}} = \mathrm{topK}_g(\{q^\top g_{m,g}\}_{m \in \mathcal{M}_{\text{selected}}, g})

该方法不是纯固定稀疏模式——它有意结合 fixed sink/local windows 与 content-routed retrieval over the middle context。类似于 MInference 的 A-shape 思想,但 HGA 添加了基于预训练 key 空间的 top-k recall 路径。

精确 Token 输出注意力

路由后,注意力模块仅连接来自以下部分的真实 token-level keys 和 values:

  1. 当前 chunk 中因果可见的部分
  2. always-visible first 和 recent chunks
  3. 打开的 groups 或选定的 routed middle chunks

Softmax 在上述选定的 token 集合上精确计算。无 summary output tokens、无 summary gates、无 per-head scalar calibration parameters。

检查点兼容性

模块与密集注意力参数兼容。在 40M 模型中,dense weights 直接复制到 routed model。在 Qwen3 FP8 实现中,routed attention wrapper 通过引用保持原始量化投影和 norms。唯一改变的操作是每个 query 允许关注的 keys 和 values 集合。

三层 K/V 存储(Tiered K/V Store, Sec. 4.3)

KVRouter 实现将 routing 与 storage 分离。Store 暴露四个操作:append closed chunk、read chunk summaries for routing、fetch group summaries for second routing level、gather exact token K/V for selected chunks/groups。

RAM-backed 实现(RamKVCacheStore)按 temperature 分区数据:

层级存储位置内容说明
HotGPU VRAMChunk summaries + always-visible first/recent chunks始终驻留 compute device,训练中保留梯度
WarmGPU VRAM (LRU)Bounded shard cache,最近路由的 token chunks自动收缩以留出 configurable VRAM headroom
Cold主机 RAM所有剩余 token K/V 和 group summaries仅在被 router 选中时传输到 compute device

该分区意味着 VRAM 消耗由模型权重、hot windows、chunk summaries 和当前 routed working set 主导。cold token record 随上下文长度增长,但对 GPU 内存预算无压力。

实现细节

Qwen3-30B-A3B FP8 Wrapper:

  • 用 QwenRoutedAttention 替换 Qwen3-30B-A3B 的所有注意力模块
  • 30.5B 参数 MoE,3.3B activated parameters,48 layers
  • GQA: 32 query heads / 4 KV heads,FP8 量化,原生 262K context
  • 不修改模型权重

40M SmallLM:

  • 8 decoder layers, hidden size 384, 6 query heads, 2 KV heads, FFN 2048
  • RoPE + GQA
  • C=64C=64, gs=16g_s=16, Kc=20K_c=20, Kg=32K_g=32
  • Triton-fused CUDA + torch.compile

四、核心创新

创新点说明理论/实验依据
层次化两路路由Chunk → Group 两级筛选,减少路由搜索空间Sec. 3.1-3.3
Mixed-RoPE summaries高频/低频 RoPE pairs 分别处理,避免相位噪声Eq.(2) 前后,Sec. 3.2
精确 token 输出路由后打开真实 token K/V,非 summary 近似Sec. 3.4
零参数兼容直接权重复制,无需校准或微调Tab. 6: +0.018 nat gap
RAM-backed KV store完整历史 K/V 存于主机 RAM,按需拉取到 VRAMSec. 4.3
三层缓存管理Hot/Warm/Cold 三级,LRU 自动收缩Sec. 4.1

五、代码实现分析

仓库: github.com/vfedosov77/HierarchicalGlobalAttention

实现架构:

  1. QwenRoutedAttention:大型模型的注意力模块替换器,在 exact mode 下运行
  2. ChunkSummaryManager:维护 chunk-level 和 group-level key summaries
  3. TieredKVStore:RAM-backed 三层 K/V 存储 + LRU 缓存管理
  4. Triton-fused Router:CUDA 融合的路由计算,减少 host-device 传输
  5. Correctness Tests:4 项数值验证(router vs SDPA, prefill vs SDPA, cache vs no-cache, decode vs full recompute)

关键实现特性:

  • 路由后注意力模块与 dense causal SDPA 的数值差异 < 10−610^{-6}
  • Vectorized prefill 与 causal SDPA 差异 4.7×10−74.7 \times 10^{-7}
  • HA cache vs no-cache 差异 1.5×10−51.5 \times 10^{-5}
  • HA decode vs full recompute 差异 2.8×10−52.8 \times 10^{-5}

硬件要求:

  • 训练:NVIDIA RTX A4000(40M SmallLM benchmark)
  • 推理演示:NVIDIA RTX 5090, 32GB VRAM(Qwen3-30B-A3B)

六、实验结果

实验设置

配置项值
大模型Qwen3-30B-A3B-Instruct-2507-FP8
小模型40M SmallLM (8 layers, d=384, 6 Q-heads, 2 KV-heads)
硬件RTX 5090 (32GB), RTX A4000
Chunk size64 tokens
Sink tokens128 (前 2 chunks)
Recent tokens512 (后 8 chunks)
Routed chunks16 (middle context)
Prefill block64 tokens
VRAM cacheBounded LRU, auto-shrink

Loss Gap vs Sparsity and Context Length (Tab. 2)

ContextSparsity 3.13%Sparsity 6.25%Sparsity 12.5%Sparsity 25%
4,096N/AN/AN/A2.456 / 2.460 (Δ<0.01\Delta < 0.01)
8,192N/A2.430 / 2.441 (Δ≈0.01\Delta \approx 0.01)2.430 / 2.437 (Δ<0.01\Delta < 0.01)N/A
16,3842.309 / 2.324 (Δ<0.02\Delta < 0.02)2.309 / 2.322 (Δ<0.02\Delta < 0.02)2.309 / 2.317 (Δ<0.01\Delta < 0.01)N/A
32,7682.204 / 2.227 (Δ>0.02\Delta > 0.02)2.204 / 2.221 (Δ<0.02\Delta < 0.02)2.204 / 2.214 (Δ<0.01\Delta < 0.01)N/A
65,5362.243 / 2.258 (Δ<0.02\Delta < 0.02)2.243 / 2.253 (Δ≈0.01\Delta \approx 0.01)N/AN/A

关键发现:

  • 12.5% 稀疏度下,32K context loss gap < 0.01 nats
  • 64K context 仅需 6.25% 预算即可保持 gap ≈ 0.01 nats
  • 3.13% 预算在 32K 是唯一超过 0.02 nats 的配置,表明需要最小路由预算

Summary Validation (Tab. 3)

ContextSparsityLoss GapStatus
4K25%< 0.01Stable
16K12.5%< 0.01Stable
32K12.5%< 0.01Stable
64K6.25%≈ 0.01Stable

Needle-in-a-Haystack at 64K Tokens (Tab. 4)

Qwen3-30B-A3B-Instruct-2507-FP8 + HGA group-level routing, 零微调。

Needle depthResultTTFT (s)
25%HIT ✓948
50%HIT ✓1053
75%HIT ✓1332
Overall3/3 = 100%–

Copy-only Routing Accuracy (Tab. 6)

40M SmallLM @ 8192 tokens, 无 HGA 微调, use_summaries=False。

ModelLoss (nats)Perplexity
Dense causal SDPA3.7051640.657
Chunk-routed, same weights3.7234441.406
Difference+0.01828+1.8%

Fine-tuning Comparison (Tab. 5)

Qwen3-0.6B fine-tuning, seq 4096, 100 steps, novel text validation.

Metric(a) Routed, full VRAM(b) Dense(c) Routed, RAM cache
Initial loss (stock)3.530 (ppl 34.12)3.530 (ppl 34.12)3.530 (ppl 34.12)
Deploy loss3.196 (ppl 24.44)3.177 (ppl 23.97)3.201 (ppl 24.55)
Fine-tune gain−0.334−0.330−0.329
Routing cost (routed−dense)+0.015—+0.017
KV attended / dense679/2048 (66.9% saved)—679/2048 (66.9% saved)
Speed~917 tok/s~570 tok/s~393 tok/s

关键发现:

  • 微调后 routing cost 极小:0.015–0.017 nats
  • 路由注意力训练吞吐 ~917 tok/s vs 密集 ~570 tok/s,提升 1.6×
  • 仅 attending 33.1% token pairs
  • “Same weights, other attention” 行显示 routed checkpoint 用 dense 评估 (3.182) 几乎等同于 dense checkpoint (3.177),确认稀疏路由微调不损害底层模型质量
  • RAM-cached 变体 (c) 在当前实现阶段较慢,但证明了分层存储路径功能正确

与长上下文位置编码的交互(Sec. 5.5)

一个意外观察是:hierarchical routing 与 dense attention 之间的剩余差异不随上下文长度快速增加。在所有评估的上下文(4K 到 64K tokens)中,validation loss 保持在 dense attention 约 0.01–0.02 nats 范围内。

这表明 routing algorithm 本身引入的近似误差很小。剩余 gap 更可能的解释是 sparse routing 与 long-context positional encoding 之间的交互:

  1. 一些最近的长上下文 LLM 首先以较短的有效上下文长度预训练,然后通过 YaRN 等方法扩展上下文窗口
  2. 在 sparse routing 下,只有部分 transformer layers 观察到 distant tokens
  3. 因此,long-context positional extrapolation 引入的任何不准确可能在 sparse routing 下比 dense attention 更加明显

为验证此假设,论文评估了 RoPE position index wrapped modulo 64K 的变体:

p←p mod 65536p \leftarrow p \bmod 65536

如果这减少了剩余 validation loss,则表明主要误差来源是 positional encoding 而非 hierarchical routing。

Training and Prefill Speed (Tab. 8)

40M SmallLM @ 12,288 tokens, RTX A4000, PyTorch 2.10.0+cu128, fp32, torch.compile=True.

ModelTrain msTrain tok/sForward msForward tok/s
HGA, Triton-fused299.8940,976102.98119,322
Dense RoPE baseline815.5615,067249.8749,177
Dense / HGA speedup2.72×–2.43×–

Correctness Checks (Tab. 7)

CheckMax absolute difference
Router, full coverage, summaries off vs. causal SDPA< 10−610^{-6}
Vectorized prefill, full token-level coverage vs. causal SDPA4.7×10−74.7 \times 10^{-7}
HA cache vs. HA no-cache1.5×10−51.5 \times 10^{-5}
HA decode vs. full recompute2.8×10−52.8 \times 10^{-5}

Fine-tuning Stability (Sec. 5.9)

Routed attention 在 exact-token mode 下进行了约 48M tokens 的 QK fine-tuning。训练稳定运行,未出现 causality leaks。这被视为稳定性结果而非headline质量数——零样本结果已表明直接权重复制无需 calibration 即接近 dense attention。

七、相关工作

固定稀疏注意力

Sparse Transformer、Longformer、BigBird 使用固定稀疏模式或 global tokens 降低注意力成本。HGA 保持固定 sink/local windows,但通过查询-键内容路由中间上下文。

基于内容的路由

Reformer 使用局部敏感哈希,Routing Transformer 在线聚类 tokens。HGA 的不同之处在于:使用预训练的 key 空间本身作为路由空间,并在路由后对精确 token K/V 进行注意力计算。

Kernel 和 Linear Attention

Performer 用随机特征映射替换 softmax 注意力。HGA 保持普通 softmax,但仅在选定的精确 tokens 上计算。

长上下文推理系统

vLLM 引入 PagedAttention 高效管理 K/V cache。MInference 识别稀疏模式(A-shape、Vertical-Slash、Block-Sparse)加速长上下文 prefill。HGA 互补:结合 sink/local windows + 内容路由 middle chunks + RAM-backed 存储抽象。

上下文扩展

YaRN 及相关方法扩展预训练模型的 RoPE 有效上下文窗口。HGA 与之兼容;论文指出在位置模 64K 的实验中, positional extrapolation 可能是剩余质量差距的主要来源。

八、局限性

  1. 稀疏近似本质:HGA 是稀疏注意力近似。如果路由器未选中包含相关 token 的 chunk/group,可能遗漏远处 token
  2. 系统演示而非全面基准:Qwen3-30B 结果目前仅为系统演示,需要系统性的长上下文检索、代码理解、文档 QA 等下游基准
  3. 路由表扫描:紧凑的 chunk-summary table 远小于 token K/V,但超长上下文应使用有界或索引路由而非朴素穷举扫描
  4. NIAH 评估有限:仅覆盖三个深度位置、一种上下文长度和三种 needle 类型
  5. 位置编码交互:剩余质量差距可能与长上下文位置编码的交互有关,而非路由算法本身

九、未来方向

  1. 层次化路由与长上下文位置编码的交互:系统性地研究两者的相互作用,可能独立改进稀疏注意力
  2. 自适应路由:动态分配额外路由容量,基于每个 query 的估计不确定性
  3. 可学习的路由 summaries:引入轻量级可训练路由表示,同时保持与预训练 checkpoint 的兼容性
  4. 更广泛的下游任务评估:长文档 QA、代码生成、检索基准、64K+ 上下文
  5. 位置模运算对齐:将 HGA 路由位置与基础训练窗口对齐而非扩展窗口

十、总结

核心贡献

  1. 即插即用分层路由:HGA 保留原始注意力投影,仅通过层次化路由改变查询可访问的历史 K/V,无需重新训练
  2. Mixed-RoPE summaries:高频/低频 RoPE pairs 分别处理,确保路由 summaries 与 RoPE-rotated queries 的相位兼容性
  3. 精确 token 输出:路由后打开真实 token K/V,在 softmax 上精确计算,无 summary gates 或校准参数
  4. RAM-backed KV store:完整历史 K/V 存于主机 RAM,按需拉取到 VRAM,使 32K context 在 32GB GPU 上可行
  5. 零样本可用:Qwen3-30B-A3B 零样本 32K 推理,Needle-in-a-Haystack 64K 上下文 100% 通过率
  6. 显著加速:40M SmallLM 上训练速度 2.72×,prefill 速度 2.43×

关键实验结论

  • 跨 4K-64K 上下文长度,层次化路由保持与密集注意力极其接近(loss gap < 0.01-0.02 nats)
  • 最新 chunk-group 路由策略在保持验证质量的同时进一步降低了路由 token 预算
  • 可与 DCA 和 YaRN 结合用于长上下文检索任务
  • 近似误差相对较小,剩余质量差距似乎更多与长上下文位置编码相关(而非路由算法本身)
  • RoPE position modulo 对齐可能是改善 quality gap 的关键方向

附图索引

编号文件名说明
Table 1—Large-model HGA demonstration (Qwen3-30B-A3B FP8 on RTX 5090)
Table 2—Dense vs. routed loss by context length and sparsity level
Table 3—Summary of HGA validation across context lengths
Table 4—Needle-in-a-Haystack results at 64K-token context
Table 5—Qwen3-0.6B fine-tuning comparison
Table 6—Copy-only dense-to-routed quality at 8192 tokens
Table 7—Selected correctness checks from the repository
Table 8—12,288-token speed benchmark for the 40M SmallLM

Note: arXiv HTML 版本不包含嵌入图像(仅包含表格和公式)。如需查看架构图和可视化结果,请参考 PDF 版本:arXiv:2606.30709 PDF