SUT 把稀疏专家装进共享层,WMT 用一半参数打平 210M 基线

2026-09-04

Mila 与 MIT-IBM 把稀疏专家和 stick-breaking 停层装进 Universal Transformer。WMT'14 英德上 66M 的 SUT 打到 29.2 BLEU,接近 210M Transformer 的 29.3;逻辑推理长度外推远强于普通 Transformer,组合泛化最难一档仍过不去。

这篇在解决什么

Vanilla Transformer(VT)每层一套独立参数。Universal Transformer(UT)把同一套块在每一层重复用。层间共享有两个已经测过的好处:参数更省,组合泛化更好。Csordás 等人在形式语言任务上看到,共享之后学到的运算不再绑死在固定深度顺序上,测试时遇到没见过的运算次序还能套用。有限深度的普通 Transformer 会被理论结果卡住表达力,UT 在层数不受限时甚至是图灵完备的。

代价按层数平方涨。L 层、每层 P 参数的 VT,一次前向大约 LP。参数总量对齐的 UT 要把那 LP 参数全塞进一个块,再跑 L 次,复杂度大约是 L²P。Takase 和 Kiyono 在 WMT 英德上测过:同样设置下 UT 训练时间大约是 VT 的两倍,显存也明显更高。Kaplan 等人的缩放曲线写过同一件事:共享参数在「参数换性能」上更划算,在「算力换性能」上更亏。目标是留住 UT 的参数效率,把平方级算力砍掉。

方法

Sparse Universal Transformer(SUT)还是层间共用一个块。块里面两处改成按 token 条件激活:

为什么注意力也要稀疏:UT 把容量做大时,注意力头和 FFN 一起膨胀,只稀疏 FFN 仍会把算力留在注意力上。MoMHA 把「多头」拆成可路由的专家,k 固定时加专家不加每次前向的乘加次数。

路由用 Mutual Information Maximization(MIM)辅助损失压负载。抬高专家被选中的边缘熵,让大家都有活干;压低条件熵,让单次路由别糊成均匀分布。

动态停层改成 stick-breaking 过程。原版 UT 用凸组合,当前层输出按 α 和上一层掺在一起,多层连乘 (1-α) 容易把梯度掐死。SUT 把停层看成概率分布,直接算「期望已停状态」:还没停的概率乘当前层隐状态,加上已经在更浅层停过的加权和。注意力的 query 用未停状态,key/value 用期望已停状态。另加 Adaptive Computation Time 损失,惩罚期望层数,把模型往少用层的方向推。训练时阈值 αthresh = 0.999,累积停层概率过线就把这个位置路由到空操作专家,这一层真的不再算。阈值训练完还能再拧,用来换推理步数。

结果

WMT'14 英德翻译,参数量对齐看:

模型参数BLEUMACs
Transformer base65M27.3604M
UT65M28.9未给出
UT base + stick-breaking64M29.31998M
SUT base66M29.2787M
Transformer big210M29.32090M
SUT big110M29.4787M
UT big + stick-breaking105M29.63707M

SUT base 用 66M 参数、787M MACs,打到 Transformer big(210M、2090M)的 29.3 BLEU 附近。同容量密 UT 仍高 0.1–0.2 BLEU,但 MACs 大约是密 UT 的四成到两成。top-k 固定,SUT 从 66M 扩到 110M,MACs 仍是 787M。Admin 60L-12L 以 256M 参数打到 30.1,SUT 没有越过这条线。

消融从 SUT base 的 29.2 BLEU 往下拆:去掉 MIM 掉到 28.9,把 MoMHA 换成普通 MoA 掉到 28.7,去掉 ACT 损失 29.0,去掉停层 29.1。翻译任务上,多头 MoE 注意力和 MIM 比停层更值钱。FFN 专家和词的共现能看出分桶,有的专家爱限定词,有的爱代词、后缀或名词,模块化有苗头,离「一个专家管一步算法」还远。

CFQ 把自然语言译成 SPARQL,用 compound divergence 衡量训练/测试在 token 组合上差多远。超参搜索把专家数收到 E=1,等于退回密 UT。无预训练时,带停层的 UT 三个 MCD 切分平均 58.4,对上 Bergen 等人的 T5-based UT 21.3、Keysers 的 Transformer 21.4。接上 T5 相对位置和 pre-LN 之后,T5-based UT 复现到 52.3,完整 UT 再抬到 58.4。预训练的 Dangle 平均 66.1,但 MACs 到 51033M,解码每个 token 都要重跑编码器。

逻辑推理任务训练只见 0–6 个算子,测 7–12 个。SUT 在 7 到 12 算子上分别是 98、97、94、90、88、81;LSTM 是 88、84、80、78、71、69;普通 Transformer 卡在 51 附近。组合泛化 A/B/C 三档由易到难,SUT 是 97/94/52,LSTM 80/60/59,Transformer 53/51/51。A、B 两档递归加共享层很管用,C 档 SUT 52,还低于 LSTM 的 59。平均停层深度随算子数上升,难样本会多想几层。

训练后拧低停层阈值:形式语言任务大约能省一半推理步,精度几乎不动;CFQ 在 0.8 阈值下跳过约 33% 步,0.8 到 0.999 之间准确率几乎持平;WMT 翻译模型停得晚,只省约 9% 还能保住 29.1 BLEU。

为什么重要

层共享如果是被 L²P 劝退的,这条补丁可以跑:稀疏专家把参数量和每次前向的乘加拆开,k 固定时加专家不加 MACs。110M 的 SUT 在 WMT 上跟 200M 级稠密模型打平,这是参数效率,不是新的翻译 SOTA。

形式语言这边,递归加共享层的归纳偏置是真的。长度外推和 A/B 组合切分上,SUT 明显强过 VT。C 档过不去,「共享层 + 停层」还没覆盖组合泛化的全部坑。停层阈值可以训练后再拧,部署时有一个不重训的算力旋钮;翻译任务上这个旋钮拧不动多少。

局限与存疑

论文自己承认:组合泛化有一部分 SUT 解不掉,更大规模有没有同样结论还没做。

另外几处对不上号。CFQ 上最好的配置就是密 UT,稀疏专家在这个小任务上帮不上忙,组合泛化的主数字其实是 UT 的,不是 SUT 的。WMT 上 SUT 始终略低于同容量密 UT,换来的是 MACs。相对 VT base,SUT base 的 MACs 是 787M 对 604M,并没有比普通 Transformer 更便宜,只是比密 UT 便宜。翻译上早停只省 9%。整篇最大模型 110M,十亿、百亿参数档会不会同样成立,没有证据。专家-词共现像词性分桶,模块化证据偏弱。

术语

原文与代码

社区讨论

全部论文解读