CS336 // FIELD MAP
← field map
LECTURE 11 · SCALE IT AND SERVE ITTatsunori Hashimoto · 2025-05-06 · 78 min

Scaling laws 2

Stanford CS336 · Spring 2025 · lecture 11 of 17

Transcript: cleaned auto-captions with timestamps

TL;DR — Lecture 9 taught you the shape of a scaling law; this one asks whether anybody actually trains a model that way. The answer, read off the three teams who published enough detail to check — Cerebras-GPT, MiniCPM, DeepSeek LLM — is that the isoFLOP part of Chinchilla replicates beautifully and everything else is a mess you have to engineer around. Two engineering moves carry the whole lecture: muP, a parameterisation that makes the optimal learning rate stop moving as you widen the model, and WSD, a trapezoid learning-rate schedule that turns an O(n²) Chinchilla sweep into roughly one training run plus a few cheap re-decays. The number to un-learn is 20 tokens per parameter: the same isoFLOP machinery that produced it now produces 39, 96 and 192 depending on whose architecture and data you use.

CS336's scaling arc has a credibility problem, and Tatsu opens by naming it: Chinchilla is curve-fitting on a log-log plot, and the labs that would know whether it survives contact with a real 70B run stopped publishing right after ChatGPT. So this lecture is a forensics exercise. Take the handful of model releases that documented their scaling process honestly, reverse-engineer the recipe each one used, and see which pieces recur. What recurs is the practitioner's actual job description: pick an architecture aspect ratio, pick a learning rate and batch size, pick a token budget — and find a way to do all three at 1/100th of the scale you intend to train at.

Outline, with timestamps

The evidence problem

The most recent model with a genuinely detailed public scaling study is, depending on how you count, still from 2022 or 2024. That is the whole framing of the lecture (01:45). Chinchilla is 2022. After that the competitive landscape changed and scaling methodology became the thing nobody talks about — Tatsu says he has asked people at frontier labs directly.

"And they're like, no, we will not tell you anything about what we do for scaling."— Tatsunori Hashimoto, 01:45

So the corpus of usable evidence is three releases that documented themselves properly — Cerebras-GPT (2023), MiniCPM (2024) and DeepSeek LLM (2024) — plus three newer ones that show a single isoFLOP plot and little else (Llama 3, Hunyuan-Large, MiniMax-01). Worth noticing what this implies about the field: the most careful published scaling work of the last three years came from a hardware vendor proving out its wafer-scale cluster and from two Chinese labs. Tatsu notes he used to have to justify covering the latter and no longer does (02:20).

Recipe one — Cerebras-GPT: buy invariance, then sweep tiny

Cerebras-GPT is a 0.1B–13B family trained on the Chinchilla ratio, and its headline claim is methodological rather than about the models: parameterise with muP and the scaling curve stops wobbling (05:36). Under standard parameterisation (SP) their measured losses oscillate around the predicted power law, because the learning rate has to be re-chosen at every scale and they never quite nail it; under muP the points sit on the fitted line, and the family scales at least as well as Pythia or GPT-J on Pile test loss.

That predictability is what makes the second move affordable. Because muP is supposed to hold hyperparameters fixed across width, they run their hyperparameter search on a 40M-parameter proxy — extremely aggressive, two to three orders of magnitude below the target — and transfer the winners upward (09:20). Tatsu flags this as the recurring shape of every recipe in the lecture: train a surrogate, then find a principled way to carry its answers up. He is openly unsure whether 40M is far enough down to still be informative for a frontier run.

The practical artefact worth stealing is the appendix table (08:48): a side-by-side of SP and muP for every layer type, precise enough to implement from. Its one-line summary — every non-embedding parameter initialised at scale 1/width, and per-layer learning rates also scaled by 1/width — is the whole of muP for an Adam user, and the second half is the part people miss. If you already use Kaiming init you are halfway there without knowing it; a global constant learning rate is the thing that has to go. Meta's Llama 4 uses a variant they call MetaP, which sets per-layer learning rates and init scales the same way (07:12).

Recipe two — MiniCPM: the schedule that made Chinchilla cheap

MiniCPM's goal inverts Chinchilla: spend a lot of compute to get an unusually good small model, which means deliberately over-training and therefore needing a defensible answer to "how many tokens is too many." Its 1.2B and 2.4B models beat the 2B class and matched many 7B models of the day (12:04). Getting there took three separate fits.

