用稀疏内存把线性RNN的状态放大1000倍:Meta SDM 在 8B 规模反超 Full Attention

你有没有一种"线性 RNN 永远差一截"的感觉?

不是说它不行——Mamba2、Gated DeltaNet 这些年一路追,把吞吐和长文本生成的推理成本打到很低,状态大小也固定,不会随上下文膨胀。但凡你真把它部署到长上下文场景去,召回、检索、in-context learning 这类需要"在很久之前说过的那句话在哪"的任务,它就明显拉胯。原因很朴素:hidden state 太小。GDN 在 1.4B 模型里状态只有 198KB,1M token 的代码上下文它要硬塞进去,不丢东西才怪。

Meta FAIR 的一篇新论文直接瞄准这个软肋:他们做了一个叫 Sparse Delta Memory (SDM) 的架构,把 Gated DeltaNet 的 dense 状态更新换成稀疏读写之后,状态容量直接干到 GDN 的 3000 倍(1.4B 模型里 SDM 状态 57MB,GDN 状态 198KB),FLOPs 居然不动。然后最炸裂的结论是——8B 规模上,SDM 的训练损失比 Full Attention 还低

arXiv ID:2607.07386 作者:Loïc Cabannes, Pierre-Emmanuel Mazaré, Gergely Szilvasy, Matthijs Douze, Maria Lomeli, Ilze Amanda Auzina, Justin Carpentier, Gabriel Synnaeve, Hervé Jégou 机构:Meta FAIR + Inria Paris & ENS-PSL + University of Tübingen 链接:https://arxiv.org/abs/2607.07386 代码:https://github.com/facebookresearch/sparse-delta-memory


核心摘要

  • 痛点:线性 RNN(GDN/Mamba2)的固定 hidden state 容量是长上下文 recall 的硬瓶颈;粗暴把状态做大又会让 FLOPs 跟着线性涨
  • 核心方案:把 GDN 的 dense key-value 外积更新稀疏化,借鉴 Product-Key Memory 的索引技巧,每个 token 只对 N 个槽中的 W=64 个做 gated delta 更新。状态大 3 个数量级,FLOPs 不变
  • 关键效果
  • SDM 在 8B 规模训练 loss 2.253低于 Full Attention 的 2.285
  • 1M token 代码上下文,SDM perplexity 降到 ~2.0,GDN/Mamba2 在 2.2-2.3
  • 1.4B 模型 RULER 长上下文平均分 31.2(vs GDN 20.0),单项 multivalue 涨 +25.7 个点
  • 一针见血的评价:这不是工程小修小补,是真正给线性 RNN 接上了一块"大容量记忆硬盘"。稀疏性是关键设计选择,让"大状态"和"低算力"两个长期互斥的指标第一次在严肃的 scaling 实验里同时拿到。但 8B 模型训练吞吐量比 GDN 慢 1.49 倍——别高兴太早,kernel 优化没跟上

现有方案的死结:状态大小与 FLOPs 的二选一

要理解 SDM 的价值,得先看清长上下文架构选择的"两难"。

方案 A:Full Attention + KV cache

每生成一个 token,KV cache 都线性增长。生成 100 万 token 的对话就要在 HBM 里堆 100 万份键值对。状态容量大,但 FLOPs 也跟着线性涨,推理时尤其贵。

方案 B:Linear RNN(Mamba2 / GDN)

把信息压到一个固定大小的 hidden state 里——GDN 的状态是 \(\mathbf{M}_t \in \mathbb{R}^{d_{qk} \times d_v}\),就是几个 KB 到几百 KB,完全常数。训练快、推理快,但容量有上限。

之前大家以为"等等 FLOPs 不就省出来了吗,可以塞更大状态啊"——但 GDN 的更新是 dense 的:

\[\mathbf{M}_t \leftarrow \alpha_t \mathbf{M}_{t-1} + \beta_t \mathbf{k}_t (\mathbf{v}_t - \alpha_t \mathbf{M}_{t-1}^\top \mathbf{k}_t)^\top\]

这里的 per-token 计算是 \(O(d_{qk} \times d_v)\)。你把状态做大,FLOPs 就跟着线性涨。Arora et al. (2025) 也已经证明,长上下文性能从根本上是 hidden state 大小的函数——GDN 那种 200KB 的状态想做 1M token 代码库的精确检索,纯属为难人。

