Matryoshka Language Model Suites
Nathan Godey, Yoav Artzi
cs.AI, cs.CL
2026-08-10
Cornell把500M、1.5B、3B嵌进一次端到端训练,七项基准平均准确率与独立训练相差不到0.5点,总算力少36%,500M草稿的投机解码吞吐提高14%–26%。
语言模型很少只发一个尺寸。厂商通常要同时交出500M、1.5B、3B这种套件,方便不同显存和延迟档位。现行做法是每个尺寸单独预训练,小模型再另外蒸一遍大模型,训练账单接近各尺寸相加。投机解码还要再挂一个独立草稿模型,KV缓存各算各的。
Godey和Artzi把套件当成一个嵌套结构来训:小模型的参数是大模型的真子集,一次前向同时出所有尺寸的logits。目标是在下游质量和困惑度不掉的前提下,把套件总训练算力压下去,并让草稿模型「长在」验证模型里面。
三个子模型宽度递增、深度各自独立,参数严格包含:θ1⊂θ2⊂θ3。和早期退出不同,这里每个出口有自己的隐层宽度和LM头,拆下来就是一份可单独部署的checkpoint。主实验做成500M / 1.5B / 3B,深度分配(24, 10, 5)、总层数39,用来对齐独立3B的KV缓存和每token FLOPs。套件总参数3.2B,独立训练三份合计5.2B,少38%。
宽度不同,小模型的输出不能直接喂给下一截。他们的接合器不加新参数:把第m个出口的输出按L2范数缩放到和新嵌入一致,再和新嵌入在通道维拼接,送进下一段Transformer。消融里拿掉范数对齐,平均PPL大约差0.2;把新嵌入换成全零,再差到0.54。
蒸馏几乎是白送的。最大模型的softmax当教师,对每个小出口做token级交叉熵,系数αd=0.3,和小模型自己的语言模型损失凸组合。离线蒸馏要么存logits要么并行跑教师,这里一次前向全有。
投机解码时,500M草稿和3B验证共享前24层的KV。草稿侧算过的缓存验证器直接复用,只跑多出来的层。
训练数据是FineWeb-Edu的350亿token,序列长度2048,AdamW,蒸馏系数来自200M代理扫描。独立基线用同数据同超参的Llama风格模型,并额外给出算力对齐的230亿token冷启动点。
token对齐(同为350亿)时,七项零样本平均准确率每个尺寸都在0.5点以内:500M是46.5对47.0,1.5B是52.2对52.4,3B是53.3对53.5。算力对齐(独立套件只训到230亿token)时,Matryoshka每个尺寸反超0.4到1.9点。域外字节PPL在1.5B和3B上优于token对齐的独立模型(2.121对2.139,2.067对2.097),500M打平。验证PPL全程贴着独立曲线,尺寸间差距在±1.4%以内。
对齐来自共享权重和在线蒸馏。训练中成对KL更低,top-1下一token一致率更高,1.5B与3B这对最终高5.7%。A100上500M草稿、3B验证、草稿长度6:贪心吞吐2650 token/s对独立套件的2100(+26%);nucleus采样下吞吐仍高约14%,接受长度平均高5%。独立的1:6草稿在nucleus下甚至慢于普通自回归;嵌套之后同一比例开始赚速度。共享缓存还把最大batch从64抬到102。
200M尺度上对比MatFormer:MatFormer只在FFN宽度上嵌套,所有粒度KV一样大;Matryoshka沿深度嵌套,小模型KV按尺寸缩。同等验证PPL约21时,Matryoshka-100M对上MatFormer-M(139M),参数和KV都更小。
这是给「要发一套尺寸」的预训练团队看的工程论文。一份3B预算几乎把500M和1.5B当副产品带出来,还附赠对齐更好的草稿模型。对投机解码尤其直接:草稿不再是另起炉灶的小模型,是验证器前缀,KV天然共享,1:6这种文献里通常亏本的比例开始划算。
它不是把一个模型切成任意宽度的旋钮。出口尺寸在训练时就钉死,能拆下来的checkpoint是离散的几档,不是MatFormer那种组合粒度。换来的是每档都有独立深度和独立KV,能按真实部署档位卸下来。
主实验停在3B和350亿token,离现在的70B套件还隔着数量级。损失在子模型间目前是均匀加和,作者在200M上发现调权重能补掉不少残差,3B套件没重扫。没有指令微调、对齐或长推理后训练,嵌套结构过了SFT会不会裂开,不知道。投机解码只报了500M/3B一对,没测1.5B当草稿。接合器的范数对齐是训练稳定性补丁,换归一化或可学习投影会不会更好,只有两个否定对照。数据是FineWeb-Edu,教育向网页,泛化到代码或多语还没测。