训练
- 实测
- GPT-2 small, medium, large 298 个词元
- 拟合
- Chinchilla Hoffmann 等 2022
- 算力
- C ≈ 6 N D
全部步骤
模型越大,损失越低
GPT-2 small、medium 和 large 用同样的方式在同样的数据上训练。在本站英文术语表的 298 个词元上(它们都没见过),随着规模增长 8 倍,损失依次降为 4.36 → 4.19 → 4.00。在对数刻度上,三个点几乎在一条直线上:这就是幂律。
这三个点上 L ∝ N^−0.040哪些词元变容易了
平均值掩盖了提升从何而来。大多数词元变化不大;少数需要知识或更长上下文的词元,对大模型来说容易了很多。
逐词元的 loss(small) − loss(large)参数与数据一起看
损失同时取决于参数量 N 和训练词元数 D。Chinchilla(Hoffmann 等 2022)用 L = E + A/N^α + B/D^β 拟合了数百次运行。每一项都随各自的预算增加而缩小,直到 E,即任何模型都消除不了的损失。
L(N, D) = 1.69 + 406/N^0.34 + 411/D^0.28如何花算力预算
训练算力约为 6 N D FLOPs。预算固定时,模型越大,见到的词元越少,损失在两者之间某处最低。按照拟合,最佳规模大约随 √C 增长,大约每个参数 20 个词元。
C ≈ 6 N D · 最佳点在 D ≈ 20 N 附近真实模型的位置
按这条规则,GPT-3 规模很大却训练不足。70B 的 Chinchilla 在同样预算下胜过了 280B 的 Gopher。LLaMA 3 8B 见过 15T 个词元,约每参数 1,900 个:故意越过了计算最优点,因为小模型推理更便宜。
每参数词元数:1.7 → 20 → 1,900
代码
E, A, B, alpha, beta = 1.69, 406.4, 410.7, 0.34, 0.28 # Chinchilla, approach 3
def loss(N, D): return E + A / N**alpha + B / D**beta
def flops(N, D): return 6 * N * D # forward + backward, per token
C = 1e23
N_best = min(np.logspace(8, 12, 2000), key=lambda N: loss(N, C / (6 * N)))
D_best = C / (6 * N_best) # about 20 tokens per parameter