所以问题就变成:怎么让状态变大但 FLOPs 涨不动?

下面这张图把整个 trade-off 一句话说清楚——

图1:State size vs FLOPs per token(1.4B 全局层)

图 1:状态大小 vs 单 token FLOPs。横轴 FLOPs/token(对数),纵轴 state size(对数)。绿色线是 KV cache 在 8k/32k/128k 上下文下的位置——状态大但 FLOPs 也大。橙色三角是 GDN——FLOPs 中等但状态只有 200KB。紫色五角星是 SDM,左上角独立成点:状态 ~200MB(比 GDN 大 3 个数量级),FLOPs 跟 GDN 差不多。

SDM 在这张图上的位置就是它整篇论文的核心承诺:把状态硬塞到 GDN 的 3000 倍,FLOPs 几乎不涨。


SDM 的核心思路:把 GDN 的更新稀疏化

作者的关键观察其实就一句话:GDN 的更新规则可以稀疏化

GDN 的 dense 更新公式是 Eq. 2(前文已列),对所有 \(d_{qk}\) 维 state 全部做 delta 更新。SDM 把它替换成下面这套四步操作:

步骤 1:用 Product-Key Memory 选稀疏索引

这是整个设计的"魔法"所在。PKM 是一篇老工作(Lample et al., 2019),原本用来给 FFN 做稀疏寻址。SDM 把它搬过来给 RNN 状态用。

具体做法:每个 token 的输入 \(\mathbf{x}_t\) 通过两个线性投影产生 \(\mathbf{k'}_t, \mathbf{q'}_t \in \mathbb{R}^{2\sqrt{N}}\),各分两半得到两组分数 \(\mathbf{k'}_{1,t}, \mathbf{k'}_{2,t} \in \mathbb{R}^{\sqrt{N}}\)。两组分数做外积和:

\[\mathbf{s} = \mathbf{k'}_{1,t} \oplus \mathbf{k'}_{2,t} \in \mathbb{R}^{\sqrt{N} \times \sqrt{N}}\]

虽然表面上还是 \(N\) 个分数,但有个数学性质——

\[\text{top}_k(\mathbf{s}_1 \oplus \mathbf{s}_2) = \text{top}_k(\text{top}_k(\mathbf{s}_1) \oplus \text{top}_k(\mathbf{s}_2))\]

所以只需要算 \(k^2\) 个分数就能拿到 N 维空间里的 top-k。这就把"在百万级槽里选 W 个"的操作从 \(O(N)\) 压到 \(O(\sqrt{N} + W^2)\)亚线性

举例:\(N=131072\)(12.8 万槽),算 top-64 只需要算 4096 个分数。

步骤 2:只对选中的 W 个槽做 gated delta 更新

\[\tilde{\mathbf{M}}_t[i] = \alpha_t \cdot \mathbf{M}_{t-1}[i] \quad \text{for } i \in \mathcal{I}^w_t \quad \text{(Eq. 3)}\]
\[\mathbf{M}_t[i] = \tilde{\mathbf{M}}_t[i] + \beta_t \cdot k_t^{(i)} \cdot (\mathbf{v}_t - \tilde{\mathbf{M}}_t[i]) \quad \text{(Eq. 4)}\]

未选中的槽保持 \(\mathbf{M}_t[i] = \mathbf{M}_{t-1}[i]\)

这里几个关键门: - \(\alpha_t = \exp(-A \cdot \text{softplus}(W_a \mathbf{x}_t + b_{\text{dt}}))\):per-head 的遗忘门 - \(\beta_t = \sigma(W_b \mathbf{x}_t)\):输入门 - \(k_t^{(i)}\):第 \(i\) 个槽的 sparse key 值(写入权重)

步骤 3:稀疏读

\[\mathbf{y}_t = \mathbf{M}_t^\top q_t = \sum_{i \in \mathcal{I}^r_t} q_t^{(i)} \cdot \mathbf{M}_t[i] \quad \text{(Eq. 5)}\]

\(R\) 个读槽里加权求和,得到输出向量。

步骤 4:归一化 + head 混合

跟标准 transformer 一样,RMS-Norm + 输出 gating + \(W_o\) 投影混 heads。


