SpiralFormer: Looped Transformers Can Learn Hierarchical Dependencies via Multi-Resolution Recursion
Chengting Yu, Xiaobo Shu, Yadao Wang, Yizhen Zhang, Haoyi Wu, You Wu, Rujiao Long, Ziheng Chen, Yuchi Xu, Wenbo Su, Bo Zheng
cs.LG
2026-02-12
阿里给循环 Transformer 加上粗到细多分辨率日程,共享层先在压缩序列上算全局再精修。1.4B 的 SpiralFormer-L 比同参数 Pythia 少约 7% FLOPs,5-shot 平均准确率高 2.44 点。
MeSH 这类机制已经能让循环 Transformer 在同等算力下追上甚至反超原版,但每一圈仍在全长 token 上做 attention。如果早期圈只需要全局粗依赖、后期才需要局部精修,每圈都跑全分辨率就是在重复付平方级代价。
SpiralFormer 把序列分辨率做成循环内部的一等公民:同一套共享 core,按日程在不同压缩长度上执行,从粗到细把层次化依赖学出来。
骨架还是 Middle-cycle(pre → loop → post)。每圈四步:按当前分辨率把 token 状态压成 chunk 级 latent,共享 core 在短序列上算,再把结果升回 token 级更新,最后做因果右移,交给拓扑更新(Anchor 或 MeSH)。
分辨率默认粗到细、每次翻倍,例如 {1/8, 1/4, 1/2, 1} 或从 1/16 起。Chunk 大小 g = floor(1/r)。下采样默认用每圈一个线性打分器做 softmax 加权聚合,上采样用输出相关的分配向量,增益取 √g,用来对齐不同 chunk 大小的更新幅度。
因果是关键。Chunk 聚合会看到块内「未来」token,所以要把升采样后的更新右移 st 位。默认 st = g−1,这是保证严格因果的最小位移,并在产生更新的块和接收更新的块之间留一个 token 的重叠。Chunk 边界默认半块偏移,避免解码时算力扎堆在固定位置。
两个规格。SpiralFormer-B 和全分辨率 LoopedFormer 共用同一套层分配,只换分辨率日程,参数几乎不变、FLOPs 下降。SpiralFormer-L 对齐非循环 Pythia 的参数量,把中间的全分辨率计算换成粗到细循环。
同样在去重 Pile 上从零预训练 2500 亿 token,序列 4096。下游 9 任务 0-shot/5-shot。
| 1.4B 模型 | 非嵌入参数 | Prefill FLOPs | 0-shot | 5-shot | Pile PPL |
| Pythia 24 层 | 1208.6M | 14.08e12 | 49.50 | 51.93 | 7.44 |
| LoopedFormer+MeSH | 805.8M | 14.08e12 | 50.56 | 52.79 | 7.39 |
| SpiralFormer-B+MeSH | 805.9M | 12.92e12 | 51.48 | 53.22 | 7.30 |
| SpiralFormer-L+MeSH | 1208.8M | 13.13e12 | 51.75 | 54.37 | 7.14 |
B 规格相对全分辨率循环,FLOPs 大约少 7% 到 11%(410M: 4.59→4.11; 1B: 9.67→8.95)。L 规格对齐参数后,1.4B 的 5-shot 比 Pythia 高 2.44 点,比 LoopedFormer+MeSH 高 1.58 点。
410M 消融里,粗到细改成细到粗,Pile PPL 从 9.00 升到 9.24,0-shot 从 44.31 掉到 43.61。MeSH 拓扑优于 Anchor。无重叠的并行位移(st=g)质量下降,但附录指出这条设置方便推理时把低分辨率计算挪出逐 token 关键路径。可学习的下/上采样比均值池化加均匀广播更好。循环层占比呈 U 型,约 30% 到 40% 验证损失最低,两端(完全不循环或过度共享)都更差。
注意力探针(410M,Pile 验证集 500 条):分辨率升高后,key-marginal entropy 下降、Local Attention Mass 上升,粗圈更散、细圈更局部。同一套探针打在全分辨率 LoopedFormer 上,圈间变化弱得多、也更无序。
循环架构先前主要在调「怎么传状态」。这篇把分辨率加成第三根轴:参数深度、计算深度之外,序列可以按圈压缩。对要压参数或压 prefill FLOPs 的预训练,B 规格用少三分之一非嵌入参数打过更大的 Pythia;L 规格用更少 FLOPs 打过同参数原版。和 MeSH 是同一组人的后续工作,两者叠在一起最好,单用 Anchor 也能跑,只是更弱。
这不是现成权重上的免费加速,全部从零预训练。推理要实现 chunk 触发的多分辨率更新和缓存,工程量比普通循环更大。
主实验停在 1.4B,没有 2.8B/6.9B 对照,谈不上已经验证过大规模。并行无重叠设置明确掉点,附录只论证了流水线可行性,没有给出补回质量的方案。注意力探针是相关证据:分辨率变了,注意力统计跟着变,但不能单独证明「模型真的在做层次推理」。没和 Huginn、Ouro 这类更大规模循环模型直接比,也没在指令微调或长上下文设定下测。Chunk 大小由分辨率日程写死,不会按内容自适应切段。