Batch size. They rerun the Kaplan critical-batch-size analysis: train a grid of (model size, batch size, token count), fit a quadratic down each vertical slice to find the loss-minimising batch, then plot that optimum against terminal loss (14:22). It comes out clean and log-linear, which gives you a usable two-step procedure: predict your target loss from a scaling law, then read the batch size off the loss. The logic is worth stating plainly because it is easy to miss — batch size is being predicted from loss, not from parameter count.

Learning rate. This is where muP is supposed to earn its keep, and MiniCPM's plot is the cleanest public evidence in the lecture: sweep learning rates across model sizes spanning orders of magnitude and the minimum sits at roughly 1e-2 the whole way, inside a wide basin with a sharp instability cliff on the right (17:08). Their constants land close to Cerebras-GPT's independently — scale_emb 12 vs 10, lr 1e-2 vs 6e-3 — which is itself a mild replication.

Token budget, and the cosine trap. To fit a Chinchilla law you need loss at many (size, tokens) pairs, and the obvious shortcut is to take intermediate checkpoints from one long run. That shortcut is wrong, and the reason is the schedule, not the model: a cosine sized for 100B tokens has a completely different learning rate at token 20B than a cosine sized for 20B tokens does at its end (19:19). Comparing them measures the schedule mismatch, not the data. Doing it correctly means a fresh run per target, which is what turns the sweep quadratic.

WSD — warmup, stable, decay — dissolves this. Warm up as usual, hold the learning rate flat for the bulk of training, then decay hard over the last ~10% (20:57). Because the stable phase is flat, it is shared by every token budget: run once to the longest target, then rewind to any earlier checkpoint and pay only for a short decay to get a properly-terminated model at that budget. One long run plus k cheap tails replaces k full runs. The loss curves look alarming — flat-ish through the stable phase, then a cliff — but at every token count the WSD minimum matches or beats cosine (23:37). Tatsu's aside here is the most interesting unresolved thing in the lecture: nearly all the gain arrives during cooldown, and nobody really knows why the optimiser needs a long high-learning-rate excursion first (30:42).

With the fits in hand, MiniCPM runs Chinchilla method 1 (lower envelope of the loss curves) and method 3 (joint fit of the two-variable law), and method 3 hands them 192 tokens per parameter (27:24). Tatsu treats that specific number as an outlier he does not trust — no one else has reproduced anything near it — but insists on the direction it points.

"20 times model size is just a starting point… feeling free to significantly increase that token to parameter ratio."— Tatsunori Hashimoto, 27:56

Recipe three — DeepSeek LLM: skip muP, fit the hyperparameters

DeepSeek's 7B/67B release took the opposite bet (32:53): assume most transformer hyperparameters are scale-invariant, use no muP, and instead fit explicit scaling laws for the two that clearly are not — batch size and learning rate. Grid over both at several compute scales, mark the near-optimal region (they use within 0.25% of the minimum loss), and extrapolate the resulting line out to the real run.

Tatsu is candid that this is the weakest link in an otherwise excellent paper. The batch-size trend is believable; the learning-rate trend is a line through a cloud.

"I could have probably fit a horizontal line and that would have also looked okay."— Tatsunori Hashimoto, 35:02

The rest of DeepSeek's recipe is textbook and lands well. They also use a WSD-style schedule, in a two-stage form: warmup, stable, then two decay steps of about 10% each, spending roughly 20% of the compute budget on cooldown, and they check it matches cosine (35:33). Their token-budget analysis is Chinchilla method 2 — plain isoFLOP quadratics, minima joined by a line — and it fits cleanly. The generalisation Tatsu draws from all three case studies is a genuinely useful prior on how much to trust a scaling plot: hyperparameter scaling laws always look noisy and tenuous; isoFLOP analyses always look beautiful (36:39). Their payoff is the predictable-scaling plot every lab wants: fit at roughly 10²⁰ FLOPs, predict the loss of the 7B and 67B runs at ~10²⁴, and hit it (37:45).

What the last year adds — and the ratio that will not sit still

Since DeepSeek, the public detail thins out sharply; DeepSeek's own v2 and v3 papers foreground MLA and low-precision systems work and report no new scaling studies (39:24). What is left is three data points, and they are all about the same quantity:

Chinchilla (2022)≈20 tokens / parameterthe origin of the rule of thumb
Llama 3 (2024)≈39–40 : 1isoFLOP redone at scale (40:30)
Hunyuan-Large (2024)≈96 : 1 per active parameterisoFLOP for an MoE (43:14)
MiniCPM (2024)≈192 : 1joint fit; treat as an upper outlier

