Graph Machine: Towards Better Pretraining via Edges

TL;DR

提出Graph Machine(GM),维护O(n)状态,通过稀疏动态路由提升预训练效率,替代75%Transformer层。

cs.LG 🔴 高级 2026-09-03 83 次浏览
Lintai Hou
图神经网络 稀疏注意力 预训练 边机制 Transformer改进

核心发现

方法论

本文提出的GM架构通过维护规模为O(n)的节点状态,利用可微更新的边(边索引与边权)实现稀疏、动态的邻域访问。引入边Referral机制,将多跳邻居关系逐步合成一跳邻居,结合边注意力(SEA)实现高效信息聚合。模型在Qwen3-0.6B基础上,将75%的密集Transformer层替换为GM稀疏层,从头训练15.7B tokens,发现即使每层只检索2个位置,模型性能仅略有下降,检索4个位置还能略微提升性能。

关键结果

  • 在15.7B tokens预训练中,采用每层仅检索2个KV位置的GM模型,其损失值仅比全密集模型高出约0.04,显示出极高的稀疏效率。
  • 引入4个检索位置时,模型性能略优于全密集Transformer,验证了边机制的有效性。
  • 相较于传统Transformer,GM显著降低了计算复杂度,检索位置比例不到0.2%,但保持了接近的性能水平。

研究意义

该研究突破了Transformer在长序列建模中的计算瓶颈,通过边机制实现O(n)复杂度,为大规模预训练模型提供了新思路。其稀疏、动态邻域访问机制,为未来高效长文本处理、图结构建模提供理论基础和实践方案,有望推动大规模语言模型的算力与存储优化。

技术贡献

核心贡献在于引入边(pointer-like对象)作为可微更新的邻域指针,结合Referral机制实现多跳邻居合成,创新性地将边索引与边权结合用于稀疏注意力,突破传统稀疏注意力固定邻域限制。模型在保持O(n)复杂度的同时,兼具灵活的邻域扩展能力,为Transformer架构提供了全新的稀疏化策略。

新颖性

首次将边机制引入大规模预训练模型中,利用可微邻域指针实现多跳邻居合成和稀疏注意力,区别于以往静态稀疏或固定邻域的设计,提供了动态、可扩展的邻域构建方式。

局限性

  • 当前模型在极长序列(超过1万Token)上的表现仍有限,边Referral机制在超大邻域时可能面临效率瓶颈。
  • 训练过程依赖高效的稀疏核实现,硬件优化需求高,实际部署仍需优化。
  • 模型在多任务下的泛化能力和迁移能力尚未充分验证,未来需结合下游任务进行系统评估。

未来方向

未来将探索多跳Referral的优化策略,提升邻域构建效率;结合硬件加速技术,降低训练和推理成本;扩展模型在多任务、多模态场景中的应用,验证其通用性和鲁棒性。

AI 总览摘要

随着大规模预训练模型的不断发展,如何在保证模型性能的同时降低计算成本成为关键挑战。传统Transformer在长序列建模中面临二次复杂度瓶颈,限制了其规模扩展。本文提出的Graph Machine(GM)架构,通过维护规模为O(n)的节点状态,利用可微更新的边机制实现稀疏、动态邻域访问,有效突破了这一瓶颈。

GM引入边Referral机制,将多跳邻居关系逐步合成为一跳邻居,结合边注意力(SEA)实现高效信息聚合。实验中,将75%的Transformer层替换为GM稀疏层,在15.7B tokens的预训练中,模型在只检索每层2个位置的情况下,性能仅略高于全密集模型,检索4个位置还能略微提升性能。这表明边机制在保持模型效果的同时,大幅降低了计算复杂度,检索比例不到0.2%。

该架构为长文本建模提供了新思路,突破了传统稠密注意力的限制,兼具灵活性和扩展性。未来,结合硬件优化和多任务验证,GM有望推动大规模语言模型的高效发展,开启长序列处理的新篇章。

深度解读

原文摘要

We introduce the Graph Machine (GM), an architecture that maintains an $O(n)$-sized state and accesses it through sparse, dynamic routing. Unlike methods with fixed-size states or sparse but static routing, GM preserves $O(n)$ complexity in its sparse layers without restricting the potentially accessible state size to $O(1)$. Instead, GM uses edges - pointer-like objects updated differentiably by a referral mechanism resembling pointer chasing. We replace 75% of the dense Transformer layers in Qwen3-0.6B with GM sparse layers and pretrain from scratch on 15.7B tokens. With only 2 of 4,096 tokens retrieved per KV head in each sparse layer, loss degrades only slightly; with 4, the best model marginally improves loss.

cs.LG