Gated Recurrent Transformers: Expressive Depth through Recurrent Modulation
Amr Hegazy, Amr Alanwar, Mostafa Elhoushi
cs.CL, cs.LG
2026-08-15
共享核心夹在prelude与coda之间,用逐元素门控在深度上循环。等算力下三层打平十二层GPT-2;大规模参数少62%,编译延迟高10%。
Transformer的深度和参数绑在一起:多一层就多一套权重。共享权重能把有效深度拉长、参数不涨,但同一套变换打在第1步和第8步的隐状态上,功能多样性会被压扁。Kaplan等人早就写过:等参数时循环更好,等算力时循环更差。后来的prelude-coda拆分、Mixture-of-Recursions、带LoRA的Relaxed Recursive Transformer,都在补「共享层在不同深度上该有不同行为」这件事。这篇要回答的是:一层共享核心能不能在深度上表现得像许多层,同时在等FLOPs和等参数两条轴上都站得住。
Gated Recurrent Transformer沿用prelude–共享核心–coda三段:记成npre + nrec × R + ncoda。prelude只跑一遍,产出固定的h^(pre);共享块循环R次;coda把最后状态映到logits。参数量跟R无关。小模型1+1×10+1对应3个独特块、循环10次,前向FLOPs对齐12层GPT-2 Small。
循环步里,当前状态和prelude拼接后经Wproj投影,加上每步重采样的噪声,再进共享块得到提案o^(r)。逐元素门g^(r)由两层MLP从归一化的当前状态和prelude算出,再混入门噪声:h^(r) = g ⊙ h^(r-1) + (1-g) ⊙ o。门偏置初始化成+4,训练开始时g≈0.98,残差几乎原样穿过,模型再学哪些位置该被覆盖。门条件化在「此刻状态、原始输入锚、随机扰动」上,同一套权重在不同步看到不同输入。训练时R在1到最大深度上均匀抽样,中间出口都能当终点,推理早停不需要辅助损失。
训练跟nanoGPT同一套:序列1024、约9.8B token、AdamW,GPT-2 Small/Medium/Large从零训。对照包括MoR、Geiping式重尾深度抽样、RRT、Ouro。
等FLOPs。Small上GRT验证损失3.14,对上稠密基线3.15,参数35M对124M;三颗种子最差3.148,仍好过基线最好的3.154。Medium 2.89对2.84(127M对354M),Large 2.77对2.71(293M对774M),中大尺度标准token预算下稠密仍略好,对MoR和重尾抽样九个尺度×预算格子全部领先。等参数时加深循环:Medium 2.76对2.84,Large 2.65对2.71。
下游九项zero-shot,Large等FLOPs均分42.08对稠密42.05;等参数44.15,+2.10。编译后Large生成延迟+10%,参数少62%,峰值显存少59%(639MB对1570MB)。均匀深度抽样让早停是顺带产物:摘要称只跑一半循环仍留约92%精度。消融里,光循环会让损失差0.107 nats,逐元素门是最大单项贡献,再降0.048。
这是用门控换功能多样性,不是再堆独特层。等FLOPs路线适合显存紧、还想对齐稠密算力的训练;等参数路线适合愿意加推理FLOPs换质量。一个checkpoint能当计算-质量旋钮用。规模停在GPT-2 Large、数据约10B token,还不是现代LLM的结论。对循环深度这条线,门控加prelude锚比LoRA解绑或token路由更省事,也在等FLOPs上把MoR、Ouro拉开了。
作者列了三条:推理深度R固定,没有按token停;门偏置和噪声幅度出了GPT-2家族可能要重调;最优共享比例随尺度变,还没有系统扫。中大尺度标准预算下稠密验证损失仍更好,摘要里「追平/反超」发生在把token预算加倍之后。朴素做法每步一份KV cache,大batch会吃掉参数红利;跨步平均KV能把B=32的显存收到稠密的0.39倍。训练数据配方写的是「diverse text」,没有公开成现代堆栈那种精细配比。