Training

Scaling laws

loss against parameters, data and compute

measured
GPT-2 small, medium, large 298 tokens
fit
Chinchilla Hoffmann et al. 2022
compute
C ≈ 6 N D

All steps

  1. Bigger model, lower loss

    GPT-2 small, medium and large were trained the same way on the same data. On 298 tokens of this site’s glossary, which none of them saw, their loss falls 4.36 → 4.19 → 4.00 as size grows 8×. On a log scale the three points sit close to a straight line: a power law.

    L ∝ N^−0.040 on these three

  2. Which tokens got easier

    The average hides where the gain comes from. Most tokens change little; a few that need knowledge or longer context get much easier for the large model.

    loss(small) − loss(large), per token

  3. Parameters and data together

    Loss depends on both the number of parameters N and training tokens D. Chinchilla (Hoffmann et al. 2022) fit L = E + A/N^α + B/D^β to hundreds of runs. Each term shrinks as its budget grows, down to E, the loss no model removes.

    L(N, D) = 1.69 + 406/N^0.34 + 411/D^0.28

  4. Spending a compute budget

    Training compute is about 6 N D FLOPs. For a fixed budget a bigger model sees fewer tokens, and loss is lowest in between. By the fit, the best size grows about as √C, with roughly 20 tokens per parameter.

    C ≈ 6 N D · best near D ≈ 20 N

  5. Where real models sit

    GPT-3 was large and under-trained by that rule. Chinchilla, at 70B, beat the 280B Gopher on the same budget. LLaMA 3 8B saw 15T tokens, about 1,900 per parameter: past the compute-optimal point on purpose, since a small model is cheaper to serve.

    tokens per parameter: 1.7 → 20 → 1,900

Code

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

Go deeper