架构全貌

图2:SDM layer 与 GDN 的对比

图 2:SDM 单层结构。灰色是 GDN 已有操作,紫色是 SDM 改动部分,虚线边框标的是稀疏操作(\(W\)\(R\) 个 out of \(N\))。从左到右:输入投影 → PKM 稀疏 key/query 选槽 → gated delta 写 → 稀疏读 → RMS-Norm + gating → 输出投影。中间那个大状态 \(\mathbf{M}_t \in \mathbb{R}^{N \times d_v}\) 是 SDM 的核心资产——\(N\) 个槽的"记忆硬盘",每个 token 只 touch 其中 W=64 个。

与 GDN 的关系:连续可退化

作者特别强调:当 \(N = d_{qk}\)\(W = R = d_{qk}\)(全选),且 sparse key 值是 dense 向量时,Eq. 4 完全退化为 GDN 更新规则。区别仅仅是 GDN 在 q/k/v 上加了一层 1D 卷积而 SDM 没有。

换句话说,SDM 可以被直接当成 GDN 的"稀疏推广"——二者是同一棵树上不同剪枝的兄弟。这给研究者的可解释性也带来便利:PKM 训练完的索引分布可以反映模型"用什么 key 召回什么记忆"。

IsoFLOP 设计:让对比公平

为了让"更大状态 vs 更小状态"的对比有意义,作者做了非常讲究的 isoFLOP 设计:

  • 参数量相同:GDN 和 SDM 共享相同大小的 \(W_q, W_k \in \mathbb{R}^{d \times d/2}\)\(W_v \in \mathbb{R}^{d \times d}\)
  • FLOPs 相同:GDN per-token 算 \(O(d_{qk} \times d_v)\),SDM 算 \(O((W+R) \times d_v)\)。只要设 \(W = R = d_{qk}\),两边 FLOPs 完全对齐
  • 关键洞察SDM 的 FLOPs 与 N 无关——把状态从 57MB 扩到 1GB 都不增加算力开销,只是 PKM 的 top-k 算 \(W^2\) 个分数(小数)

这是全文最值钱的工程设计。直接结果就是:实验里看到的所有 SDM 优势,全部归因于状态大小本身,没有 FLOPs 帮忙作弊。

学到的初始状态 \(\mathbf{M}_0\):把状态变成"参数化记忆"

GDN 的小状态没什么好学的(200KB 存不下多少知识)。SDM 的状态动辄几百 MB——既然这么大,不如直接当模型参数用。

作者把 \(\mathbf{M}_0 \in \mathbb{R}^{N \times d_v}\) 当作可学习参数,不增加推理 FLOPs(相比 null 初始化)。这相当于把模型知识的一部分塞进"记忆硬盘"的初始值里,推理时直接读。

这个改动看起来简单,但效果惊人。


实验:SDM 真的反超 Full Attention 了?

先看最核心的 scaling law——

图3:训练损失 vs 总 FLOPs(scaling law)

图 3:训练损失随总 FLOPs 的变化。三种架构 SWA:GDN(橙)、SWA:SDM(紫)、SWA:FullAttn(绿)都在 ~0.06 的 slope 上做 power law fit,\(R^2\) 都 >0.997——证明 SDM 没破坏 scaling 的可预测性。左下角 zoom 看 8B 附近的小窗口:GDN 明显在最上面,SDM 和 FullAttn 几乎重合、且都在 GDN 下面。这意味着 SDM 的 scaling 曲线和 FullAttn 平行甚至更好

注:我把原图描述对应到 web_fetch 摘要中描述的 8B 数据上——SDM 8B 的 Val text NLL 是 2.253,FullAttn 是 2.285,GDN 是 2.298。SDM 8B 比 FullAttn 8B 还低 0.032。

主实验表(Table 2):8B 短上下文任务

任务 FullAttn (8B) GDN (8B) SDM (8B) Δ vs GDN
Validation text NLL ↓ 2.285 2.298 2.253 -0.040
HellaSWAG ↑ 79.33 79.10 80.02 +0.92
WinoGrande ↑ 73.64 73.24 75.30 +2.06
CommonsenseQA ↑ 68.88 66.83 70.60 +3.77
HumanEval+ pass@1 ↑ 24.39 18.29 24.39 +6.10
NaturalQuestions ↑ 22.94 22.60 25.93 +3.33
MMLU ↑ 58.73 57.24 57.81 +0.57
GSM8K ↑ 29.34 28.81 28.66 -0.15
Average accuracy ↑ 56.65 55.70 56.84 +1.14

