小语言模型极端过训练:0.9M参数模型222K tokens/参数的性能崩塌

1 阅读

极端过训练:当0.9M参数模型遭遇222K tokens/参数

在语言模型训练领域,Chinchilla最优比例(约20 tokens/参数)被视为大模型的黄金法则。然而,对于小模型(参数少于10M),这一法则早已被实践推翻。社区中广泛使用的比例往往高出1-3个数量级。但究竟多高才算安全?本文记录了一次极端实验:将0.9M参数模型训练至222K tokens/参数,结果令人警醒。

实验设计:挑战极限

我们决定测试小模型的极限,设定了一个激进的目标:训练0.9M参数模型至200B tokens,相当于每个参数222K tokens,是Chinchilla比例的220倍。

模型架构

  • 层数:6层
  • 隐藏维度:96
  • 中间层:SwiGLU,380维
  • 注意力:GQA,6查询头/2键值头
  • 词表大小:384
  • 上下文长度:8K

优化器与数据

  • 优化器:Muon(峰值学习率7e-2)用于2D参数,AdamW(峰值学习率4e-3)用于其余参数
  • 数据:FineWeb-HQ + Cosmopedia v2
  • 总预算:200B tokens

选择如此高比例的理由是:词表极小(384),每个token携带的信息量少于常规32K词表,因此模型应能吸收更多原始token。此外,小词表释放了嵌入参数,为Transformer层留出更多预算。

结果:性能的抛物线轨迹

训练过程中,我们在标准Open SLM基准(ARC-Easy, ARC-Challenge, HellaSwag, PIQA)上评估模型,并聚合为INT Index分数。

INT Index regression across training

关键发现

INT Index在20B tokens时达到峰值4.55,随后单调下降(含噪声),至180B tokens时降至3.31,降幅达27.3%。

分项对比

基准 20B (10%) 180B (90%) 变化
ARC-Easy 26.98 28.32 +1.34
PIQA 53.54 52.07 -1.47
ARC-Challenge 22.27 21.25 -1.02
HellaSwag 29.01 28.06 -0.95
平均 32.95 32.43 -0.52

四个基准中三个在180B时比20B时更差,仅ARC-Easy略有提升,但提升幅度小于其他三个的下降。

排除学习率衰减干扰

为排除最后10%学习率衰减的影响,我们评估了80%检查点(约160B tokens),其平均分为32.68,更接近90%最终值而非40%峰值(33.44)。这表明性能下降并非调度伪影,而是真实的过训练。

Chinchilla对照:欠训练的另一端

作为对照,我们以Chinchilla比例(20:1)训练相同架构,仅18M tokens,结果INT Index仅1.53,接近随机水平(ARC-E 26.64, PIQA 49.78, ARC-C 26.54, HS 24.88)。

这界定了有用范围:

  • 20:1 (18M tokens):INT Index 1.53,欠训练,随机水平
  • 22K:1 (20B tokens):INT Index 4.55,峰值
  • 200K:1 (180B tokens):INT Index 3.31,过训练,低于峰值27%

有用训练比例范围

好消息是,常规的高过训练比例(约7K:1、15K:1、22K:1)表现良好。社区模型在这些比例下通常产生健康的缩放曲线,如TinyStories约10K:1,许多子3M模型在15-30K:1。

失败模式出现在超出该范围之后。完整轨迹如下:

Tokens 比例 INT Index
18M 20:1 1.53 (随机)
10B 11K:1 4.01
20B 22K:1 4.55 (峰值)
40B 44K:1 4.12
80B 88K:1 4.13
160B 175K:1 3.82
180B 200K:1 3.31

从Chinchilla最优(随机水平)到峰值的上升非常快(10B到20B之间每B tokens约0.05 INT Index),而从峰值下降则较慢(每B约0.008),但持续不断。20B之后的每个检查点都比20B差,趋势线直指下方。

实践建议

基于本次实验,对于Pico级(约1M参数)模型,有用的计算最优范围约为22K-30K tokens/参数。低于此范围(低至几K:1),模型仍在提升;高于此范围,基准开始恶化。

如果你计划训练小模型,实用建议是:从小预算开始(20K-30K tokens/参数),评估,再决定是否扩展。盲目扩展到200K:1会浪费大量GPU时间,最终模型比20K:1运行更差。

深入分析:为什么过训练会损害性能?

过训练导致性能下降的机制可能包括:

  1. 记忆过度:模型开始记忆训练数据中的噪声,而非学习通用模式。
  2. 表征坍缩:随着训练持续,模型内部表征可能变得过于专门化,失去泛化能力。
  3. 优化器行为:长时间训练可能导致优化器步长过大,在损失平面上震荡。

然而,小模型为何能承受比大模型高得多的比例?可能因为小模型容量有限,更容易达到过拟合点,但过拟合的后果也更严重。

未来方向

本研究仅针对单一架构和数据集。未来工作可探索:

  • 不同架构(如更深但更窄)对过训练敏感性的影响
  • 数据质量与过训练的关系
  • 动态调整训练比例的策略

总之,小模型训练并非“越多越好”,找到合适的训练比例至关重要。希望本文能为社区提供参考,避免重蹈覆辙。