参数、算力、KV 缓存三样一分不许多花,让中间一半的层再跑一遍——省下 15% 训练算力,而第二遍干的是把注意力从句首那个黑洞里拽回来
SMELT: Scaling Laws for Compute-Matched MoE Looped Transformers
9 月 1 日 arXiv 上挂出一篇 35 页的论文,作者来自清华大学、字节跳动 Seed、开源研究社区 M-A-P 和 TokenWave.AI。它盯的是一个 2019 年就被提出、最近两年又火起来的老想法:与其把 Transformer 越堆越深,不如让同一批层跑两遍。这类模型叫循环 Transformer(Looped Transformer)——权重只存一份,前向传播时反复经过它,等于用同样的参数换来更长的串行计算链。过去两年不断有论文报告,循环模型能在算术、多跳归纳、数学这些任务上追平甚至超过比它大好几倍的普通模型。
这篇论文做的事,是把那句话按住重新验一遍,结论是:优势是真的,但远没有大家以为的那么大。问题出在比较方式上——把一个 12 层模型循环成 24 层执行,权重确实只存了一半,可它每个 token 花掉的算力、需要的 KV 缓存,跟一个真正的 24 层模型没什么区别。省的是显存,不是计算。作者把三样预算同时卡死之后重新比,循环剩下的架构红利是训练算力省 6.8% 到 18%:不再是以小博大的神话,但是一份真实、可复现、而且附了机制解释的收益。
「省参数」不等于「省算力」
先说清楚这三样为什么都得卡死:总参数决定模型能记住多少知识;每 token FLOPs 决定训练和推理要花多少钱;KV 缓存(推理时缓存下来的键、值向量,占显存的大头)决定线上能服务多长的上下文。以前的循环类论文通常只固定其中一样——固定存储的参数量去加循环次数,或者拿它跟更大的非共享模型比参数效率。两种做法都把「架构本身更好」和「偷偷多花了算力」混在了一起。今年有一篇工作固定了 FLOPs,算出 r 次循环大约只相当于 r 的 0.46 次方个独立层;但固定 FLOPs 又会顺带砍掉循环模型的参数量,那它赢不了也可能是因为被砍笨了,而不是循环这件事没用。
让三样同时对齐是这篇能做成的关键,而这件事之所以做得成,靠的是 MoE(混合专家:每层备一大堆专家网络,每个 token 只激活其中几个,于是总参数量和单 token 算力可以脱钩)。循环多跑的那几层要花 FLOPs,就把隐藏维度收窄还回去;收窄之后总参数少了,就多加专家补回来;多执行几层会撑大 KV 缓存,就调小注意力头维度、提高 GQA 分组比。论文里给了个实例:200M 那一档,基线是 12 层、隐藏维 1280、每层 192 个专家;对齐后的循环模型把隐藏维压到 1056、专家加到 288,中间 6 层跑两遍、总共执行 18 层。最后每 token FLOPs 差 2.9%、总参数差 0.4%、KV 缓存差不到 4%。三样都卡在误差之内,剩下的差异才好归给架构。
三次消融,锁出一个配方
接下来是三个问题,全在 200M 规模上扫。第一,循环哪些层:从完全不循环一路扫到整摞层全循环,验证损失在「中间一半」处最低——把整个模型都循环反而不如只循环中间六层。第二,深宽比:循环模型偏好比基线更大的有效深宽比,最优解是物理 12 层、执行 18 层。第三,循环几次:两次最好,三次和四次都退步,因为预算是死的,多跑一遍就得把模型压得更瘦,瘦下去的损失盖过了深度的收益。
三条规则拼起来就是论文的名字 SMELT——Sparse MoE Transformer, middle layers Loop Twice。值得注意的是这个配方有多保守:没有自适应停机,没有按 token 决定循环几次,没有给每一遍配不同的低秩适配器,就是老老实实把中间一半的层原样再跑一遍。这些更花哨的变体作者全列在了未来工作里,一个都没做——在一篇要下缩放律结论的论文里,这个取舍是对的。
省下的那 15% 是怎么量出来的
验证部分是全文最扎实的地方。四个规模(激活参数 100M 到 1.6B,最大一档的总非嵌入参数 540 亿)乘四档稀疏度,32 次预训练、96 组预算对齐的对照、192 个评测端点,单次训练约 2150 亿 token 且不重复数据,两种架构喂的是完全相同的 token 序列。然后给基线和 SMELT 各自拟合一条 Chinchilla 式的缩放律,再反过来问:同一个损失值,两边各要花多少算力。
拟合出的指数是全文的核心数字:SMELT 的容量指数 0.3892、数据指数 0.7011,基线是 0.3703 和 0.6594,两个都更大,合成的前沿指数从 0.237 抬到 0.250,高 5.5%。翻成账单:10²⁰ FLOPs 的预算下省 6.8%–10.0%,10²¹ 下省 14.7%–18.0%。差距是随算力放大的,这比省下的绝对值更重要。但也要看清边界——拟合窗口只到 2.2×10²¹,再往上那一行作者自己打了外推标记,10²² 的自助法置信区间下界甚至压到了 0。
下游表现比验证损失更好看:DCLM Completion 上 96 组对照全胜,DCLM Core 83 胜 13 负。而且增益不是均匀摊开的——按数据类别拆,代码 20.4% 最高,金融 16.8%、数学与 STEM 16.6%,知识和网页文本 14.9%、14.8% 垫底;按文档长度拆,最长的四档增益是最短四档的 1.52 倍,而单纯加参数、加专家的对照组没有这个倾斜;按上下文示例拆,零样本时两边只差 0.9 个百分点,一给示例就拉到 1.9 个百分点。越有结构、越长、越依赖回头去翻前文的任务,多跑那一遍越值。
第二遍不是重算,是回头看
论文最好看的部分在后面:他们把第二遍究竟在干什么拆开看了。其一,专家路由没有简单重复——在大专家池下,两遍选中的 top-8 专家只重叠 2 到 3 个,但远高于随机路由的期望,说明模型固定复用一小撮核心专家,其余的换掉。其二,第二遍写得更重:残差流的更新范数在全部 16 个格子里都比第一遍大,倍数落在 1.2 到 3.5 之间;而且两次写入方向是正相关的(同一物理层跨两遍的余弦 0.56,不匹配的层对只有 0.16)——第二遍不是推翻第一遍,是把第一遍立起来的信号加粗。
其三是我认为最漂亮的一条:Q 和 K 的跨遍余弦相似度稳在 0.89–0.93,V 却掉到 0.65–0.74。也就是说,第二遍看的还是同一批位置,只是从那些位置读回来的东西变了。循环因此更像一次「精修」,而不是简单地加容量。
括号匹配(Dyck 语言)的案例把这件事看得最清楚。这里要先解释一个术语:注意力 sink,指模型习惯把大量注意力权重堆在句首这类没有实际内容的位置上,像个泄压阀,而在标准 Transformer 里它还会越往深处越严重。在这个任务上,第一遍有 0.60 的注意力质量压在句首 BOS 上,示例答案只分到 0.24;第二遍句首掉到 0.02,示例答案涨到 0.85——sink 让出去的,几乎不多不少正好是示例拿到的。作者在 100 万 token 的通用留出集上复测,结论同样成立,而且是反着「越深 sink 越大」的常规趋势走的。这是全篇唯一一处把宏观的百分比和微观的机制接上的地方。
本文为 AKL AI Club 原创撰写的导读,不是原文翻译;著作权归原文作者所有。 篇目由编辑独立选取,来源均经人工核实。