Graph Machine:通过边实现更好的预训练

Graph Machine: Towards Better Pretraining via Edges

arXiv: 2609.02881v1

论文信息

标题: Graph Machine: Towards Better Pretraining via Edges

作者: Lintai Hou

发布日期: 2026-09-02

arXiv ID: 2609.02881v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:Transformer 的 O(n) 状态虽然避免了历史压缩,但密集注意力造成 O(n²) 访问,而现有稀疏注意力又往往采用静态访问模式,本文探索能否在保持 O(n) 状态的同时,实现动态、O(1) 大小的访问。
  • 核心方法:提出图机器(Graph Machine, GM),用可微更新的整数边索引和浮点边权重构成稀疏图,通过 “referral” 机制在相邻层间传播和重定向指针,再用这些指针驱动的稀疏注意力聚合信息。
  • 关键结果:在 Qwen3-0.6B 骨干上,将 75% 的密集 Transformer 层替换为 GM 稀疏层,每个 KV 头仅检索 4,096 个位置中的 2 或 4 个;最佳配置 Hyperion-K16-R3-S 的最终测试损失比 Qwen3 低约 0.003(表 3),同时参考与注意力计算量减少约 19%(表 2)。
  • 主要局限:论文实验规模较小(0.6B 参数、15.7B tokens),只评估了预训练测试损失,未覆盖下游任务;原型实现速度在 H100 上通常比 Qwen3 慢数倍,作者认为很依赖内核实现。
  • 适合读者:对高效注意力、稀疏路由、图谱神经网络与语言模型架构融合感兴趣的 LLM 研究者、架构设计者和推理优化工程师。

论文背景和研究动机

作者首先按状态大小、单步访问大小和动态寻址位数对序列模型进行分类。RNN 和 SSM 保持 O(1) 状态,因此必须压缩历史;标准 Transformer 保持 O(n) 键值状态,但每步访问 O(n) 位置,导致 O(n²) 复杂度。滑动窗口等稀疏注意力把访问降到 O(1),但这些位置的集合在 token 到来前就已固定,因此是 “稀疏但静态” 的。作者用信息论观点指出:若要在 O(n) 状态上做动态、O(1) 大小的访问,每个新 token 需提供 Θ(log n) 位地址信息,用于直接指定要读取的一小块状态。在常见序列长度下,这些地址位可以装进 int64,实际存储成本近似恒定。

Graph Machine 正是落在这一新类别:维持 O(n) 个节点特征和边状态,单步访问 O(1) 个条目,并使用 O(log n) 位动态地址。作者认为,这相当于把 “注意力中传递地址、解析地址” 的核心机制显式化为可微的边索引和边权重,而不是让注意力用大量内积分数隐式地做这件事。

核心方法和技术细节

GM 维护三类状态:形状为 n×dn\times d 的节点特征,以及形状均为 n×k×s1n\times k\times s_1 的边索引和边权重。每条边包含 s1s_1 个成员位置;边索引给出目标节点,边权重给出对应非负权重,并在成员位置上归一化。输入 token 的初始边指向自身和最近的 k−1k-1 个前序 token,初始权重集中在单一目标上,其余成员位置用零索引和零权重填充。

GM 稀疏层包含两个新子模块。稀疏边传递(SER) 做两跳 referral:给定两个稀疏邻接矩阵 A1A_1 和 A2A_2,得到 A′=Sparsify⁡s1(A1A2)A'=\operatorname{Sparsify}_{s_1}(A_1 A_2)。在路径层面,这相当于把边 e1=(n1→w1n2)e_1=(n_1\xrightarrow{w_1} n_2) 与 e2=(n2→w2n3)e_2=(n_2\xrightarrow{w_2} n_3) 组合为 e1∘e2=(n1→w1w2n3)e_1\circ e_2=(n_1\xrightarrow{w_1 w_2} n_3)。具体实现中,SER 先将 kk 条存储边混合成 2k2k 个通道,形成 kk 对两条腿;每条腿分别稀疏化成 s2s_2、s3s_3 个成员,然后用第一腿的索引去检索第二腿的索引和权重,得到最多 s2s3s_2 s_3 个候选路径,最后通过 Sparsify 选出 s1s_1 个成员。

稀疏边注意力(SEA) 则负责特征聚合。它把 kk 条存储边混合成 gg 个通道(每个 KV 头一个),并稀疏化到 s4s_4 个位置。这些位置的键和值参与查询-键评分,同时混合边权重的对数作为额外偏置加入 softmax:

aij=softmax⁡j∈N(i)(τqi⊤kjdh+log⁡wij)a_{ij}=\operatorname{softmax}_{j\in\mathcal{N}(i)}\left(\tau\frac{q_i^\top k_j}{\sqrt{d_h}}+\log w_{ij}\right)

等价于在概率空间中做查询-键因子与边权重的乘积专家组合。作者还引入了 “混合” 操作:先对边权重施加温度缩放 Tτ(a)=aτ∑jajτT_\tau(a)=\frac{a^\tau}{\sum_j a_j^\tau},再通过节点特征投影得到混合矩阵,最后再缩放和稀疏化。该操作类似于 CNN 中的通道混合,用于在保持稀疏性的同时增加非线性。Sparsify 本身会合并重复索引、保留最高权重项并重新归一化;由于是硬 top-ss 选择,只有保留下来的权重会获得梯度。