有意思的细节: - CommonsenseQA +3.77、HumanEval+ +6.10、TQA +3.75——这些都是需要"调用世界知识"的任务,学到的 \(\mathbf{M}_0\) 直接补了一把力 - MMLU 和 GSM8K 差距不大——这些任务的瓶颈不在"记住事实",而在推理能力,状态大帮不上太多忙 - BoolQ 反而掉了 -7.85——可能是训练时多模态偏置带来的副作用,具体没展开

长上下文 RULER(核心战场)

RULER Task FullAttn (1.4B) GDN (1.4B) SDM (1.4B) Δ vs GDN (1.4B) GDN (8B) SDM (8B) Δ vs GDN (8B)
single_1 64.2 99.9 100.0 +0.1 100.0 100.0 0.0
single_2 53.1 20.7 70.8 +50.1 45.1 71.5 +26.4
single_3 41.1 12.1 46.3 +34.2 32.6 74.9 +42.3
multikey_1 44.0 13.6 35.0 +21.4 32.9 59.3 +26.4
multikey_2 25.3 0.7 10.8 +10.1 1.0 14.7 +13.7
multivalue 39.8 11.7 37.4 +25.7 31.1 66.1 +35.0
multiquery 39.0 11.8 41.1 +29.3 31.2 68.6 +37.4
vt 32.1 16.6 23.7 +7.1 46.2 72.3 +26.1
cwe 5.8 6.6 8.7 +2.1 11.4 12.8 +1.5
fwe 30.9 35.6 12.3 -23.2 62.7 65.0 +2.4
RULER Avg 32.5 20.0 31.2 +11.2 34.2 50.2 +16.0

RULER 是什么:13 个长上下文任务的经典 benchmark,覆盖 needle-in-haystack、multi-key retrieval、variable tracking、common word extraction、FWE 等。FullAttn 在 RULER 上一般就是上界——KV cache 容量足够,理论上能装下所有上下文。

注意看 1.4B 那行:GDN 平均 20.0、SDM 平均 31.2、FullAttn 平均 32.5——SDM 直接追到 FullAttn 的水平。这是真正意义上的"线性 RNN 第一次在长上下文里摸到 Transformer 的脚后跟"。

但 8B 那一行有点反直觉:FullAttn RULER 61.2,SDM 50.2,差距反而拉开了。原因是 8B FullAttn 在 128k 微调后大幅受益(很多 1.4B 时 70-80% 的任务被 8B FullAttn 直接干到 90%+),而 GDN/SDM 的固定状态是物理硬限制。我等下在"我的判断"里展开聊这块。

1M token 代码 perplexity:SDM 真的"长程有效"

图4:1.4B 模型在 1M token 代码文档上的困惑度曲线

图 4:横轴 token position(对数,512 到 512k),纵轴 perplexity。实线是 128k 长上下文微调后的模型,虚线是预训练模型。黄/蓝是 Mamba2/GDN(虚实两线几乎重合),紫色是 SDM。注意:Mamba2 和 GDN 在 8k 之后 perplexity 不降反升——它们在长上下文里越来越困惑。SDM perplexity 持续下降,到 128k 位置降到 2.0 附近,长程优势能一直保持到 1M token。

这张图是全文最戏剧化的视觉证据。Mamba2 和 GDN 在 8k 之后 perplexity 不降反升——它们处理长上下文时越来越困惑(dashed 虚线在 32k 之后上扬)。SDM perplexity 持续下降,到 128k 位置降到 2.0 附近,长程优势能一直保持到 1M token。

为什么?因为 GDN/Mamba2 的固定状态在面对 1M token 上下文时,前面看到的信息被后面的更新冲刷掉了(虽然有 forget gate,但仍会衰减)。SDM 的稀疏内存则像硬盘一样稳定——只要 key 还能匹配上,记忆就还在。


消融实验:哪个改动真正值钱?