Read the spread the right way. What replicates across every one of these is the procedure — sweep isoFLOP, fit quadratics, join the minima — not the constant it returns. The constant absorbs architecture, data quality, optimiser tuning and MoE sparsity all at once, so a single number quoted without its setup is close to meaningless. And there is a serving-economics thumb on the scale that the compute-optimal framing ignores entirely: everyone would rather ship a smaller model trained longer, because that is the one that is cheap to serve (43:46).

Two other uses of scaling laws show up in passing and both are worth knowing. Llama 3 fits a two-stage chain from compute → negative log-likelihood → downstream benchmark accuracy, sigmoid-shaped, and uses it to predict 405B benchmark numbers from small runs plus Llama 2 points (41:33) — the answer to "we don't actually care about log-loss." MiniMax-01 uses scaling curves as an architecture decision procedure: fit Chinchilla method 1 for softmax attention, for their linear "lightning attention", and for the hybrid, and ship the hybrid once the curves overlap (44:55). Tatsu notes that plot is standard in linear-attention papers, but rare at production scale.

An alternative to WSD gets a nod too: Gadre et al. show the loss penalty for over-training past the Chinchilla ratio is itself predictable, so you can fit the penalty at small scale and extrapolate rather than re-decaying checkpoints (24:09). Tatsu has not seen it used for a large run, but it is the same problem attacked from the other side.

muP from first principles

The second half asks what muP actually is, because most write-ups only give the recipe. The whole thing follows from two demands on a network of width n, holding depth fixed (50:58):

Violate either and widening the model makes activations blow up or vanish, at init or after the first step. Tatsu points out the intellectual lineage explicitly: this is renormalisation-group thinking imported into deep learning — take a limit, insist the observable quantities stay finite, and read off what the couplings must be (65:53).

A1 gives the initialisation. Take a deep linear network hl = Wlhl−1 with Wl ~ N(0, σl). Random matrix theory says its operator norm concentrates at σl(√nl + √nl−1), and at init Wl is independent of hl−1, so ‖hl‖ ≈ ‖Wl‖*‖hl−1‖. Choose σ_l = (1/√n_{l−1})·min(1, √(n_l/n_{l−1})) and the induction closes: ‖hl−1‖ = Θ(√nl−1) implies ‖hl‖ = √nl + lower order (52:37). One over root fan-in, plus a correction that only bites when fan-out is smaller than fan-in. Tatsu flags his own hand-waving — the minimum singular value argument is not uniform — and says so out loud rather than hiding it.

A2 gives the learning rate. An SGD step on a linear layer is a rank-one outer product ΔWl = −ηl∇hlℓ · hl−1⊤, and expanding Δhl = WlΔhl−1 + ΔWl(hl−1 + Δhl−1) leaves exactly one unknown: how big ‖ΔWl‖* is. The extra assumption that closes it is that the loss decrease per step is Θ(1) — a well-behaved optimiser should not improve less and less as the model widens (60:58). Chain Δℓ = Θ(‖ΔWl‖*‖∇Wlℓ‖*) through and you get η_l = Θ(n_l / n_{l−1}) for SGD (62:02).

Here is the trap, and Tatsu names it before the class can object: for a standard transformer MLP that ratio is just 4, a constant. The SGD result therefore says almost nothing — which is precisely why muP looks like SP under SGD. Redo the derivation for Adam, whose update is normalised rather than proportional to the gradient, and you get η_l = Θ(1 / n_{l−1}) instead (63:06). That is the whole difference from standard practice: not the initialisation, which a correct Kaiming init already gets right, but the fact that a single global learning rate is wrong for Adam and must be scaled per layer by 1/fan-in. Embeddings are the exception — one-hot inputs mean their norms do not grow with vocabulary, so they get treated separately (64:45).

One consequence that catches people: muP papers typically scale attention logits by 1/d rather than 1/√d, for the same activation-size reasons (70:47). If you port a muP recipe into a standard implementation and leave 1/√d in place, you are not running the thing the theory describes.

Does muP survive real architectures?

The last segment covers an independent large-scale ablation study of µ-transfer (69:40). Setup: a normal autoregressive transformer, width swept 128 → 512 → 2048 with depth held fixed, learning-rate sweep at every width. Success means the optimum found at the smallest width is still the optimum at the largest. Baseline result: it transfers (72:26). Then the interesting part — perturb everything modern transformers do that the derivation never covered:

SurvivesSwiGLU and squared ReLU (both also just better); batch size ×4 up or down; zero-query init; SP-style 1/M unembedding
Breakslearnable RMSNorm gains; sign-based optimisers (Lion); strong weight decay (0.1)

The failure modes are informative in different ways. Lion breaking is expected — muP's constants are derived for AdamW's update geometry, so a sign-gradient optimiser has no reason to inherit them (75:12). RMSNorm gains breaking is annoying but survivable, since dropping the gains costs little accuracy. Weight decay is the one that actually hurts: 0.1 is a completely ordinary setting, and it is the closest thing to a real muP failure in the study (75:45).

Two caveats to hold onto. The study varies width only, while real scale-ups grow depth too — so this is the friendly case. And the honest comparison is not "muP is exact" but "SP is much worse": reuse a small-model learning rate under SP at width 2048 and the run degenerates, whereas muP's hero run at 10B parameters kept the same optimal learning rate at 2⁻⁶ (76:17). Tatsu's verdict is deliberately unexcited: useful, evidently easier to tune, adopted by Meta for Llama 4 — and still not a consensus.

Carry this away: muP is not a correctness requirement, it is a search-cost reduction. Nail the learning rate at your target scale and you never needed it (68:35) — its value is that you can afford to find that learning rate on a model 100× smaller. WSD is the same trade in the data dimension: it buys you a Chinchilla sweep for roughly the price of one run. Both are answers to the same question, which is the real question of this lecture — how do I make decisions about a run I cannot afford to repeat?

What you build with this

This lecture closes the scaling half of Assignment 3: Scaling (handout), where you spend a fixed FLOP budget querying a training API to fit your own isoFLOP curves and predict the compute-optimal configuration at a scale you are not allowed to train at. Everything in the DeepSeek section is directly applicable there: fit quadratics per compute scale, join the minima, and be suspicious of any hyperparameter trend that could equally well be a horizontal line. Tatsu says explicitly that the assignment is where you find out whether Chinchilla's method holds up (01:10).

It is also where Assignment 4: Data opens (handout, leaderboard) — you build a Common Crawl filtering, deduplication and quality-classification pipeline and are scored on the loss of a model trained on what you produced. The handoff between the two is the lecture's own loose end: every ratio above moved partly because data quality moved, and A4 is where you get to move it yourself.

Supporting materials, verified

Exercises

  1. muP learning-rate transfer, from scratch code — Build a small autoregressive transformer with configurable width. (1) Implement SP and muP as a switch: muP sets non-embedding init to 1/√fan-in and per-layer Adam learning rates to base_lr/fan-in, with embeddings excluded. (2) Sweep 8 log-spaced learning rates at widths 128, 512 and 2048, short runs. (3) Plot loss vs learning rate, one curve per width, for each parameterisation. (4) Add learnable RMSNorm gains and re-run the muP sweep. A good answer shows the SP optimum sliding left roughly with 1/width while the muP optimum stays put, and reproduces the gains-break-transfer result — and says how many FLOPs the muP path would have saved versus sweeping at the largest width.
  2. Rewind-and-decay: WSD for a cheap Chinchilla sweep code — (1) Train one model with a WSD schedule to your longest token target, checkpointing during the stable phase. (2) For 4 earlier token budgets, restart from the matching checkpoint and decay to zero over ~10% of that budget. (3) Separately train 4 cosine runs each sized to its own target. (4) Overlay the terminal losses. Report the compute ratio between the two paths and whether the WSD losses match, beat, or trail cosine at each budget. A good answer also shows what you would have concluded from naively reading intermediate checkpoints of a single cosine run — the wrong answer this whole technique exists to prevent.
  3. Reconcile the token-to-parameter ratios — Pull the reported compute-optimal ratios from Chinchilla, Llama 3, Hunyuan-Large and MiniCPM, and for each record: the fitting method used (envelope / isoFLOP / joint fit), the architecture, whether it is dense or MoE, and whether the ratio is per total or per active parameter. A good answer separates differences that come from the method from differences that come from the model, and states which of the four numbers you would actually plan a run around and why.
  4. Audit the weakest fit in the lecture — Take DeepSeek's fitted learning-rate scaling law. Estimate, from the width of the loss basin their own grid search shows, how much worse a run would be if you replaced their fitted line with a constant learning rate extrapolated from the smallest scale. A good answer produces a rough tolerance — "the fit has to be right within a factor of k or it costs you nothing" — and uses it to say whether Tatsu's scepticism is a real problem or a stylistic complaint.
Next: L12 Evaluation · Back to the map.