实验中,GM 与 Transformer 按 3:13:1 稀疏-密集比例混合,采用 7 组 [S,S,D,S][S,S,D,S] 块堆叠。基线是 Qwen3-0.6B。所有 GLM 都使用相同骨干、训练超参数和随机种子,在 FineWeb-Edu 的 15.7B token 子集上从头预训练。

创新点和贡献

该工作的主要贡献在于提出并验证了 “通过边进行动态路由” 的稀疏注意力范式。与固定窗口注意力不同,GM 的边索引是持久的、可微更新的状态,可以在多步 referral 中把地址信息沿图结构传播,类似于指针追踪。这使得稀疏注意力的 “支持集” 不再由位置先验决定,而是由模型可学习的指针状态决定。

第二个贡献是把边表征做成稀疏坐标格式,并设计了 Sparsify、Mix、SER、SEA 等兼容现代 GPU 的算子。前作 GM-1 的密集边表示需要三次时间和二次空间,只能处理数百节点;GM-2 的稀疏表示使其能扩展到真实 LLM 训练规模。作者还引入了 refresh 和 realignment:refresh 允许从初始边或最近密集注意力层重新引入结构信息;realignment 让 SEA 用检索到的索引和最终注意力权重回写部分存储边,把近似 referral 与特征证据结合起来。

工程上,本文展示了:在 Qwen3-0.6B 上替换掉大多数密集层后,GM 混合模型用远低于密集因果注意力的 KV 访问量,仍能保持接近甚至略优的测试损失。表 2 显示,Theia 类模型 SEA 的 KV 访问仅为密集因果注意力的 0.098%,Hyperion 类为 0.195%;同时参考与注意力计算量相对 Qwen3 减少约 10%–30%。

实验结果分析

在本次实验设置下(表 3),Qwen3 最终测试损失为 2.587。Theia 类(稀疏预算 s4=2s_4=2)最佳配置最终损失约比 Qwen3 高 0.014;Hyperion 类(s4=4s_4=4)最佳配置 Hyperion-K16-R3-S 最终损失低约 0.003。该结果说明:在该数据集和模型规模下,将 75% 的密集层替换为每 KV 头只检索 2 或 4 个位置的 GM 稀疏层,不会实质损害语言模型预训练质量;当检索 4 个位置时,最佳模型甚至出现轻微损失改善。

从消融角度看,referral 步骤对性能很重要:Theia-K24-R3 相比无 referral 的 Theia-K24 在结束时损失低 0.026。增加存储边数也有帮助:Hyperion-K24-R3 比 K16-R3 低约 0.003,但代价是交叉边混合复杂度随 kk 二次增长。dense refresh 带来一定收益:Hyperion-K16-R3-S 比 K16-R3 在结束时低 0.004。不过作者也指出,参数数量和估计计算量并不是性能的强预测因子;例如 Hyperion-K16-R3-S 虽然参数比 Qwen3 多 11%,但参考与注意力计算量少 19%,仍取得最佳结果。

必须注意,这些结论都基于小规模预训练和单一测试损失指标,论文没有提供下游任务数据,因此只能视为该实验设置下的有限结论,作者也承认需要更大规模和更多评估来确认。

实践建议

对于希望把 GM 稀疏层集成到现有 LLM 架构中的团队,以下实践方向可作参考:

首先,从与注意力头数匹配的边数开始。论文中 Hyperion-K16-R3-S 使用 16 条边、4 个检索位置、3 个 referral 步骤和 8 条 dense-refresh 边,在计算与质量之间取得了较好平衡。建议在目标骨干上先固定边数等于注意力头数,再逐步试验 k=24k=24 或 3232 的配置。

其次,重视 custom kernel 开发。论文的原型实现用 PyTorch 加自定义 Triton 稀疏化内核,在 H100 上通常慢于 Qwen3,但作者提到初步实验表明某些配置在 RTX 4090 上已接近 Qwen3 训练吞吐。若要实际部署,需要针对 Sparsify、索引聚集和混合操作做深度融合,特别是减少稀疏化带来的内核启动和显存往返开销。

第三,保留 refresh 和 realignment。论文的初步实验表明二者都有好处,dense refresh 在 Hyperion-K16-R3 上带来约 0.004 的损失改善。实际应用中,密集层提供的 refresh 边可以作为一种全局结构注入,而 realignment 能让注意力证据修正边分布,缓解近似 referral 的误差积累。

最后,不要只盯着测试损失。作者未评估下游任务,因此落地前应在自有任务上验证 GM 稀疏层对推理速度、显存占用和长序列能力的实际影响。若目标场景以长上下文或低延迟推理为主,GM 的 0.098%–0.195% KV 访问比例可能带来显著收益,但前提是稀疏内核足够高效。建议先用小规模预训练和全面下游评测确认成本-质量边界,再决定是否替换更大模型中的密集层。