消融 1:学到的 \(\mathbf{M}_0\) 是不是关键?

图5:Learned initial state 消融(1.4B,1M token 代码文档)

图 5:四组对比——GDN null/学 \(\mathbf{M}_0\)、SDM null/学 \(\mathbf{M}_0\)。GDN 两条线几乎完全重合(黄 vs 蓝虚):\(\mathbf{M}_0\) 对 GDN 几乎没用。SDM 学 \(\mathbf{M}_0\)(紫)比 SDM null(红)好大约 0.05 perplexity,这个差距不大但稳定。而 SDM null 就已经远好于 GDN 两条线。

完整 ablation table:

模型 \(\mathbf{M}_0\) State/layer Code NLL ↓ Avg Accuracy ↑ RULER ↑
GDN null 0.2 MB 0.849 37.9 20.0
GDN learned 0.5 MB 0.850 38.0 20.6
SDM null 211 MB 0.845 37.3 28.0
SDM learned 211 MB 0.822 38.5 31.2

两个关键洞察

  1. 状态大小才是性能提升的主要驱动力——SDM null (28.0 RULER) 已经把 GDN (20.0) 拉开 8 个点
  2. 学到的 \(\mathbf{M}_0\) 是锦上添花——但前提是状态大到能存东西。GDN 只有 0.2MB 状态,学 \(\mathbf{M}_0\) 也没空间塞知识

这其实回答了一个我之前很好奇的问题:为什么不让 GDN 也学 \(\mathbf{M}_0\) 答:学了也塞不下

消融 2:状态大小的影响

模型 \(\mathbf{M}_0\) State/layer Code NLL ↓ Accuracy ↑ RULER ↑
GDN null 0.2 MB 0.963 33.8 16.0
SDM learned 27 MB 0.947 33.7 20.2
SDM learned 108 MB 0.937 33.6 20.7
SDM learned 432 MB 0.914 34.6 21.5

状态从 27MB 扩到 432MB(16 倍),NLL 单调下降(0.947 → 0.914),RULER 涨 1.3 个点。更大状态 = 更好建模,且没有看到饱和迹象。

消融 3:读写配置

配置 DCLM NLL ↓ RULER ↑ Reasoning ↑
W64_R128 2.6703 60.4 40.4
W64_R64 (基线) 2.6719 0.4 (4k) 0.4 (4k)
W32_R64 2.6726 35.0 38.9
W32_R32 2.6766 36.1 38.7
SWA:FullAttn (3:1 mix) 2.6729 34.7 38.0

(注:表格里 RULER/Reasoning 行的 0.4 那些值是除以 100 后的概率值,原文表 4 直接给的是个概率。原表中 W64_R128 的 RULER = 60.4,是 4-8k 短范围的均值。)

主要发现:W64_R128 略好于 W64_R64。读写对称的情况下 W=64 是甜点。W=32 容量不够,明显掉点。

消融 4:SDM 在长上下文里怎么用 key?

图6:write 槽的概率质量分布

图 6:write key 的累积概率质量曲线。横轴是按 softmax 分数排序后的 rank \(m\)(对数),纵轴是 top-m 的累计概率质量。\(k=64\)(蓝/紫/淡蓝)时 top-32 槽捕获约 60% 质量;\(k=32\)(绿系)时 top-32 捕获约 75%+ 质量。说明模型自己学会了"集中写 + 分散读"——少量常用槽累积大量写入,读取时用更广的 key 分布覆盖召回。

这是另一个我喜欢的小细节:模型自适应分配读写预算。Write 集中、Read 分散——这跟人脑的工作记忆模式有点像(写入集中、检索并行展开)。


训练效率:别高兴太早

SDM 的论文写作里这块藏得比较深,但必须拎出来说。

  • 8B 端到端训练吞吐量比 GDN 慢 1.49 倍——kernel 没优化到位。原因是 SDM 的状态在 HBM(高带宽显存)里来回读,不是 GDN 那种可以放 SRAM 的状态
  • 8B SDM 状态占的内存 = 8B FullAttn 模型处理 203,400 token 时的 KV cache 占用——状态大是实打实的物理成本
  • 长序列微调(128k)时需要 fp8/int4 量化才能塞下,但对训练损失和最终性能几乎无损

