发表机构
Georgia Institute of Technology; University of California, Berkeley; Massachusetts Institute of Technology(佐治亚理工学院; 加州大学伯克利分校; 麻省理工学院)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
提出SMat-Attention,通过VC维可调的因果掩码族统一softmax与线性注意力,实现次二次预填充和常数时间解码,实验验证路由表达能力及对Mamba-2等模型的改进。
AI 中文摘要
长上下文序列模型面临一个根本性的权衡:softmax注意力以二次成本实现灵活的token级交互,而线性注意力通过将历史压缩为固定大小的状态,获得线性时间训练和常数时间解码。在这项工作中,我们探讨能否通过一种可调的结构概念来连接这两种机制。为此,我们通过一族具有结构化长距离路由的因果掩码引入结构化矩阵注意力(SMat-Attention),其行支撑集的VC维数为$d$。在我们的构造中,$d=1$恢复标准因果掩码,增大$d$允许更丰富的子集路由模式。我们给出了分块前向和后向算法以实现硬件高效性。对于长度为$T$的序列,尽管掩码是稠密的,硬路由构造在我们指定的族中仅需$O(T^{2-3/d}+T)$工作量。在固定视界流式处理中,远处前缀后的解码每个token耗时常数,使用$O(T^{1-1/d})$个缓存状态。因此,SMat-Attention使VC维数成为控制访问模式复杂度、预填充成本和解码内存的显式旋钮。实验上,子集路由和规则辅助多键检索实验展示了掩码的路由表达能力。扩展到Mamba-2和Gated DeltaNet,使用带top-$k$查询读取的学习路由,保持次二次预填充,在多种设置下优于骨干模型的召回准确率,并达到可比的小规模语言建模性能。
英文摘要
Long-context sequence models face a fundamental tradeoff: softmax attention uses flexible token-level interactions at quadratic cost, whereas linear attention obtains linear-time training and constant-time decoding by compressing history into a fixed-size state. In this work, we ask whether we can connect these regimes through a tunable notion of structure. To this end, we introduce Structured Matrix Attention (SMat-Attention) via a family of causal masks with structured long-range routing whose row supports have VC-dimension $d$. In our construction, $d=1$ recovers the standard causal mask, and increasing $d$ permits richer subset-routing patterns. We give chunkwise forward and backward algorithms to enable hardware-efficiency. For sequences of length $T$, the hard-routing construction takes $O(T^{2-3/d}+T)$ work, despite the mask being dense, for our prescribed family. In fixed-horizon streaming, decoding after the distant prefix takes constant time per token using $O(T^{1-1/d})$ cached states. SMat-Attention therefore makes VC-dimension an explicit knob governing access-pattern complexity, prefill cost, and decoding memory. Empirically, subset-routing and rule-assisted multi-key retrieval experiments illustrate the masks' routing expressiveness. Extensions to Mamba-2 and Gated DeltaNet using learned routing with top-$k$ query reads retain subquadratic prefill, improve recall accuracy over the backbones in several settings, and achieve comparable small-scale language-modeling performance.