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 |
核心贡献:
- 提出 HGA(Hierarchical Global Attention)——一种无需重新训练即可替代预训练长上下文 Transformer 中密集因果注意力的即插即用分层路由算法
- 设计两级层次化路由(chunk-level → group-level),使用预训练的 key 空间本身作为路由空间,无需额外可训练参数
- 实现精确 token 级注意力输出:路由后打开真实 token 的 K/V,在 softmax 上精确计算,无 summary output tokens、summary gates 或标度校准参数
- 在 Qwen3-30B-A3B-Instruct-2507-FP8 上实现 32K-token 上下文零样本推理(单张 RTX 5090, 32GB VRAM),无需微调
- Needle-in-a-Haystack 检索测试在 64K-token 上下文下达到 100% 通过率(3/3 depth),仅 1.9% 稀疏度
- 在 40M SmallLM 上,直接权重复制仅产生 +0.018 nat loss gap,Triton-fused 实现训练速度提升 2.72×,prefill 速度提升 2.43×
二、核心思想
问题定义
预训练长上下文 LLM 的实际限制不仅在于 注意力计算,更在于必须驻留在加速器内存中的密集 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 | × | × | × |
| Reformer | LSH hashing | Summary tokens | × | × | × |
| Routing Transformer | 在线聚类 | Summary tokens | × | × | × |
| MInference | 固定稀疏模式 | 精确 token | × | × | × |
| vLLM | PagedAttention | 精确 token | × | × | GPU-only |
| HGA | 层次化 chunk→group | 精确 token | × | ✓ | ✓ (RAM) |
三、技术架构
整体框架
HGA 将序列划分为固定长度 的 chunks,每个 chunk 进一步细分为大小为 的 groups。关闭的 chunks 维护三类数据:
- Chunk key summary:第一级路由的候选集
- Group key summaries:第二级路由(当启用 group routing 时)
- 原始 token-level K/V:路由后的精确注意力输出
只有最后一项用于最终的注意力输出。
核心公式
标准因果注意力(Background)
精确 token 级感受野随 增长,总查询-键比较次数为 ,K/V cache 与所有先前 token 成正比。
Token 级计算成本
其中 和 分别是 first 和 recent chunks 数量, 是固定路由 token 预算(如 或 )。
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 中间位置旋转
这些 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 打分:
- Chunk-group 变体:对选定 chunks 内的 group summaries 打分,打开 top- groups 到精确 token-level K/V
- Qwen3 exact-token 变体:选定的 chunks 直接按 KV head 获取 token K/V
该方法不是纯固定稀疏模式——它有意结合 fixed sink/local windows 与 content-routed retrieval over the middle context。类似于 MInference 的 A-shape 思想,但 HGA 添加了基于预训练 key 空间的 top-k recall 路径。
精确 Token 输出注意力
路由后,注意力模块仅连接来自以下部分的真实 token-level keys 和 values:
- 当前 chunk 中因果可见的部分
- always-visible first 和 recent chunks
- 打开的 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 分区数据:
| 层级 | 存储位置 | 内容 | 说明 |
|---|---|---|---|
| Hot | GPU VRAM | Chunk summaries + always-visible first/recent chunks | 始终驻留 compute device,训练中保留梯度 |
| Warm | GPU 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
- , , ,
- 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,按需拉取到 VRAM | Sec. 4.3 |
| 三层缓存管理 | Hot/Warm/Cold 三级,LRU 自动收缩 | Sec. 4.1 |
五、代码实现分析
仓库: github.com/vfedosov77/HierarchicalGlobalAttention
实现架构:
- QwenRoutedAttention:大型模型的注意力模块替换器,在 exact mode 下运行
- ChunkSummaryManager:维护 chunk-level 和 group-level key summaries
- TieredKVStore:RAM-backed 三层 K/V 存储 + LRU 缓存管理
- Triton-fused Router:CUDA 融合的路由计算,减少 host-device 传输
- Correctness Tests:4 项数值验证(router vs SDPA, prefill vs SDPA, cache vs no-cache, decode vs full recompute)
关键实现特性:
- 路由后注意力模块与 dense causal SDPA 的数值差异 <
- Vectorized prefill 与 causal SDPA 差异
- HA cache vs no-cache 差异
- HA decode vs full recompute 差异
硬件要求:
- 训练: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 size | 64 tokens |
| Sink tokens | 128 (前 2 chunks) |
| Recent tokens | 512 (后 8 chunks) |
| Routed chunks | 16 (middle context) |
| Prefill block | 64 tokens |
| VRAM cache | Bounded LRU, auto-shrink |
Loss Gap vs Sparsity and Context Length (Tab. 2)
| Context | Sparsity 3.13% | Sparsity 6.25% | Sparsity 12.5% | Sparsity 25% |
|---|---|---|---|---|
| 4,096 | N/A | N/A | N/A | 2.456 / 2.460 () |
| 8,192 | N/A | 2.430 / 2.441 () | 2.430 / 2.437 () | N/A |
| 16,384 | 2.309 / 2.324 () | 2.309 / 2.322 () | 2.309 / 2.317 () | N/A |
| 32,768 | 2.204 / 2.227 () | 2.204 / 2.221 () | 2.204 / 2.214 () | N/A |
| 65,536 | 2.243 / 2.258 () | 2.243 / 2.253 () | N/A | N/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)
| Context | Sparsity | Loss Gap | Status |
|---|---|---|---|
| 4K | 25% | < 0.01 | Stable |
| 16K | 12.5% | < 0.01 | Stable |
| 32K | 12.5% | < 0.01 | Stable |
| 64K | 6.25% | ≈ 0.01 | Stable |
Needle-in-a-Haystack at 64K Tokens (Tab. 4)
Qwen3-30B-A3B-Instruct-2507-FP8 + HGA group-level routing, 零微调。
| Needle depth | Result | TTFT (s) |
|---|---|---|
| 25% | HIT ✓ | 948 |
| 50% | HIT ✓ | 1053 |
| 75% | HIT ✓ | 1332 |
| Overall | 3/3 = 100% | – |
Copy-only Routing Accuracy (Tab. 6)
40M SmallLM @ 8192 tokens, 无 HGA 微调, use_summaries=False。
| Model | Loss (nats) | Perplexity |
|---|---|---|
| Dense causal SDPA | 3.70516 | 40.657 |
| Chunk-routed, same weights | 3.72344 | 41.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 loss | 3.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 / dense | 679/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 之间的交互:
- 一些最近的长上下文 LLM 首先以较短的有效上下文长度预训练,然后通过 YaRN 等方法扩展上下文窗口
- 在 sparse routing 下,只有部分 transformer layers 观察到 distant tokens
- 因此,long-context positional extrapolation 引入的任何不准确可能在 sparse routing 下比 dense attention 更加明显
为验证此假设,论文评估了 RoPE position index wrapped modulo 64K 的变体:
如果这减少了剩余 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.
| Model | Train ms | Train tok/s | Forward ms | Forward tok/s |
|---|---|---|---|---|
| HGA, Triton-fused | 299.89 | 40,976 | 102.98 | 119,322 |
| Dense RoPE baseline | 815.56 | 15,067 | 249.87 | 49,177 |
| Dense / HGA speedup | 2.72× | – | 2.43× | – |
Correctness Checks (Tab. 7)
| Check | Max absolute difference |
|---|---|
| Router, full coverage, summaries off vs. causal SDPA | < |
| Vectorized prefill, full token-level coverage vs. causal SDPA | |
| HA cache vs. HA no-cache | |
| HA decode vs. full recompute |
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 可能是剩余质量差距的主要来源。
八、局限性
- 稀疏近似本质:HGA 是稀疏注意力近似。如果路由器未选中包含相关 token 的 chunk/group,可能遗漏远处 token
- 系统演示而非全面基准:Qwen3-30B 结果目前仅为系统演示,需要系统性的长上下文检索、代码理解、文档 QA 等下游基准
- 路由表扫描:紧凑的 chunk-summary table 远小于 token K/V,但超长上下文应使用有界或索引路由而非朴素穷举扫描
- NIAH 评估有限:仅覆盖三个深度位置、一种上下文长度和三种 needle 类型
- 位置编码交互:剩余质量差距可能与长上下文位置编码的交互有关,而非路由算法本身
九、未来方向
- 层次化路由与长上下文位置编码的交互:系统性地研究两者的相互作用,可能独立改进稀疏注意力
- 自适应路由:动态分配额外路由容量,基于每个 query 的估计不确定性
- 可学习的路由 summaries:引入轻量级可训练路由表示,同时保持与预训练 checkpoint 的兼容性
- 更广泛的下游任务评估:长文档 QA、代码生成、检索基准、64K+ 上下文
- 位置模运算对齐:将 HGA 路由位置与基础训练窗口对齐而非扩展窗口
十、总结
核心贡献
- 即插即用分层路由:HGA 保留原始注意力投影,仅通过层次化路由改变查询可访问的历史 K/V,无需重新训练
- Mixed-RoPE summaries:高频/低频 RoPE pairs 分别处理,确保路由 summaries 与 RoPE-rotated queries 的相位兼容性
- 精确 token 输出:路由后打开真实 token K/V,在 softmax 上精确计算,无 summary gates 或校准参数
- RAM-backed KV store:完整历史 K/V 存于主机 RAM,按需拉取到 VRAM,使 32K context 在 32GB GPU 上可行
- 零样本可用:Qwen3-30B-A3B 零样本 32K 推理,Needle-in-a-Haystack 64K 上下文 100% 通过率
- 显著加速: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