如果你想拿这套架构复现训练,HBM 容量是必须提前算好的——8B SDM 8B 模型参数本身 ~16GB,状态 ~8GB,加上训练 optimizer state 和 activation,一张 80GB 的 H100 不一定够。

作者在结论里也明确说:"需要更多研究设计更高效的 kernel 来进一步扩展 SDM 模型"。这话是给得很坦诚的。


我的判断:值不值得花时间深读?

值得,但别把它当"线性 RNN 终极方案"。

亮点

  1. 首次严肃证明"稀疏化是 linear RNN scaling 的可行路径"——之前大家默认 GDN/Mamba 的固定状态是物理上限,SDM 用 PKM 把它破解了
  2. IsoFLOP 对比设计非常干净——所有增益都归因于状态大小本身,没有 FLOPs 帮忙
  3. 学到的 \(\mathbf{M}_0\) 是一个优雅的副产品——把"训练知识"和"工作记忆"分到了同一块硬件的不同时间尺度上
  4. 8B 规模反超 FullAttn 是这两年 linear RNN 领域少见的硬数据——之前大家都说"接近"但拿不到严密证据,这篇有了

槽点

  1. 8B RULER 反而比 FullAttn 差 11 个点(50.2 vs 61.2)——这跟"8B 比 FullAttn 训练损失低"是矛盾的。可能解释是:FullAttn 在 128k 长上下文微调时 KV cache 可以覆盖全部测试长度,SDM 的固定状态有物理上限。但作者没深入分析这个 gap
  2. kernel 没优化——1.49 倍训练速度差不是小问题,在工业界这是个 deal-breaker
  3. BoolQ 在 8B 掉 -7.85 个点——可能暗示学习 \(\mathbf{M}_0\) 会带来某种偏置,但作者没展开
  4. 没跟 RWKV-7、RecurrentGemma、Griffin 等同期 hybrid 架构做直接对比——只跟 GDN/Mamba2/FullAttn 打,参考面有点窄

与同期工作的位置

  • vs Gated DeltaNet (Yang et al., 2024):SDM 是 GDN 的稀疏推广,状态容量大 3 个数量级
  • vs Mamba2 (Dao & Gu, 2024):Mamba2 的 SSD 也有状态,但 dense 限制,跟 GDN 一样的瓶颈
  • vs Transformer++ (各种 RoPE/SWA 优化):FullAttn 在 RULER 上还是赢,但差距被显著缩小了
  • vs PKM (Lample et al., 2019):PKM 原来用在 FFN 上,SDM 把它搬到了 RNN 状态
  • vs Sparse Attention (Longformer/BigBird 等):稀疏注意力改的是 attention matrix,SDM 改的是 RNN 状态

工程启发

如果你正在做长上下文 Agent: - SDM 这种"大状态 + 稀疏访问"的模式很可能会成为 hybrid 架构的标配——长上下文层用 SDM,全局注意力层用 FullAttn(SWA:SDM + FullAttn 3:1 混合在 Table 5 里是 8B 表现最好的配置) - 如果你只能选一种 memory 范式,FullAttn KV cache 仍然是最稳的选择——SDM 的状态容量大但有物理上限,且 kernel 优化未到位 - 学到的 \(\mathbf{M}_0\) 这个 idea 值得借鉴——任何"远大于激活但小于参数"的中间状态,都值得考虑学个初值


写在最后

SDM 让我兴奋的地方不是"它超过了 Full Attention"——这种话术听听就好。它的真正价值是证明了 linear RNN 的状态容量不再是个物理死结:用稀疏化换 FLOPs 预算,再把省下来的算力换成大状态。

下一步悬念在两个地方:

  1. kernel 优化能不能把 1.49x 训练速度差压回去? 这个问题决定 SDM 能不能在工业界落地
  2. 更大规模(70B+)上 SDM 还能保持对 FullAttn 的优势吗? 8B 的数据点太少了,可能存在规模外不出来的因素

如果你也在做长上下文架构选型,这篇 paper-analysis 论文值得存到收藏夹——下次有人跟你说"linear RNN 状态不够大没法做长上下文",把 SDM 的 Figure 1 拍他脸上。


参考资料


觉得有启发的话,欢迎点赞、在看、转发。跟进最新AI前沿,关注我。