primer.ml.reasoning
Reasoning models: thinking before answering
Run: python -m primer.ml.reasoning
New to the notation? primer.notation explains every symbol used here from
zero. This lesson builds on sampling from primer.ml.inference and on
reinforcement learning from primer.ml.reinforcement.
Level 1: The practitioner's guide
In one sentence. A reasoning model writes intermediate steps before it answers, buying accuracy with tokens, and the practitioner's job is to decide, task by task, how many of those tokens to pay for and how to check what comes back.
When you need it. You need thinking when a task has steps that depend on each other: multi-step arithmetic, a proof, a plan, code that must satisfy several constraints at once, a diagnosis from several clues. The tell is a model that answers instantly and confidently, and is wrong in a way that a moment's working would have caught. This lesson's toy model shows the shape of the gain: answering in one token it is right on one-step sums and on about 5% of the rest; allowed to write one step per token it is right every time, at one token per step. You do not need thinking for lookups, formatting, classification or short factual answers: there the extra tokens cost money and latency and buy nothing, and published work finds reasoning models spending hundreds of tokens on problems like 2 + 3 (Chen et al., 2024). The question is never "thinking on or off" but "how much, and checked how".
Your options. From the cheapest to the most certain:
| Option | What it does | What it buys | What it costs | Where it lives |
|---|---|---|---|---|
| No thinking | The model answers at once | Lowest latency and cost | Wrong on anything with more serial steps than one pass can do | The request |
| Ask for steps in the prompt | "Think step by step", or a worked example with its steps | A real gain on arithmetic and logic from any capable model (Kojima et al.; Wei et al.) | Longer answers; the steps are visible and count as output | Your prompt |
| A reasoning model with a thinking budget | A model trained to think, with a cap or an effort setting on how long | Accuracy that rises roughly with the logarithm of the budget, then flattens | Thinking tokens are billed as output; latency grows with the chain | The request: Anthropic's API takes a budget_tokens target (minimum 1,024) or an effort level |
| Sample several and vote | Several independent chains; return the most common final answer | A solid gain when the model is right more often than any one wrong answer | n times the tokens; no extra waiting if run side by side; needs answers that compare exactly | Your code |
| Sample several and verify | Several chains; keep the one that passes a check (tests, an exact answer, a proof checker, or a learned verifier) | With a reliable check, "right sometimes" becomes "right almost always": pass@n | n times the tokens plus the checker; a learned verifier can be gamed | Your code and your checker |
| Train with verifiable rewards | Reinforcement learning on problems a program can grade | A model whose long, self-checking chains emerge on their own (DeepSeek-R1) | A training run, a graded problem set, and reward design that decides how long it thinks | Training |
How to choose. Route by two questions: does the task need serial steps, and can a program check the answer?
- Easy or latency-critical: no thinking, or the smallest budget the API allows, and measure whether accuracy moves at all.
- Hard and checkable (code with tests, math with a known answer, a schema to validate): think, sample several, keep what passes. This is where spare compute turns into accuracy almost for free; in this lesson's experiment, one sample is right 38% of the time, eight with a step checker 96%.
- Hard and not checkable (judgement, writing, open-ended analysis): think with a capped budget, and vote only when answers can be compared exactly. Gains here are smaller and voting flattens after a few samples.
- Choosing between a bigger model and more thinking: on problems a small model solves sometimes, Snell et al. (2024) found test-time compute can beat a 14× larger model at matched FLOPs, and adapting the budget per prompt beats a flat budget by more than 4×.
- Whatever you pick, measure accuracy and cost per task, and raise the budget only where the measurement says it pays.
What it costs. Thinking is paid per token, as output. Eight chains of 2,000 tokens at an illustrative \$10 per million output tokens cost 16 cents a question; one chain costs 2 cents, and both take 40 seconds at 50 tokens per second when the eight run side by side (this lesson's cost formula). Thinking wider costs money; thinking longer costs money and time. Anthropic's extended thinking documentation puts the practical range at about 1,024 tokens for simple tasks and 16,000 or more for complex ones, with diminishing returns, and recommend batch processing above 32,000 because the requests run long enough to hit timeouts; they also bill thinking as output and report it separately in the response's usage. Changing the budget between requests invalidates prompt caching, because the budget is rendered into the prompt, so hold it stable within a conversation.
What breaks.
- Overthinking. Rewarded for correctness alone, nothing tells a model to stop, so easy questions get long chains. Route easy work away from thinking and cap budgets everywhere else.
- Votes that agree on the same mistake. Samples from one model share its blind spots; in this lesson's simulation, 31 votes with a 0.9 shared error rate are no better than one. Diversity (different prompts, a tool that computes) helps more than more samples.
- Voting on a minority-right model. A yes/no question answered right 40% of the time gets worse with more votes: 0.32 at five. Voting amplifies the common answer, right or wrong.
- A verifier with gaps. More samples mean more chances to find a wrong answer the checker likes; Cobbe et al. (2021) saw best-of-n accuracy fall again after a few hundred samples, and Brown et al. (2024) found voting and reward models plateau beyond several hundred samples where no automatic check exists. Programs beat learned checkers where a program exists.
- Slips compound. At a 2% slip rate a 50-step chain is clean 36% of the time and a 200-step chain 2%. Long chains need per-step checks or a lower slip rate, not just more length.
- The chain is not a confession. Turpin et al. (2023) planted a hidden bias in prompts; answers followed it and the written reasoning never mentioned it. Read the chain as evidence, not as the cause.
- A length penalty that bites harder than its size. In this lesson's training toy, GRPO's division by the group's spread turns a 0.01 penalty into a full advantage, and the easiest problems shrink to one token at the cost of accuracy. Reward shape decides how long a model thinks.
In the wild. Chain-of-thought prompting (Wei et al.) and "let's think
step by step" (Kojima et al.) are the prompt-level versions;
self-consistency (Wang et al.) is the vote; Cobbe et al. trained the first
outcome verifiers on GSM8K and Lightman et al. trained a process reward
model on 800,000 step labels. DeepSeek-R1 showed that reinforcement
learning with verifiable rewards alone produces self-reflection and
verification in the chain, using the GRPO recipe from DeepSeekMath (Shao
et al.). Hosted APIs expose the budget dial directly: Anthropic's
extended thinking takes a budget_tokens target and returns thinking
blocks alongside the answer, with newer models replacing the fixed budget
by an effort setting the model spends adaptively (its extended thinking
documentation). Brown et al.'s Large Language Monkeys is the reference
for how far repeated sampling scales when a checker exists: on SWE-bench
Lite, from 15.9% of issues solved with one sample to 56% with 250.
Go deeper. Level 2 builds each piece with a toy model that can do one addition per token: why writing steps adds serial computation, the budget-accuracy curve, the voting formula and its failure under shared mistakes, outcome and process verifiers with pass@n, a reinforcement learning loop in which longer thinking emerges from correctness alone, and the arithmetic of compounding slips and cost. If you only needed to set a budget and a check, you are done.
Level 2: How it works, from scratch
Ask someone to multiply 37 by 48 in their head, instantly, and they will probably guess. Hand them a pencil and a minute and they will get it right. They didn't get smarter in that minute. They got room: somewhere to write 37 × 8 = 296 and 37 × 40 = 1,480, so each small step could lean on the last.
A language model is in the same position. It can do a fixed amount of work for each token it writes: one pass through all of its layers. A reasoning model is a model that has learned to use the pencil. Before it answers, it writes out intermediate steps, checks them, backs up when one is wrong, and only then commits to an answer. This lesson builds each piece from scratch:
- why writing steps helps at all (chain of thought);
- why more thinking can buy more accuracy (test-time compute);
- sampling several attempts and voting on the answer (self-consistency);
- checking attempts, either the final answer or every step (verifiers);
- how reinforcement learning teaches a model to think this way;
- what thinking costs, and where it still fails.
1. Why thinking out loud helps
Everyday picture. Picture a clerk who, each time they glance at a page, can do exactly one addition and write down one number. Hand them a column of four numbers and allow one glance, and they have to guess. Allow three glances and a margin to write in, and they get it right every time. The margin is the whole trick: what they write on one glance is there to read on the next.
Tiny example. Our toy model is that clerk. Every token it writes costs one forward pass (one run of the whole model to produce one token), and each forward pass can do one addition. Ask it for 3 + 5 + 8 + 2, which needs three additions.
| Pass | Answer at once (1 token) | Think first, then answer |
|---|---|---|
| 1 | adds 3 + 5 = 8, has no pass left for 8 and 2, estimates them at 4.5 each: 8 + 9 = 17 ✗ | writes 3+5=8 |
| 2 | reads 8, writes 8+8=16 |
|
| 3 | reads 16, answers 16 + 2 = 18 ✓ |
The weights are identical in both columns. The only difference is that the second column wrote its running totals down, where the next pass could read them. Writing intermediate steps before the answer is called chain of thought.
flowchart LR subgraph D["Answer at once: 1 token"] direction LR q1["3+5+8+2 ="] --> p1["pass 1<br/>3+5 = 8<br/>no pass left"] --> a1["17, a guess"] end subgraph C["Think first: 3 tokens"] direction LR q2["3+5+8+2 ="] --> s1["pass 1<br/>writes 3+5=8"] --> s2["pass 2<br/>reads 8<br/>writes 8+8=16"] --> s3["pass 3<br/>reads 16<br/>answers 18"] end
Reading it: each box is one forward pass, and each arrow is text handed to the next pass. In the top row there is only one box, so all the work has to fit in it, and it doesn't. In the bottom row every box does one small step and leaves its result in the text. Nothing inside the model carries a running total from one token to the next except what has been written, so the written steps are the model's working memory.
The rule behind the table: the number of steps a model can do one after another grows with the number of tokens it writes.
Level 3: the formula and its symbols
$$ S = L \times T $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $S$ | serial steps: how many steps the model can chain, each using the result of the one before | 1 or 3 |
| $L$ | serial steps one forward pass can do; for a transformer, roughly its number of layers | 1 addition in the toy |
| $T$ | tokens generated, the answer included; each token is one forward pass | 1 at once, 3 thinking first |
| $\times$ | multiply |
In words: "the length of the chain of steps a model can work through equals the steps per token times the number of tokens it writes."
With the numbers: the toy has $L = 1$ and the sum needs 3 additions. Answering at once, $T = 1$, so $S = 1$: two additions short, so it guesses. Thinking first, $T = 3$, so $S = 3$: exactly enough. A 32-layer model answering at once has at most 32 layers of one-after-another work; with a 1,000-token chain of thought it has up to 32,000.
Level 3: in Python
In Python:
# the toy: 1 addition per pass; 3 + 5 + 8 + 2 needs 3 additions
L = 1
# answering at once: T = 1, so S = 1, not enough
L * 1 # → 1
# the one pass adds 3 + 5, then estimates the unreached 8 and 2 at 4.5 each
(3 + 5) + round(4.5 * 2) # → 17
# thinking first: T = 3, so S = 3, exactly enough
L * 3 # → 3
t1 = 3 + 5 # → 8
t2 = t1 + 8 # → 16
t2 + 2 # → 18
# a 32-layer model: answering at once, then after a 1,000-token chain of thought
32 * 1 # → 32
32 * 1000 # → 32000
A layer is not literally one addition, so $L$ is a rough count, but the shape of the rule holds: a model with a fixed depth can only do a fixed amount of one-after-another work per token, and problems that need more must spread it over more tokens. Li et al. (2024) proved a version of this for transformers: with enough chain-of-thought tokens they can solve inherently serial problems that a fixed-depth transformer answering at once cannot.
Reading it: on the left, the x-axis is how many additions a sum needs and the y-axis is how often the toy gets it right. Answering in one token (red), it is perfect on one-addition sums and then collapses to the few percent of lucky estimates. Writing steps (blue), it is right every time. On the right is the price: the blue line climbs one token per addition, while the red line stays at one. Chain of thought doesn't make the model smarter per token; it lets the model buy more tokens.
That price is real: each token costs roughly $2N$ floating-point operations
for a model with $N$ parameters (see primer.ml.inference), so a
1,000-token chain costs a thousand answers' worth of compute. In practice
this is why simply asking a model to "think step by step" (Kojima et al.,
2022) or showing it worked examples with steps (Wei et al., 2022) improved
accuracy on arithmetic and logic puzzles, and why every reasoning model
writes a long chain before its answer.
In code: solve is the toy model: it spends one forward pass per written step, one on the answer, and estimates anything it never reached. It returns an Attempt holding the steps, the answer and the tokens used.
2. Test-time compute: buying accuracy with tokens
Everyday picture. An exam has questions of very different difficulty. Give yourself one minute per question and you finish only the easiest. Two minutes, and the next tier falls. Each doubling of time unlocks one more tier, until you can finish everything and extra time buys nothing.
Tiny example. Test-time compute is computation spent while answering, as opposed to while training. The simplest dial is a thinking budget: the most tokens the model may write before it must answer. Take six problems that need 1, 2, 4, 8, 16 and 32 additions, and a budget of 8 tokens. Four of them fit (1, 2, 4, 8). The other two must be guessed, and say a guess is right 10% of the time. Accuracy is 4/6 + 2/6 × 0.1 = 0.70.
There are two ways to spend test-time compute, and the rest of this lesson uses both.
flowchart LR P[Problem] --> LONG["Think longer<br/>one chain, bigger budget"] P --> WIDE["Think wider<br/>n chains side by side"] LONG --> A1[Answer] WIDE --> PICK["Pick one:<br/>vote or verifier"] --> A2[Answer]
Reading it: the top path spends compute sequentially: one chain, allowed to run longer. It helps when the problem needs many steps in a row. The bottom path spends it in parallel: several independent chains, then a rule that picks one answer. It helps when the model is right sometimes but not reliably. The top path costs waiting time; the bottom path costs money but, run side by side, no extra waiting. Sections 3 and 4 are about the "Pick one" box.
Level 3: the formula and its symbols
$$ \text{acc}(B) = F(B) + \big(1 - F(B)\big)\, g $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $B$ | the thinking budget, in tokens | 8 |
| $F(B)$ | the share of problems that need at most $B$ steps, so fit in the budget | 4/6 = 0.667 |
| $1 - F(B)$ | the share that doesn't fit and must be guessed | 2/6 = 0.333 |
| $g$ | the chance a guess happens to be right | 0.1 |
| $\text{acc}(B)$ | the share of problems answered correctly with budget $B$ | 0.70 |
In words: "accuracy is the share of problems that fit in the budget, plus a lucky share of the ones that don't."
With the numbers: $F(8) = 4/6 = 0.667$, so $\text{acc}(8) = 0.667 + 0.333 \times 0.1 = 0.70$. Doubling to $B = 16$ lets a fifth level fit: $F = 5/6$, $\text{acc} = 0.85$. Every doubling adds the same $1/6 \times 0.9 = 0.15$, until $B = 32$ fits everything and $\text{acc} = 1$.
Level 3: in Python
In Python:
needs = [1, 2, 4, 8, 16, 32]
g = 0.1
B = 8
# F(B): the share of problems that fit
F = sum(d <= B for d in needs) / len(needs)
round(F, 3) # → 0.667
# acc(B) = F(B) + (1 − F(B)) · g
round(F + (1 - F) * g, 3) # → 0.7
# each doubling of B lets one more of the six levels fit
round((1 / len(needs)) * (1 - g), 3) # → 0.15
Reading it: on the left, the x-axis is the budget on a doubling (log) scale. The grey line is the formula; the blue dots are the toy model from section 1 solving real random sums. On a log axis a straight line means "each doubling adds the same amount", and that is what both show until $B = 32$, where every problem fits and the line goes flat. The dots sit a little below the line at small budgets because a guess about many unreached numbers is right less often than 10%. On the right is what was actually spent: always less than the budget, because easy problems stop early, and nothing more past 32.
The straight line on a log axis is built into this toy, because its difficulties are spaced by doubling. Real problem sets are also spread over many scales of difficulty, and published reasoning models show the same shape over a useful range: accuracy rising roughly in proportion to the logarithm of thinking tokens, then leveling off. Snell et al. (2024) found that spending test-time compute adaptively (more on harder prompts) beats spending it evenly, and that on problems a small model can sometimes solve, extra test-time compute can stand in for a much larger model. In practice, model APIs expose this dial as a "reasoning effort" setting or a maximum number of thinking tokens.
In code: budget_accuracy is the formula, and budget_sweep runs solve on random sums at each budget and reports accuracy and tokens actually spent.
3. Sampling many and voting: self-consistency
Everyday picture. Unsure of an answer, you ask five friends separately. If four of them say the same thing, you trust it. Each friend can be wrong, but it is unlikely that most of them are wrong in the same way, unless they all read the same wrong article.
Tiny example. A model that samples its tokens (see primer.ml.inference)
writes a different chain each time. Ask the toy's noisy cousin for
3 + 5 + 8 + 2 three times and it answers 18, 17, 18. Keep only the final
answers and take the most common one: 18. Sampling several chains of thought
and returning the most common final answer is self-consistency (Wang et
al., 2022), also called majority voting.
flowchart LR Q["3+5+8+2 = ?"] --> S1["chain 1 ends in 18"] & S2["chain 2 ends in 17"] & S3["chain 3 ends in 18"] S1 & S2 & S3 --> V["count final answers<br/>18: two votes, 17: one"] --> A["answer 18"]
Reading it: the question fans out into independent chains, each free to take a different route. The chains themselves are thrown away; only their last lines meet in the counting box. That is why voting needs answers that can be compared exactly (a number, a multiple-choice letter): two essays are never identical, so there would be nothing to count.
How much does voting help? Take the simplest case, a yes/no question: every wrong vote lands on the same wrong answer, and each vote is right with the same chance $p$, independently of the others.
Level 3: the formula and its symbols
$$ P_{\text{vote}}(n) = \sum_{k=\lceil n/2 \rceil}^{n} \binom{n}{k}\, p^{k}\, (1-p)^{n-k} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $n$ | number of sampled votes (odd, so there are no ties) | 3 |
| $p$ | chance one vote is right | 0.6 |
| $k$ | how many of the $n$ votes are right | 2 or 3 |
| $\lceil n/2 \rceil$ | $n/2$ rounded up: the smallest number of votes that wins | 2 |
| $\binom{n}{k}$ | "$n$ choose $k$": how many ways to pick which $k$ of the $n$ votes are the right ones | $\binom{3}{2} = 3$ |
| $p^{k}(1-p)^{n-k}$ | the chance of one particular pattern: those $k$ right, the other $n-k$ wrong | $0.6^2 \times 0.4 = 0.144$ |
| $\sum_{k=\lceil n/2 \rceil}^{n}$ | add up over every winning count of right votes | $k = 2, 3$ |
| $P_{\text{vote}}(n)$ | the chance the majority is right | 0.648 |
In words: "the vote is right when at least half the votes are right, so add up the chance of every such count: the number of ways to get that many right, times the chance of each way."
With the numbers: $p = 0.6$, $n = 3$. Two right: $3 \times 0.6^2 \times 0.4 = 0.432$. Three right: $0.6^3 = 0.216$. Together, 0.648, up from 0.6 for one vote. Five votes give 0.683. But a solver right only 40% of the time gets worse with five votes: 0.317. Voting amplifies whichever answer is most common, right or wrong. (This is the Condorcet jury theorem, from 1785.)
Level 3: in Python
In Python:
import math
p, n = 0.6, 3
# one term per winning count k = 2, 3
terms = [math.comb(n, k) * p**k * (1 - p)**(n - k) for k in range(math.ceil(n / 2), n + 1)]
[round(t, 3) for t in terms] # → [0.432, 0.216]
round(sum(terms), 3) # → 0.648
# five votes
round(sum(math.comb(5, k) * 0.6**k * 0.4**(5 - k) for k in range(3, 6)), 3) # → 0.683
# a 40% solver on a yes/no question: voting makes it worse
round(sum(math.comb(5, k) * 0.4**k * 0.6**(5 - k) for k in range(3, 6)), 3) # → 0.317
Reading it: the x-axis is how many samples vote; the y-axis is how often the vote is right. The solid lines are the formula. Above 0.5 (green, blue) more votes push accuracy towards 1; below 0.5 (solid red) they push it towards 0. The dashed red line is the same 40% solver on a question with a numeric answer, where its mistakes scatter over ten different wrong values. No wrong value gets more than a few percent of the votes, so 40% is the biggest pile and the vote climbs past 0.9. That is the situation self-consistency relies on in math problems: the right answer only needs to be the most common one, not a majority.
When voting fails: shared mistakes. Everything above assumed the votes were independent. Samples from one model are not: they share its training, its blind spots and its reading of the question. Model that as a shared draw: each sample copies one common answer with probability $\rho$ (rho), and otherwise answers on its own. Each single sample is still right 60% of the time, so only the correlation changes.
Reading it: every line starts at 0.6 on the left, because one sample is one sample. Independent samples (blue) climb to nearly 1. At $\rho = 0.6$ (amber) and $\rho = 0.9$ (red) the lines go flat at 0.6: whenever the shared answer is wrong, the 60% or 90% of votes that copy it outnumber the at most 40% × 0.6 = 24% (or 6%) that can independently land on the right answer. Thirty-one correlated votes are one opinion, repeated. At $\rho = 0.3$ (green) the right answer's share, 0.7 × 0.6 = 42%, still beats the shared wrong answer's 0.3 + 0.7 × 0.4 / 5 = 35.6%, so the vote does win in the long run, just slowly.
In practice, self-consistency gives a solid gain for a few samples and then flattens, because a model's mistakes on a question are correlated. Diversity helps (different prompts, different models, a tool that computes instead of guessing), and voting only works where answers can be compared exactly.
In code: majority_vote picks the most common answer, majority_accuracy is the formula (with ties on even $n$ counted as a coin flip), and correlated_vote_accuracy simulates votes with scattered wrong answers and a shared-mistake rate.
4. Verifiers: checking the answer, or checking every step
Everyday picture. One teacher looks only at the boxed answer at the bottom of the page. Another marks every line of working. The first can't tell a lucky guess from understanding, and when the answer is wrong, can't tell you where you went wrong. The second can do both.
Tiny example. A verifier is anything that scores a candidate solution. Here are two chains the toy wrote for 3 + 5 + 8 + 2 (true answer 18), each with slips:
| Chain | Final answer | Outcome check: is it 18? | Step check: first line that is false |
|---|---|---|---|
3+5=8, 8+8=17, 17+2=19 |
19 | reward 0 | step 2: 8 + 8 is 16 |
3+5=9, 9+8=16, 16+2=18 |
18 | reward 1 | step 1: 3 + 5 is 8 |
An outcome verifier looks only at the final answer: either a learned outcome reward model (ORM) that predicts whether it is right, or a check against a known answer. It gives the second chain full marks, although two slips just happened to cancel. A process verifier, or process reward model (PRM), scores every step. It catches the second chain's slip and points at exactly where the first one went wrong.
flowchart LR subgraph O["Outcome verifier"] direction LR o1["3+5=8"] --> o2["8+8=17"] --> o3["17+2=19"] --> oc{"final 19<br/>equals 18?"} --> orr["reward 0<br/>but where did it go wrong?"] end subgraph P["Process verifier"] direction LR p1["3+5=8<br/>ok"] --> p2["8+8=17<br/>wrong"] --> p3["17+2=19<br/>ok"] --> pr["first bad step: 2"] end
Reading it: both rows read the same chain. The outcome verifier jumps straight to the diamond at the end and returns one bit. The process verifier stamps every box. Notice that step 3, 17 + 2 = 19, is marked ok: it is correct arithmetic on a wrong input. A step checker judges each step on its own terms, and the first bad step is where the chain left the rails.
Verifiers power best-of-n: sample $n$ chains, keep the one the verifier scores highest. With a perfect verifier, best-of-n is right whenever any of the $n$ samples is right, a number called pass@n.
Level 3: the formula and its symbols
$$ \text{pass@}n = 1 - (1 - p)^{n} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $p$ | chance one sample is right | 0.3 |
| $n$ | number of samples drawn | 5 |
| $1 - p$ | chance one sample is wrong | 0.7 |
| $(1 - p)^{n}$ | chance all $n$ are wrong, multiplying because samples are independent | $0.7^5 = 0.168$ |
| $\text{pass@}n$ | chance at least one of the $n$ is right | 0.832 |
In words: "the chance that at least one sample is right is one minus the chance that every sample is wrong."
With the numbers: a model right 30% of the time is wrong on five tries in a row with chance $0.7^5 = 0.168$, so pass@5 = 0.832. A 30% model with a perfect checker and five tries beats an 80% model on one try. In the experiment below, each chain has 5 additions that each slip with chance 0.2, so a chain is clean with chance $0.8^5 = 0.33$.
Level 3: in Python
In Python:
p, n = 0.3, 5
# (1 − p)^n: every one of the five is wrong
round((1 - p) ** n, 5) # → 0.16807
round(1 - (1 - p) ** n, 5) # → 0.83193
# the experiment's chains: 5 additions, each clean with chance 0.8
round(0.8 ** 5, 3) # → 0.328
Reading it: the x-axis doubles the number of chains sampled; the y-axis is accuracy. One sample (red) stays at 0.38 whatever $n$ is: about a third of chains are clean, and a few more land on 18 by luck. Voting (blue) climbs slowly, because wrong answers bunch on near misses one or two away from the truth. Best-of-n with the step checker (green) hugs the dotted pass@n ceiling: at 8 chains, 0.96 against 0.99. The small gap is the lucky chains: their slips cancelled, so pass@n counts them as right, but the checker refuses them. That is the checker doing its job.
In practice, the best verifiers are programs: unit tests for code, an exact match for a math answer, a proof checker. Where no program can check, a learned verifier stands in. Lightman et al. (2023) trained a process reward model on 800,000 human labels of individual steps and found it picked correct solutions far more reliably than an outcome reward model. A learned verifier can be fooled, though, and more samples mean more chances to find a wrong answer it likes: Cobbe et al. (2021) saw best-of-n accuracy start to fall again after a few hundred samples.
In code: check_step is the toy's process verifier for one line, first_bad_step runs it over a chain, and outcome_reward compares only the final answer with a reference. noisy_chain writes chains whose slips carry forward, pass_at_n is the formula, and verifier_experiment compares one sample, voting, best-of-n with the step checker, and the pass@n ceiling.
5. Learning to reason with reinforcement learning
Everyday picture. A workbook with the answers printed in the back. Nobody shows you how to solve anything. You try each problem several ways, check the back, and do more of whatever worked. Over hundreds of problems you discover good habits on your own: write things down, double-check the tricky step, don't stop too early.
Tiny example. Prompting a model to "think step by step" gets it to
write steps; a reasoning model is trained to write good ones. The recipe
behind models such as DeepSeek-R1 is reinforcement learning (see
primer.ml.reinforcement) with two ingredients:
- Verifiable rewards. Train on problems whose answers a program can check: math with a known final answer, code with unit tests. The reward is 1 if the final answer is right and 0 if not. Nobody grades the chain itself.
- Group comparisons. For each problem, sample a group of $G$ attempts and judge each one against its own group. With $G = 4$ and rewards 1, 0, 0, 1, the group averages 0.5, so the two right attempts are above average (+1) and the two wrong ones below (−1). The model is nudged towards whatever the right ones did, including how they reasoned. This is the heart of GRPO (group relative policy optimization).
flowchart LR P["problem with a checkable answer"] --> G["sample G attempts,<br/>each with its own chain of thought"] G --> R["check each final answer:<br/>reward 1 or 0"] R --> A["advantage: reward minus group mean,<br/>divided by group spread"] A --> U["make above-average attempts more likely,<br/>below-average ones less"] U -->|next batch| P
Reading it: follow the loop clockwise. Nothing in it ever says "think longer" or "check your work"; the only signal is whether the last line was right. Whatever habits the chains of right attempts share get reinforced, batch after batch. The group is what makes this work without a separate model estimating how good an attempt "should" be: the other attempts at the same problem are the baseline.
Level 3: the formula and its symbols
$$ A_i = \frac{r_i - \operatorname{mean}(r_1, \ldots, r_G)}{\operatorname{std}(r_1, \ldots, r_G)} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $G$ | how many attempts are sampled for one problem | 4 |
| $i$ | which attempt, from 1 to $G$ | 1, 2, 3, 4 |
| $r_i$ | attempt $i$'s reward: 1 if its final answer is right, 0 if not | 1, 0, 0, 1 |
| $\operatorname{mean}(\ldots)$ | the average reward in the group | 0.5 |
| $\operatorname{std}(\ldots)$ | the standard deviation: the typical distance of a reward from the mean (see primer.notation) |
0.5 |
| $A_i$ | attempt $i$'s advantage: how much better than its group it did, in units of the group's spread | +1, −1, −1, +1 |
In words: "an attempt's advantage is how far its reward sits above its group's average, measured in units of how spread out the group's rewards are."
With the numbers: rewards 1, 0, 0, 1 have mean 0.5 and standard deviation 0.5, so the advantages are (1 − 0.5)/0.5 = +1 and (0 − 0.5)/0.5 = −1. If all four attempts are right, every reward equals the mean and the spread is 0: every advantage is 0, and the problem teaches nothing. The same holds when all four are wrong. Training therefore needs problems the model solves sometimes.
Level 3: in Python
In Python:
import statistics
r = [1, 0, 0, 1]
mu = statistics.mean(r) # → 0.5
sigma = statistics.pstdev(r) # → 0.5
[(r_i - mu) / sigma for r_i in r] # → [1.0, -1.0, -1.0, 1.0]
# a group that all succeeded: spread 0, nothing to learn
statistics.pstdev([1, 1, 1, 1]) # → 0.0
The toy: learning how long to think. Our policy (the model's rule for choosing what to do) makes one choice per problem: how many tokens to think for, out of 1, 2, 4, 8, 16 or 32. It keeps a separate choice for each difficulty (problems needing 1, 2, 4 or 8 steps), and it starts out preferring short answers, like a model never rewarded for thinking: 63% of the time it answers in 1 token. A chain shorter than the steps needed must guess (right 10% of the time). Otherwise each step gets as many tries as the length allows, and a try slips 10% of the time: one pass through the steps, or two (the first attempt and a recheck that catches a slip), or more. The reward is 1 for a right answer and 0 for a wrong one. That is all.
Reading it: the x-axis is training time. On the left, the blue line (reward for correctness only) shows the average thinking length rising from 2 tokens to about 7.6; on the right, accuracy rising from 0.41 to 0.96. Nobody asked for longer answers: longer answers were simply right more often, so they were reinforced. This is the toy version of what DeepSeek-R1 reported at scale: the length of its chains grew steadily through reinforcement learning, and behaviours like re-checking and backing up appeared without being taught. The red and green lines add a price per token; they come next.
Reading it: each group of bars is one difficulty; bar height is how many tokens the trained model spends on it. The black tick is the steps the problem needs, the grey tick twice that. The blue bars (correctness only) sit on the grey ticks: the model learned to spend more on harder problems, and to leave room for one recheck of every step, which lifts the hardest problems from 0.43 to 0.92. That extra room is self-correction in miniature: the chain gets longer because checking pays.
In code: group_advantages is the formula, success_probability is the toy's chance of solving a problem at a given length, and train_reasoner runs the whole loop: sample a group of lengths per difficulty, reward, compute advantages, and nudge the policy.
The price of thinking
Rewarded for correctness alone, nothing ever tells the model to stop: any extra length that helps even slightly gets reinforced. Real reasoning models show this as overthinking: hundreds of tokens spent on "what is 2 + 3?". The usual remedy is to charge for length: subtract a small penalty per token from the reward. The expected reward for a chain of length $L$ on a problem needing $d$ steps (with $L \ge d$) becomes:
Level 3: the formula and its symbols
$$ \mathbb{E}[r] = \left(1 - \varepsilon^{\lfloor L/d \rfloor}\right)^{d} - \lambda L $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\mathbb{E}[r]$ | the expected (average) reward for this length | 0.763 |
| $d$ | steps the problem needs | 8 |
| $L$ | thinking length, in tokens | 16 |
| $\lfloor L/d \rfloor$ | $L/d$ rounded down: how many tries each step gets | 2 |
| $\varepsilon$ | chance one try at a step slips | 0.1 |
| $\varepsilon^{\lfloor L/d \rfloor}$ | chance every try at one step slips, so the step fails | 0.01 |
| $\left(1 - \ldots\right)^{d}$ | chance all $d$ steps come out right | 0.923 |
| $\lambda$ | lambda: the penalty per token of thinking | 0.01 |
| $\lambda L$ | the total charge for the chain | 0.16 |
In words: "the reward for a length is how likely it is to get every step right, given the tries it allows, minus a small charge for every token."
With the numbers: for an 8-step problem, 8 tokens give $0.9^8 - 0.08 = 0.43 - 0.08 = 0.35$; 16 tokens give $0.99^8 - 0.16 = 0.923 - 0.16 = 0.763$; 32 tokens give $0.999 - 0.32 = 0.679$. Sixteen wins: one recheck is worth paying for, a third is not. For a 1-step problem, 1 token gives $0.9 - 0.01 = 0.89$, 2 tokens $0.99 - 0.02 = 0.97$, 4 tokens $0.96$. Two wins.
Level 3: in Python
In Python:
eps, lam = 0.1, 0.01
# an 8-step problem: chance of success at 8, 16 and 32 tokens
[round((1 - eps ** (L // 8)) ** 8, 3) for L in (8, 16, 32)] # → [0.43, 0.923, 0.999]
# minus λL: 16 tokens wins
[round((1 - eps ** (L // 8)) ** 8 - lam * L, 3) for L in (8, 16, 32)] # → [0.35, 0.763, 0.679]
# a 1-step problem at 1, 2 and 4 tokens: 2 tokens wins
[round((1 - eps ** (L // 1)) ** 1 - lam * L, 3) for L in (1, 2, 4)] # → [0.89, 0.97, 0.96]
Now look back at the two figures. With the penalty and no division by the spread (green), training finds exactly those best lengths, 2 tokens for the easiest problems and 16 for the hardest. With GRPO's division by the spread (red), the easiest problems shrink to 1 token, and accuracy settles at 0.92, below the 0.96 of correctness alone. Why? In a group where every attempt succeeded, rewards differ only by the penalty: 0.99 for 1 token, 0.98 for 2. Their spread is tiny, so dividing by it inflates that 0.01 difference into advantages of +1 and −1, as loud as the difference between solved and failed. The penalty ends up far stronger than its size. Liu et al. (2025) analyse biases like this in GRPO; the lesson for anyone training or budgeting a reasoning model is that the exact shape of the reward decides how long the model thinks.
6. Where reasoning still fails, and how to budget it
Everyday picture. A long line of dominoes falls all the way only if every single one is placed right. Add more dominoes and the chance that one is misplaced grows, however careful you are with each.
Slips compound
Tiny example. If each step of a chain slips 2% of the time, a 10-step chain is clean 82% of the time, and a 50-step chain only 36% of the time.
Level 3: the formula and its symbols
$$ P(\text{no slip}) = (1 - \varepsilon)^{k} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\varepsilon$ | chance one step slips | 0.02 |
| $1 - \varepsilon$ | chance one step is right | 0.98 |
| $k$ | steps in the chain | 50 |
| $(1-\varepsilon)^{k}$ | chance all $k$ are right, multiplying because each step is a separate chance to slip | 0.364 |
In words: "the chance a whole chain is clean is the chance one step is right, multiplied by itself once per step."
With the numbers: $0.98^{10} = 0.817$, $0.98^{50} = 0.364$, $0.98^{200} = 0.018$.
Level 3: in Python
In Python:
eps = 0.02
round((1 - eps) ** 10, 3) # → 0.817
round((1 - eps) ** 50, 3) # → 0.364
round((1 - eps) ** 200, 3) # → 0.018
Reading it: the x-axis is the length of the chain; the y-axis is the
chance it contains no slip at all. Every curve falls, and the only thing
that flattens one is a lower slip rate per step. That is why length alone
is not reasoning: the trained model in section 5 got better by spending
its extra tokens on rechecking, which lowers the effective slip rate, and
why verifiers that check each step matter. The same law governs agents
taking many actions; see primer.agents.planning.
In code: steps_all_right is the formula.
What thinking costs
Tiny example. Eight sampled chains of 2,000 tokens each, at an illustrative \$10 per million output tokens, cost 16 cents per question. Run side by side at 50 tokens per second, the user waits 40 seconds whether you sample one chain or eight.
Level 3: the formula and its symbols
$$ \text{dollars} = \frac{n \cdot T \cdot c}{10^{6}}, \qquad \text{seconds} = \frac{T}{v} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $n$ | chains sampled for one question | 8 |
| $T$ | tokens in each chain | 2,000 |
| $c$ | price per million output tokens | \$10 |
| $10^{6}$ | one million: turns a per-million price into a per-token one | |
| $v$ | generation speed, tokens per second | 50 |
In words: "money grows with every token of every chain; waiting time, with chains run side by side, grows only with the length of one chain."
With the numbers: 8 × 2,000 × 10 / 1,000,000 = \$0.16, and 2,000 / 50 = 40 seconds. One chain instead of eight costs \$0.02 and still takes 40 seconds.
Level 3: in Python
In Python:
n, T, c, v = 8, 2000, 10.0, 50
round(n * T * c / 10**6, 2) # → 0.16
T / v # → 40.0
# one chain instead of eight: an eighth of the money, the same wait
round(1 * T * c / 10**6, 2) # → 0.02
Thinking wider costs money; thinking longer costs money and time. See
primer.agents.cost for pricing requests and measuring cost per successful
task.
In code: reasoning_cost returns both numbers for one question.
Other ways reasoning fails
- The chain is not a transcript. The written reasoning need not be the real cause of the answer. Turpin et al. (2023) nudged models towards an answer with a hidden bias in the prompt; the answers followed the bias, and the written explanations never mentioned it. Treat a chain of thought as evidence about the model's reasoning, not a faithful record of it.
- Overthinking. Long chains on easy questions waste tokens and time, and a model can talk itself out of a right first answer.
- No checker, less progress. Reinforcement learning with verifiable rewards works where a program can check the answer. Open-ended writing, judgement calls and long projects have no cheap checker, and gains there are smaller.
- Checkers get gamed. A verifier with gaps is a target: code that
special-cases the unit tests passes them. This is reward hacking; see
primer.ml.reinforcement.
Budgeting it
Everyday picture. You wouldn't convene a committee to decide what to have for lunch, and you wouldn't let one person sign off a bridge design alone. Match the effort to the stakes and to whether the result can be checked.
flowchart TD Q[Incoming task] --> E{"Easy, or<br/>latency-critical?"} E -->|yes| N["No thinking,<br/>or a small budget"] E -->|no| C{"Can a program<br/>check the answer?"} C -->|"yes: tests, math"| BV["Think, sample n,<br/>keep what passes the check"] C -->|no| B["Think with a capped budget;<br/>vote if answers compare exactly"] N & BV & B --> M["Measure accuracy and cost per task;<br/>raise budgets only where it pays"]
Reading it: start at the top with each incoming task. The first
question routes easy or time-critical work away from thinking entirely,
because that is where overthinking wastes the most. The second asks whether
a program can check the answer: if so, sampling several chains and keeping
the one that passes turns spare compute into accuracy almost for free
(section 4). If not, cap the budget and vote only when answers can be
compared. Everything ends in the same box: measure, because the right
budget is an empirical question per task, not a constant. See
primer.agents.planning for decomposing long tasks into checkable steps.
In 20 seconds
- Chain of thought: every written token is another forward pass, so writing steps gives a fixed-depth model more serial computation, and the text is its working memory.
- Test-time compute: spend more tokens at answer time, either one longer chain or many chains; accuracy often grows with the logarithm of the budget until the problems run out.
- Self-consistency: sample several chains and vote on the final answer; it helps when the right answer is the most common one and mistakes are independent, and stalls when they are shared.
- Verifiers: outcome checks score the answer, process checks score each step; best-of-n with a good checker approaches pass@n = 1 − (1 − p)ⁿ.
- RL with verifiable rewards: reward right final answers, compare each attempt with its group (GRPO), and longer, self-checking reasoning emerges because it pays; a length penalty keeps it from overthinking.
- Cost and limits: thinking is paid per token, slips compound over long chains, and the written chain is not a guaranteed account of why the model answered.
Self-test questions
How can writing its reasoning out make a model more accurate, when its weights don't change? Each token is one forward pass with a fixed amount of serial computation. A problem that needs more sequential steps than one pass can do can't be solved in a single token. Writing intermediate results spreads the work over many passes, and the written text carries each result to the next pass: it is the model's working memory.
What is test-time compute, and what are the two basic ways to spend it? Computation spent while answering rather than while training. Think longer (one chain with a bigger thinking budget, which costs time and money) or think wider (many chains in parallel, then vote or verify, which costs money but no extra waiting when run side by side).
Why does majority voting over samples help, and when does it stop helping? If each sample is independently right more often than any single wrong answer appears, the right answer becomes the biggest pile as samples grow. It stops helping when samples share the same mistakes (they are one opinion repeated), when the model is usually wrong in one consistent way, or when answers can't be compared exactly.
What is the difference between an outcome reward model and a process reward model? An outcome reward model scores only the final answer, so it can't tell a lucky answer from a sound one and can't say where a wrong chain went wrong. A process reward model scores each step, catching errors that cancel out and pointing at the first bad step.
What is pass@n, and why is it a ceiling for best-of-n? The chance that at least one of n samples is right, 1 − (1 − p)ⁿ. A best-of-n picker can only choose among the samples it has, so even a perfect verifier can't do better than "some sample was right".
How does reinforcement learning with verifiable rewards teach a model to reason, if nobody grades the reasoning? Each problem gets a group of sampled attempts, each rewarded 1 or 0 by a program that checks the final answer. Attempts better than their group's average are made more likely, worse ones less likely. Whatever the right attempts had in common, including longer chains and rechecking, is reinforced. A group where every attempt scores the same teaches nothing, so the useful problems are ones the model solves only sometimes.
Why do reasoning models sometimes spend thousands of tokens on easy questions, and what reins that in? A reward for correctness alone never says stop: any extra length that helps slightly is reinforced. A per-token penalty in training, a thinking budget at answer time, and routing easy tasks to little or no thinking all rein it in. The exact shape of the reward matters: dividing by the group's spread, as GRPO does, can magnify a tiny length penalty.
Is a model's chain of thought an accurate account of why it answered? Not necessarily. Experiments that plant a hidden bias in a prompt show answers following the bias while the written reasoning never mentions it. The chain is useful evidence and often helpful, but it is not a guaranteed record of the computation.
The papers behind this lesson
- Wei et al., Chain-of-Thought Prompting Elicits Reasoning in Large Language Models (2022): https://arxiv.org/abs/2201.11903. Showed that prompting with a few worked examples that include intermediate steps makes large models far better at arithmetic, commonsense and symbolic reasoning. Annotated companion
- Wang et al., Self-Consistency Improves Chain of Thought Reasoning in Language Models (2022): https://arxiv.org/abs/2203.11171. Introduced sampling many chains of thought and taking a majority vote over their final answers. Annotated companion
- Cobbe et al., Training Verifiers to Solve Math Word Problems (2021): https://arxiv.org/abs/2110.14168. Introduced the GSM8K dataset and showed that a trained verifier picking the best of many sampled solutions beats fine-tuning alone. Annotated companion
- Lightman et al., Let's Verify Step by Step (2023): https://arxiv.org/abs/2305.20050. Showed that process supervision, a reward model trained on labels for each step, selects correct solutions more reliably than outcome supervision. Annotated companion
- Snell et al., Scaling LLM Test-Time Compute Optimally can be More Effective than Scaling Model Parameters (2024): https://arxiv.org/abs/2408.03314. Measured how to spend test-time compute (longer revisions or verifier-guided search) and showed that adapting it to each prompt's difficulty pays.
- Shao et al., DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models (2024): https://arxiv.org/abs/2402.03300. Introduced GRPO, reinforcement learning that uses a group of sampled answers as its own baseline. Annotated companion
- DeepSeek-AI, DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning (2025): https://arxiv.org/abs/2501.12948. Showed that reinforcement learning with verifiable rewards alone makes long, self-checking chains of thought emerge. Annotated companion
Further reading
- Wei et al., Chain-of-Thought Prompting (2022): https://arxiv.org/abs/2201.11903
- Kojima et al., Large Language Models are Zero-Shot Reasoners ("let's think step by step", 2022): https://arxiv.org/abs/2205.11916
- Nye et al., Show Your Work: Scratchpads for Intermediate Computation with Language Models (2021): https://arxiv.org/abs/2112.00114
- Li et al., Chain of Thought Empowers Transformers to Solve Inherently Serial Problems (2024): https://arxiv.org/abs/2402.12875
- Wang et al., Self-Consistency (2022): https://arxiv.org/abs/2203.11171
- Brown et al., Large Language Monkeys: Scaling Inference Compute with Repeated Sampling (2024): https://arxiv.org/abs/2407.21787
- Cobbe et al., Training Verifiers to Solve Math Word Problems (2021): https://arxiv.org/abs/2110.14168
- Uesato et al., Solving math word problems with process- and outcome-based feedback (2022): https://arxiv.org/abs/2211.14275
- Lightman et al., Let's Verify Step by Step (2023): https://arxiv.org/abs/2305.20050
- Snell et al., Scaling LLM Test-Time Compute Optimally (2024): https://arxiv.org/abs/2408.03314
- Muennighoff et al., s1: Simple test-time scaling (2025): https://arxiv.org/abs/2501.19393
- Zelikman et al., STaR: Bootstrapping Reasoning With Reasoning (2022): https://arxiv.org/abs/2203.14465
- Shao et al., DeepSeekMath (GRPO, 2024): https://arxiv.org/abs/2402.03300
- DeepSeek-AI, DeepSeek-R1 (2025): https://arxiv.org/abs/2501.12948
- Liu et al., Understanding R1-Zero-Like Training: A Critical Perspective (2025): https://arxiv.org/abs/2503.20783
- Chen et al., Do NOT Think That Much for 2+3=? On the Overthinking of o1-Like LLMs (2024): https://arxiv.org/abs/2412.21187
- Turpin et al., Language Models Don't Always Say What They Think (2023): https://arxiv.org/abs/2305.04388
1r""" 2# Reasoning models: thinking before answering 3 4Run: `python -m primer.ml.reasoning` 5 6New to the notation? `primer.notation` explains every symbol used here from 7zero. This lesson builds on sampling from `primer.ml.inference` and on 8reinforcement learning from `primer.ml.reinforcement`. 9 10## Level 1: The practitioner's guide 11 12**In one sentence.** A reasoning model writes intermediate steps before it 13answers, buying accuracy with tokens, and the practitioner's job is to 14decide, task by task, how many of those tokens to pay for and how to check 15what comes back. 16 17**When you need it.** You need thinking when a task has steps that depend 18on each other: multi-step arithmetic, a proof, a plan, code that must 19satisfy several constraints at once, a diagnosis from several clues. The 20tell is a model that answers instantly and confidently, and is wrong in a 21way that a moment's working would have caught. This lesson's toy model 22shows the shape of the gain: answering in one token it is right on 23one-step sums and on about 5% of the rest; allowed to write one step per 24token it is right every time, at one token per step. You do not need 25thinking for lookups, formatting, classification or short factual 26answers: there the extra tokens cost money and latency and buy nothing, 27and published work finds reasoning models spending hundreds of tokens on 28problems like 2 + 3 (Chen et al., 2024). The question is never "thinking 29on or off" but "how much, and checked how". 30 31**Your options.** From the cheapest to the most certain: 32 33| Option | What it does | What it buys | What it costs | Where it lives | 34|---|---|---|---|---| 35| No thinking | The model answers at once | Lowest latency and cost | Wrong on anything with more serial steps than one pass can do | The request | 36| Ask for steps in the prompt | "Think step by step", or a worked example with its steps | A real gain on arithmetic and logic from any capable model (Kojima et al.; Wei et al.) | Longer answers; the steps are visible and count as output | Your prompt | 37| A reasoning model with a thinking budget | A model trained to think, with a cap or an effort setting on how long | Accuracy that rises roughly with the logarithm of the budget, then flattens | Thinking tokens are billed as output; latency grows with the chain | The request: Anthropic's API takes a `budget_tokens` target (minimum 1,024) or an effort level | 38| Sample several and vote | Several independent chains; return the most common final answer | A solid gain when the model is right more often than any one wrong answer | n times the tokens; no extra waiting if run side by side; needs answers that compare exactly | Your code | 39| Sample several and verify | Several chains; keep the one that passes a check (tests, an exact answer, a proof checker, or a learned verifier) | With a reliable check, "right sometimes" becomes "right almost always": pass@n | n times the tokens plus the checker; a learned verifier can be gamed | Your code and your checker | 40| Train with verifiable rewards | Reinforcement learning on problems a program can grade | A model whose long, self-checking chains emerge on their own (DeepSeek-R1) | A training run, a graded problem set, and reward design that decides how long it thinks | Training | 41 42**How to choose.** Route by two questions: does the task need serial 43steps, and can a program check the answer? 44 45- Easy or latency-critical: no thinking, or the smallest budget the API 46 allows, and measure whether accuracy moves at all. 47- Hard and checkable (code with tests, math with a known answer, a schema 48 to validate): think, sample several, keep what passes. This is where 49 spare compute turns into accuracy almost for free; in this lesson's 50 experiment, one sample is right 38% of the time, eight with a step 51 checker 96%. 52- Hard and not checkable (judgement, writing, open-ended analysis): think 53 with a capped budget, and vote only when answers can be compared exactly. 54 Gains here are smaller and voting flattens after a few samples. 55- Choosing between a bigger model and more thinking: on problems a small 56 model solves sometimes, Snell et al. (2024) found test-time compute can 57 beat a 14× larger model at matched FLOPs, and adapting the budget per 58 prompt beats a flat budget by more than 4×. 59- Whatever you pick, measure accuracy and cost per task, and raise the 60 budget only where the measurement says it pays. 61 62**What it costs.** Thinking is paid per token, as output. Eight chains of 632,000 tokens at an illustrative \$10 per million output tokens cost 16 64cents a question; one chain costs 2 cents, and both take 40 seconds at 50 65tokens per second when the eight run side by side (this lesson's cost 66formula). Thinking wider costs money; thinking longer costs money and 67time. Anthropic's extended thinking documentation puts the practical range at about 681,024 tokens for simple tasks and 16,000 or more for complex ones, with 69diminishing returns, and recommend batch processing above 32,000 because 70the requests run long enough to hit timeouts; they also bill thinking as 71output and report it separately in the response's usage. Changing the 72budget between requests invalidates prompt caching, because the budget is 73rendered into the prompt, so hold it stable within a conversation. 74 75**What breaks.** 76 77- **Overthinking.** Rewarded for correctness alone, nothing tells a model 78 to stop, so easy questions get long chains. Route easy work away from 79 thinking and cap budgets everywhere else. 80- **Votes that agree on the same mistake.** Samples from one model share 81 its blind spots; in this lesson's simulation, 31 votes with a 0.9 shared 82 error rate are no better than one. Diversity (different prompts, a 83 tool that computes) helps more than more samples. 84- **Voting on a minority-right model.** A yes/no question answered right 85 40% of the time gets worse with more votes: 0.32 at five. Voting 86 amplifies the common answer, right or wrong. 87- **A verifier with gaps.** More samples mean more chances to find a wrong 88 answer the checker likes; Cobbe et al. (2021) saw best-of-n accuracy 89 fall again after a few hundred samples, and Brown et al. (2024) found 90 voting and reward models plateau beyond several hundred samples where no 91 automatic check exists. Programs beat learned checkers where a program 92 exists. 93- **Slips compound.** At a 2% slip rate a 50-step chain is clean 36% of 94 the time and a 200-step chain 2%. Long chains need per-step checks or a 95 lower slip rate, not just more length. 96- **The chain is not a confession.** Turpin et al. (2023) planted a hidden 97 bias in prompts; answers followed it and the written reasoning never 98 mentioned it. Read the chain as evidence, not as the cause. 99- **A length penalty that bites harder than its size.** In this lesson's 100 training toy, GRPO's division by the group's spread turns a 0.01 penalty 101 into a full advantage, and the easiest problems shrink to one token at 102 the cost of accuracy. Reward shape decides how long a model thinks. 103 104**In the wild.** Chain-of-thought prompting (Wei et al.) and "let's think 105step by step" (Kojima et al.) are the prompt-level versions; 106self-consistency (Wang et al.) is the vote; Cobbe et al. trained the first 107outcome verifiers on GSM8K and Lightman et al. trained a process reward 108model on 800,000 step labels. DeepSeek-R1 showed that reinforcement 109learning with verifiable rewards alone produces self-reflection and 110verification in the chain, using the GRPO recipe from DeepSeekMath (Shao 111et al.). Hosted APIs expose the budget dial directly: Anthropic's 112extended thinking takes a `budget_tokens` target and returns thinking 113blocks alongside the answer, with newer models replacing the fixed budget 114by an effort setting the model spends adaptively (its extended thinking 115documentation). Brown et al.'s *Large Language Monkeys* is the reference 116for how far repeated sampling scales when a checker exists: on SWE-bench 117Lite, from 15.9% of issues solved with one sample to 56% with 250. 118 119**Go deeper.** Level 2 builds each piece with a toy model that can do one 120addition per token: why writing steps adds serial computation, the 121budget-accuracy curve, the voting formula and its failure under shared 122mistakes, outcome and process verifiers with pass@n, a reinforcement 123learning loop in which longer thinking emerges from correctness alone, 124and the arithmetic of compounding slips and cost. If you only needed to 125set a budget and a check, you are done. 126 127## Level 2: How it works, from scratch 128 129Ask someone to multiply 37 by 48 in their head, instantly, and they will 130probably guess. Hand them a pencil and a minute and they will get it right. 131They didn't get smarter in that minute. They got room: somewhere to write 13237 × 8 = 296 and 37 × 40 = 1,480, so each small step could lean on the last. 133 134A language model is in the same position. It can do a fixed amount of work 135for each token it writes: one pass through all of its layers. A **reasoning 136model** is a model that has learned to use the pencil. Before it answers, it 137writes out intermediate steps, checks them, backs up when one is wrong, and 138only then commits to an answer. This lesson builds each piece from scratch: 139 1401. why writing steps helps at all (**chain of thought**); 1412. why more thinking can buy more accuracy (**test-time compute**); 1423. sampling several attempts and voting on the answer (**self-consistency**); 1434. checking attempts, either the final answer or every step (**verifiers**); 1445. how reinforcement learning teaches a model to think this way; 1456. what thinking costs, and where it still fails. 146 147## 1. Why thinking out loud helps 148 149**Everyday picture.** Picture a clerk who, each time they glance at a page, 150can do exactly one addition and write down one number. Hand them a column of 151four numbers and allow one glance, and they have to guess. Allow three 152glances and a margin to write in, and they get it right every time. The 153margin is the whole trick: what they write on one glance is there to read on 154the next. 155 156**Tiny example.** Our toy model is that clerk. Every token it writes costs 157one **forward pass** (one run of the whole model to produce one token), and 158each forward pass can do one addition. Ask it for 3 + 5 + 8 + 2, which needs 159three additions. 160 161| Pass | Answer at once (1 token) | Think first, then answer | 162|---|---|---| 163| 1 | adds 3 + 5 = 8, has no pass left for 8 and 2, estimates them at 4.5 each: 8 + 9 = **17** ✗ | writes `3+5=8` | 164| 2 | | reads 8, writes `8+8=16` | 165| 3 | | reads 16, answers 16 + 2 = **18** ✓ | 166 167The weights are identical in both columns. The only difference is that the 168second column wrote its running totals down, where the next pass could read 169them. Writing intermediate steps before the answer is called **chain of 170thought**. 171 172```mermaid 173flowchart LR 174 subgraph D["Answer at once: 1 token"] 175 direction LR 176 q1["3+5+8+2 ="] --> p1["pass 1<br/>3+5 = 8<br/>no pass left"] --> a1["17, a guess"] 177 end 178 subgraph C["Think first: 3 tokens"] 179 direction LR 180 q2["3+5+8+2 ="] --> s1["pass 1<br/>writes 3+5=8"] --> s2["pass 2<br/>reads 8<br/>writes 8+8=16"] --> s3["pass 3<br/>reads 16<br/>answers 18"] 181 end 182``` 183 184**Reading it:** each box is one forward pass, and each arrow is text handed 185to the next pass. In the top row there is only one box, so all the work has 186to fit in it, and it doesn't. In the bottom row every box does one small 187step and leaves its result in the text. Nothing inside the model carries a 188running total from one token to the next except what has been written, so 189the written steps *are* the model's working memory. 190 191The rule behind the table: the number of steps a model can do one after 192another grows with the number of tokens it writes. 193 194$$ 195S = L \times T 196$$ 197 198**Symbols** 199 200| Symbol | Meaning here | In the example | 201|---|---|---| 202| $S$ | serial steps: how many steps the model can chain, each using the result of the one before | 1 or 3 | 203| $L$ | serial steps one forward pass can do; for a transformer, roughly its number of layers | 1 addition in the toy | 204| $T$ | tokens generated, the answer included; each token is one forward pass | 1 at once, 3 thinking first | 205| $\times$ | multiply | | 206 207**In words:** "the length of the chain of steps a model can work through 208equals the steps per token times the number of tokens it writes." 209 210**With the numbers:** the toy has $L = 1$ and the sum needs 3 additions. 211Answering at once, $T = 1$, so $S = 1$: two additions short, so it guesses. 212Thinking first, $T = 3$, so $S = 3$: exactly enough. A 32-layer model 213answering at once has at most 32 layers of one-after-another work; with a 2141,000-token chain of thought it has up to 32,000. 215 216**In Python:** 217 218```python 219# the toy: 1 addition per pass; 3 + 5 + 8 + 2 needs 3 additions 220L = 1 221# answering at once: T = 1, so S = 1, not enough 222L * 1 # → 1 223# the one pass adds 3 + 5, then estimates the unreached 8 and 2 at 4.5 each 224(3 + 5) + round(4.5 * 2) # → 17 225# thinking first: T = 3, so S = 3, exactly enough 226L * 3 # → 3 227t1 = 3 + 5 # → 8 228t2 = t1 + 8 # → 16 229t2 + 2 # → 18 230# a 32-layer model: answering at once, then after a 1,000-token chain of thought 23132 * 1 # → 32 23232 * 1000 # → 32000 233``` 234 235A layer is not literally one addition, so $L$ is a rough count, but the 236shape of the rule holds: a model with a fixed depth can only do a fixed 237amount of one-after-another work per token, and problems that need more 238must spread it over more tokens. Li et al. (2024) proved a version of this 239for transformers: with enough chain-of-thought tokens they can solve 240inherently serial problems that a fixed-depth transformer answering at once 241cannot. 242 243 244 245**Reading it:** on the left, the x-axis is how many additions a sum needs 246and the y-axis is how often the toy gets it right. Answering in one token 247(red), it is perfect on one-addition sums and then collapses to the few 248percent of lucky estimates. Writing steps (blue), it is right every time. On 249the right is the price: the blue line climbs one token per addition, while 250the red line stays at one. Chain of thought doesn't make the model smarter 251per token; it lets the model buy more tokens. 252 253That price is real: each token costs roughly $2N$ floating-point operations 254for a model with $N$ parameters (see `primer.ml.inference`), so a 2551,000-token chain costs a thousand answers' worth of compute. In practice 256this is why simply asking a model to "think step by step" (Kojima et al., 2572022) or showing it worked examples with steps (Wei et al., 2022) improved 258accuracy on arithmetic and logic puzzles, and why every reasoning model 259writes a long chain before its answer. 260 261**In code:** `solve` is the toy model: it spends one forward pass per written step, one on the answer, and estimates anything it never reached. It returns an `Attempt` holding the steps, the answer and the tokens used. 262 263## 2. Test-time compute: buying accuracy with tokens 264 265**Everyday picture.** An exam has questions of very different difficulty. 266Give yourself one minute per question and you finish only the easiest. Two 267minutes, and the next tier falls. Each doubling of time unlocks one more 268tier, until you can finish everything and extra time buys nothing. 269 270**Tiny example.** **Test-time compute** is computation spent while 271answering, as opposed to while training. The simplest dial is a **thinking 272budget**: the most tokens the model may write before it must answer. Take 273six problems that need 1, 2, 4, 8, 16 and 32 additions, and a budget of 2748 tokens. Four of them fit (1, 2, 4, 8). The other two must be guessed, and 275say a guess is right 10% of the time. Accuracy is 4/6 + 2/6 × 0.1 = **0.70**. 276 277There are two ways to spend test-time compute, and the rest of this lesson 278uses both. 279 280```mermaid 281flowchart LR 282 P[Problem] --> LONG["Think longer<br/>one chain, bigger budget"] 283 P --> WIDE["Think wider<br/>n chains side by side"] 284 LONG --> A1[Answer] 285 WIDE --> PICK["Pick one:<br/>vote or verifier"] --> A2[Answer] 286``` 287 288**Reading it:** the top path spends compute **sequentially**: one chain, 289allowed to run longer. It helps when the problem needs many steps in a row. 290The bottom path spends it **in parallel**: several independent chains, then 291a rule that picks one answer. It helps when the model is right sometimes 292but not reliably. The top path costs waiting time; the bottom path costs 293money but, run side by side, no extra waiting. Sections 3 and 4 are about 294the "Pick one" box. 295 296$$ 297\text{acc}(B) = F(B) + \big(1 - F(B)\big)\, g 298$$ 299 300**Symbols** 301 302| Symbol | Meaning here | In the example | 303|---|---|---| 304| $B$ | the thinking budget, in tokens | 8 | 305| $F(B)$ | the share of problems that need at most $B$ steps, so fit in the budget | 4/6 = 0.667 | 306| $1 - F(B)$ | the share that doesn't fit and must be guessed | 2/6 = 0.333 | 307| $g$ | the chance a guess happens to be right | 0.1 | 308| $\text{acc}(B)$ | the share of problems answered correctly with budget $B$ | 0.70 | 309 310**In words:** "accuracy is the share of problems that fit in the budget, 311plus a lucky share of the ones that don't." 312 313**With the numbers:** $F(8) = 4/6 = 0.667$, so 314$\text{acc}(8) = 0.667 + 0.333 \times 0.1 = 0.70$. Doubling to $B = 16$ lets a 315fifth level fit: $F = 5/6$, $\text{acc} = 0.85$. Every doubling adds the 316same $1/6 \times 0.9 = 0.15$, until $B = 32$ fits everything and $\text{acc} = 1$. 317 318**In Python:** 319 320```python 321needs = [1, 2, 4, 8, 16, 32] 322g = 0.1 323B = 8 324# F(B): the share of problems that fit 325F = sum(d <= B for d in needs) / len(needs) 326round(F, 3) # → 0.667 327# acc(B) = F(B) + (1 − F(B)) · g 328round(F + (1 - F) * g, 3) # → 0.7 329# each doubling of B lets one more of the six levels fit 330round((1 / len(needs)) * (1 - g), 3) # → 0.15 331``` 332 333 334 335**Reading it:** on the left, the x-axis is the budget on a doubling (log) 336scale. The grey line is the formula; the blue dots are the toy model from 337section 1 solving real random sums. On a log axis a straight line means 338"each doubling adds the same amount", and that is what both show until 339$B = 32$, where every problem fits and the line goes flat. The dots sit a 340little below the line at small budgets because a guess about many 341unreached numbers is right less often than 10%. On the right is what was 342actually spent: always less than the budget, because easy problems stop 343early, and nothing more past 32. 344 345The straight line on a log axis is built into this toy, because its 346difficulties are spaced by doubling. Real problem sets are also spread over 347many scales of difficulty, and published reasoning models show the same 348shape over a useful range: accuracy rising roughly in proportion to the 349logarithm of thinking tokens, then leveling off. Snell et al. (2024) found 350that spending test-time compute adaptively (more on harder prompts) beats 351spending it evenly, and that on problems a small model can sometimes solve, 352extra test-time compute can stand in for a much larger model. In practice, 353model APIs expose this dial as a "reasoning effort" setting or a maximum 354number of thinking tokens. 355 356**In code:** `budget_accuracy` is the formula, and `budget_sweep` runs `solve` on random sums at each budget and reports accuracy and tokens actually spent. 357 358## 3. Sampling many and voting: self-consistency 359 360**Everyday picture.** Unsure of an answer, you ask five friends separately. 361If four of them say the same thing, you trust it. Each friend can be wrong, 362but it is unlikely that most of them are wrong *in the same way*, unless 363they all read the same wrong article. 364 365**Tiny example.** A model that samples its tokens (see `primer.ml.inference`) 366writes a different chain each time. Ask the toy's noisy cousin for 3673 + 5 + 8 + 2 three times and it answers 18, 17, 18. Keep only the final 368answers and take the most common one: 18. Sampling several chains of thought 369and returning the most common final answer is **self-consistency** (Wang et 370al., 2022), also called **majority voting**. 371 372```mermaid 373flowchart LR 374 Q["3+5+8+2 = ?"] --> S1["chain 1 ends in 18"] & S2["chain 2 ends in 17"] & S3["chain 3 ends in 18"] 375 S1 & S2 & S3 --> V["count final answers<br/>18: two votes, 17: one"] --> A["answer 18"] 376``` 377 378**Reading it:** the question fans out into independent chains, each free to 379take a different route. The chains themselves are thrown away; only their 380last lines meet in the counting box. That is why voting needs answers that 381can be compared exactly (a number, a multiple-choice letter): two essays are 382never identical, so there would be nothing to count. 383 384How much does voting help? Take the simplest case, a yes/no question: every 385wrong vote lands on the same wrong answer, and each vote is right with the 386same chance $p$, independently of the others. 387 388$$ 389P_{\text{vote}}(n) = \sum_{k=\lceil n/2 \rceil}^{n} \binom{n}{k}\, p^{k}\, (1-p)^{n-k} 390$$ 391 392**Symbols** 393 394| Symbol | Meaning here | In the example | 395|---|---|---| 396| $n$ | number of sampled votes (odd, so there are no ties) | 3 | 397| $p$ | chance one vote is right | 0.6 | 398| $k$ | how many of the $n$ votes are right | 2 or 3 | 399| $\lceil n/2 \rceil$ | $n/2$ rounded up: the smallest number of votes that wins | 2 | 400| $\binom{n}{k}$ | "$n$ choose $k$": how many ways to pick which $k$ of the $n$ votes are the right ones | $\binom{3}{2} = 3$ | 401| $p^{k}(1-p)^{n-k}$ | the chance of one particular pattern: those $k$ right, the other $n-k$ wrong | $0.6^2 \times 0.4 = 0.144$ | 402| $\sum_{k=\lceil n/2 \rceil}^{n}$ | add up over every winning count of right votes | $k = 2, 3$ | 403| $P_{\text{vote}}(n)$ | the chance the majority is right | 0.648 | 404 405**In words:** "the vote is right when at least half the votes are right, so 406add up the chance of every such count: the number of ways to get that many 407right, times the chance of each way." 408 409**With the numbers:** $p = 0.6$, $n = 3$. Two right: $3 \times 0.6^2 \times 0.4 = 4100.432$. Three right: $0.6^3 = 0.216$. Together, **0.648**, up from 0.6 for 411one vote. Five votes give 0.683. But a solver right only 40% of the time 412gets *worse* with five votes: 0.317. Voting amplifies whichever answer is 413most common, right or wrong. (This is the Condorcet jury theorem, from 1785.) 414 415**In Python:** 416 417```python 418import math 419p, n = 0.6, 3 420# one term per winning count k = 2, 3 421terms = [math.comb(n, k) * p**k * (1 - p)**(n - k) for k in range(math.ceil(n / 2), n + 1)] 422[round(t, 3) for t in terms] # → [0.432, 0.216] 423round(sum(terms), 3) # → 0.648 424# five votes 425round(sum(math.comb(5, k) * 0.6**k * 0.4**(5 - k) for k in range(3, 6)), 3) # → 0.683 426# a 40% solver on a yes/no question: voting makes it worse 427round(sum(math.comb(5, k) * 0.4**k * 0.6**(5 - k) for k in range(3, 6)), 3) # → 0.317 428``` 429 430 431 432**Reading it:** the x-axis is how many samples vote; the y-axis is how often 433the vote is right. The solid lines are the formula. Above 0.5 (green, blue) 434more votes push accuracy towards 1; below 0.5 (solid red) they push it 435towards 0. The dashed red line is the same 40% solver on a question with a 436numeric answer, where its mistakes scatter over ten different wrong values. 437No wrong value gets more than a few percent of the votes, so 40% is the 438biggest pile and the vote climbs past 0.9. That is the situation 439self-consistency relies on in math problems: the right answer only needs to 440be the most common one, not a majority. 441 442**When voting fails: shared mistakes.** Everything above assumed the votes 443were independent. Samples from one model are not: they share its training, 444its blind spots and its reading of the question. Model that as a shared 445draw: each sample copies one common answer with probability $\rho$ (rho), 446and otherwise answers on its own. Each single sample is still right 60% of 447the time, so only the correlation changes. 448 449 450 451**Reading it:** every line starts at 0.6 on the left, because one sample is 452one sample. Independent samples (blue) climb to nearly 1. At $\rho = 0.6$ 453(amber) and $\rho = 0.9$ (red) the lines go flat at 0.6: whenever the shared 454answer is wrong, the 60% or 90% of votes that copy it outnumber the at most 45540% × 0.6 = 24% (or 6%) that can independently land on the right answer. 456Thirty-one correlated votes are one opinion, repeated. At $\rho = 0.3$ 457(green) the right answer's share, 0.7 × 0.6 = 42%, still beats the shared 458wrong answer's 0.3 + 0.7 × 0.4 / 5 = 35.6%, so the vote does win in the 459long run, just slowly. 460 461In practice, self-consistency gives a solid gain for a few samples and then 462flattens, because a model's mistakes on a question are correlated. Diversity 463helps (different prompts, different models, a tool that computes instead of 464guessing), and voting only works where answers can be compared exactly. 465 466**In code:** `majority_vote` picks the most common answer, `majority_accuracy` is the formula (with ties on even $n$ counted as a coin flip), and `correlated_vote_accuracy` simulates votes with scattered wrong answers and a shared-mistake rate. 467 468## 4. Verifiers: checking the answer, or checking every step 469 470**Everyday picture.** One teacher looks only at the boxed answer at the 471bottom of the page. Another marks every line of working. The first can't 472tell a lucky guess from understanding, and when the answer is wrong, can't 473tell you where you went wrong. The second can do both. 474 475**Tiny example.** A **verifier** is anything that scores a candidate 476solution. Here are two chains the toy wrote for 3 + 5 + 8 + 2 (true answer 47718), each with slips: 478 479| Chain | Final answer | Outcome check: is it 18? | Step check: first line that is false | 480|---|---|---|---| 481| `3+5=8`, `8+8=17`, `17+2=19` | 19 | reward 0 | step 2: 8 + 8 is 16 | 482| `3+5=9`, `9+8=16`, `16+2=18` | 18 | reward 1 | step 1: 3 + 5 is 8 | 483 484An **outcome verifier** looks only at the final answer: either a learned 485**outcome reward model (ORM)** that predicts whether it is right, or a check 486against a known answer. It gives the second chain full marks, although two 487slips just happened to cancel. A **process verifier**, or **process reward 488model (PRM)**, scores every step. It catches the second chain's slip and 489points at exactly where the first one went wrong. 490 491```mermaid 492flowchart LR 493 subgraph O["Outcome verifier"] 494 direction LR 495 o1["3+5=8"] --> o2["8+8=17"] --> o3["17+2=19"] --> oc{"final 19<br/>equals 18?"} --> orr["reward 0<br/>but where did it go wrong?"] 496 end 497 subgraph P["Process verifier"] 498 direction LR 499 p1["3+5=8<br/>ok"] --> p2["8+8=17<br/>wrong"] --> p3["17+2=19<br/>ok"] --> pr["first bad step: 2"] 500 end 501``` 502 503**Reading it:** both rows read the same chain. The outcome verifier jumps 504straight to the diamond at the end and returns one bit. The process 505verifier stamps every box. Notice that step 3, 17 + 2 = 19, is marked ok: it 506is correct arithmetic on a wrong input. A step checker judges each step on 507its own terms, and the first bad step is where the chain left the rails. 508 509Verifiers power **best-of-n**: sample $n$ chains, keep the one the verifier 510scores highest. With a perfect verifier, best-of-n is right whenever *any* 511of the $n$ samples is right, a number called **pass@n**. 512 513$$ 514\text{pass@}n = 1 - (1 - p)^{n} 515$$ 516 517**Symbols** 518 519| Symbol | Meaning here | In the example | 520|---|---|---| 521| $p$ | chance one sample is right | 0.3 | 522| $n$ | number of samples drawn | 5 | 523| $1 - p$ | chance one sample is wrong | 0.7 | 524| $(1 - p)^{n}$ | chance all $n$ are wrong, multiplying because samples are independent | $0.7^5 = 0.168$ | 525| $\text{pass@}n$ | chance at least one of the $n$ is right | 0.832 | 526 527**In words:** "the chance that at least one sample is right is one minus the 528chance that every sample is wrong." 529 530**With the numbers:** a model right 30% of the time is wrong on five tries 531in a row with chance $0.7^5 = 0.168$, so pass@5 = **0.832**. A 30% model with 532a perfect checker and five tries beats an 80% model on one try. In the 533experiment below, each chain has 5 additions that each slip with chance 0.2, 534so a chain is clean with chance $0.8^5 = 0.33$. 535 536**In Python:** 537 538```python 539p, n = 0.3, 5 540# (1 − p)^n: every one of the five is wrong 541round((1 - p) ** n, 5) # → 0.16807 542round(1 - (1 - p) ** n, 5) # → 0.83193 543# the experiment's chains: 5 additions, each clean with chance 0.8 544round(0.8 ** 5, 3) # → 0.328 545``` 546 547 548 549**Reading it:** the x-axis doubles the number of chains sampled; the y-axis 550is accuracy. One sample (red) stays at 0.38 whatever $n$ is: about a third 551of chains are clean, and a few more land on 18 by luck. Voting (blue) climbs 552slowly, because wrong answers bunch on near misses one or two away from 553the truth. Best-of-n with the step 554checker (green) hugs the dotted pass@n ceiling: at 8 chains, 0.96 against 5550.99. The small gap is the lucky chains: their slips cancelled, so pass@n 556counts them as right, but the checker refuses them. That is the checker 557doing its job. 558 559In practice, the best verifiers are programs: unit tests for code, an exact 560match for a math answer, a proof checker. Where no program can check, a 561learned verifier stands in. Lightman et al. (2023) trained a process reward 562model on 800,000 human labels of individual steps and found it picked 563correct solutions far more reliably than an outcome reward model. A learned 564verifier can be fooled, though, and more samples mean more chances to find 565a wrong answer it likes: Cobbe et al. (2021) saw best-of-n accuracy start to 566fall again after a few hundred samples. 567 568**In code:** `check_step` is the toy's process verifier for one line, `first_bad_step` runs it over a chain, and `outcome_reward` compares only the final answer with a reference. `noisy_chain` writes chains whose slips carry forward, `pass_at_n` is the formula, and `verifier_experiment` compares one sample, voting, best-of-n with the step checker, and the pass@n ceiling. 569 570## 5. Learning to reason with reinforcement learning 571 572**Everyday picture.** A workbook with the answers printed in the back. 573Nobody shows you how to solve anything. You try each problem several ways, 574check the back, and do more of whatever worked. Over hundreds of problems 575you discover good habits on your own: write things down, double-check the 576tricky step, don't stop too early. 577 578**Tiny example.** Prompting a model to "think step by step" gets it to 579write steps; a reasoning model is *trained* to write good ones. The recipe 580behind models such as DeepSeek-R1 is reinforcement learning (see 581`primer.ml.reinforcement`) with two ingredients: 582 583- **Verifiable rewards.** Train on problems whose answers a program can 584 check: math with a known final answer, code with unit tests. The reward is 585 1 if the final answer is right and 0 if not. Nobody grades the chain 586 itself. 587- **Group comparisons.** For each problem, sample a group of $G$ attempts 588 and judge each one against its own group. With $G = 4$ and rewards 589 1, 0, 0, 1, the group averages 0.5, so the two right attempts are above 590 average (+1) and the two wrong ones below (−1). The model is nudged 591 towards whatever the right ones did, including how they reasoned. 592 This is the heart of **GRPO** (group relative policy optimization). 593 594```mermaid 595flowchart LR 596 P["problem with a checkable answer"] --> G["sample G attempts,<br/>each with its own chain of thought"] 597 G --> R["check each final answer:<br/>reward 1 or 0"] 598 R --> A["advantage: reward minus group mean,<br/>divided by group spread"] 599 A --> U["make above-average attempts more likely,<br/>below-average ones less"] 600 U -->|next batch| P 601``` 602 603**Reading it:** follow the loop clockwise. Nothing in it ever says "think 604longer" or "check your work"; the only signal is whether the last line was 605right. Whatever habits the chains of right attempts share get reinforced, 606batch after batch. The group is what makes this work without a separate 607model estimating how good an attempt "should" be: the other attempts at the 608same problem are the baseline. 609 610$$ 611A_i = \frac{r_i - \operatorname{mean}(r_1, \ldots, r_G)}{\operatorname{std}(r_1, \ldots, r_G)} 612$$ 613 614**Symbols** 615 616| Symbol | Meaning here | In the example | 617|---|---|---| 618| $G$ | how many attempts are sampled for one problem | 4 | 619| $i$ | which attempt, from 1 to $G$ | 1, 2, 3, 4 | 620| $r_i$ | attempt $i$'s reward: 1 if its final answer is right, 0 if not | 1, 0, 0, 1 | 621| $\operatorname{mean}(\ldots)$ | the average reward in the group | 0.5 | 622| $\operatorname{std}(\ldots)$ | the **standard deviation**: the typical distance of a reward from the mean (see `primer.notation`) | 0.5 | 623| $A_i$ | attempt $i$'s **advantage**: how much better than its group it did, in units of the group's spread | +1, −1, −1, +1 | 624 625**In words:** "an attempt's advantage is how far its reward sits above its 626group's average, measured in units of how spread out the group's rewards 627are." 628 629**With the numbers:** rewards 1, 0, 0, 1 have mean 0.5 and standard 630deviation 0.5, so the advantages are (1 − 0.5)/0.5 = +1 and 631(0 − 0.5)/0.5 = −1. If all four attempts are right, every reward equals the 632mean and the spread is 0: every advantage is 0, and the problem teaches 633nothing. The same holds when all four are wrong. Training therefore needs 634problems the model solves *sometimes*. 635 636**In Python:** 637 638```python 639import statistics 640r = [1, 0, 0, 1] 641mu = statistics.mean(r) # → 0.5 642sigma = statistics.pstdev(r) # → 0.5 643[(r_i - mu) / sigma for r_i in r] # → [1.0, -1.0, -1.0, 1.0] 644# a group that all succeeded: spread 0, nothing to learn 645statistics.pstdev([1, 1, 1, 1]) # → 0.0 646``` 647 648**The toy: learning how long to think.** Our policy (the model's rule for 649choosing what to do) makes one choice per problem: how many tokens to think 650for, out of 1, 2, 4, 8, 16 or 32. It keeps a separate choice for each 651difficulty (problems needing 1, 2, 4 or 8 steps), and it starts out 652preferring short answers, like a model never rewarded for thinking: 63% of 653the time it answers in 1 token. A chain shorter than the steps needed must 654guess (right 10% of the time). Otherwise each step gets as many tries as the 655length allows, and a try slips 10% of the time: one pass through the steps, 656or two (the first attempt and a recheck that catches a slip), or more. The 657reward is 1 for a right answer and 0 for a wrong one. That is all. 658 659 660 661**Reading it:** the x-axis is training time. On the left, the blue line 662(reward for correctness only) shows the average thinking length rising from 6632 tokens to about 7.6; on the right, accuracy rising from 0.41 to 0.96. 664Nobody asked for longer answers: longer answers were simply right more 665often, so they were reinforced. This is the toy version of what DeepSeek-R1 666reported at scale: the length of its chains grew steadily through 667reinforcement learning, and behaviours like re-checking and backing up 668appeared without being taught. The red and green lines add a price per 669token; they come next. 670 671 672 673**Reading it:** each group of bars is one difficulty; bar height is how many 674tokens the trained model spends on it. The black tick is the steps the 675problem needs, the grey tick twice that. The blue bars (correctness only) 676sit on the grey ticks: the model learned to spend *more on harder 677problems*, and to leave room for one recheck of every step, which lifts the 678hardest problems from 0.43 to 0.92. That extra room is self-correction in 679miniature: the chain gets longer because checking pays. 680 681**In code:** `group_advantages` is the formula, `success_probability` is the toy's chance of solving a problem at a given length, and `train_reasoner` runs the whole loop: sample a group of lengths per difficulty, reward, compute advantages, and nudge the policy. 682 683### The price of thinking 684 685Rewarded for correctness alone, nothing ever tells the model to *stop*: 686any extra length that helps even slightly gets reinforced. Real reasoning 687models show this as **overthinking**: hundreds of tokens spent on "what is 6882 + 3?". The usual remedy is to charge for length: subtract a small penalty 689per token from the reward. The expected reward for a chain of length $L$ on 690a problem needing $d$ steps (with $L \ge d$) becomes: 691 692$$ 693\mathbb{E}[r] = \left(1 - \varepsilon^{\lfloor L/d \rfloor}\right)^{d} - \lambda L 694$$ 695 696**Symbols** 697 698| Symbol | Meaning here | In the example | 699|---|---|---| 700| $\mathbb{E}[r]$ | the **expected** (average) reward for this length | 0.763 | 701| $d$ | steps the problem needs | 8 | 702| $L$ | thinking length, in tokens | 16 | 703| $\lfloor L/d \rfloor$ | $L/d$ rounded down: how many tries each step gets | 2 | 704| $\varepsilon$ | chance one try at a step slips | 0.1 | 705| $\varepsilon^{\lfloor L/d \rfloor}$ | chance *every* try at one step slips, so the step fails | 0.01 | 706| $\left(1 - \ldots\right)^{d}$ | chance all $d$ steps come out right | 0.923 | 707| $\lambda$ | lambda: the penalty per token of thinking | 0.01 | 708| $\lambda L$ | the total charge for the chain | 0.16 | 709 710**In words:** "the reward for a length is how likely it is to get every step 711right, given the tries it allows, minus a small charge for every token." 712 713**With the numbers:** for an 8-step problem, 8 tokens give 714$0.9^8 - 0.08 = 0.43 - 0.08 = 0.35$; 16 tokens give 715$0.99^8 - 0.16 = 0.923 - 0.16 = 0.763$; 32 tokens give 716$0.999 - 0.32 = 0.679$. Sixteen wins: one recheck is worth paying for, a 717third is not. For a 1-step problem, 1 token gives $0.9 - 0.01 = 0.89$, 2 718tokens $0.99 - 0.02 = 0.97$, 4 tokens $0.96$. Two wins. 719 720**In Python:** 721 722```python 723eps, lam = 0.1, 0.01 724# an 8-step problem: chance of success at 8, 16 and 32 tokens 725[round((1 - eps ** (L // 8)) ** 8, 3) for L in (8, 16, 32)] # → [0.43, 0.923, 0.999] 726# minus λL: 16 tokens wins 727[round((1 - eps ** (L // 8)) ** 8 - lam * L, 3) for L in (8, 16, 32)] # → [0.35, 0.763, 0.679] 728# a 1-step problem at 1, 2 and 4 tokens: 2 tokens wins 729[round((1 - eps ** (L // 1)) ** 1 - lam * L, 3) for L in (1, 2, 4)] # → [0.89, 0.97, 0.96] 730``` 731 732Now look back at the two figures. With the penalty and no division by the 733spread (green), training finds exactly those best lengths, 2 tokens for the 734easiest problems and 16 for the hardest. With GRPO's division by the spread 735(red), the easiest problems shrink to **1** token, and accuracy settles at 7360.92, below the 0.96 of correctness alone. Why? In a group where every attempt succeeded, 737rewards differ only by the penalty: 0.99 for 1 token, 0.98 for 2. Their 738spread is tiny, so dividing by it inflates that 0.01 difference into 739advantages of +1 and −1, as loud as the difference between solved and 740failed. The penalty ends up far stronger than its size. Liu et al. (2025) 741analyse biases like this in GRPO; the lesson for anyone training or 742budgeting a reasoning model is that the exact shape of the reward decides 743how long the model thinks. 744 745## 6. Where reasoning still fails, and how to budget it 746 747**Everyday picture.** A long line of dominoes falls all the way only if 748every single one is placed right. Add more dominoes and the chance that one 749is misplaced grows, however careful you are with each. 750 751### Slips compound 752 753**Tiny example.** If each step of a chain slips 2% of the time, a 10-step 754chain is clean 82% of the time, and a 50-step chain only 36% of the time. 755 756$$ 757P(\text{no slip}) = (1 - \varepsilon)^{k} 758$$ 759 760**Symbols** 761 762| Symbol | Meaning here | In the example | 763|---|---|---| 764| $\varepsilon$ | chance one step slips | 0.02 | 765| $1 - \varepsilon$ | chance one step is right | 0.98 | 766| $k$ | steps in the chain | 50 | 767| $(1-\varepsilon)^{k}$ | chance all $k$ are right, multiplying because each step is a separate chance to slip | 0.364 | 768 769**In words:** "the chance a whole chain is clean is the chance one step is 770right, multiplied by itself once per step." 771 772**With the numbers:** $0.98^{10} = 0.817$, $0.98^{50} = 0.364$, 773$0.98^{200} = 0.018$. 774 775**In Python:** 776 777```python 778eps = 0.02 779round((1 - eps) ** 10, 3) # → 0.817 780round((1 - eps) ** 50, 3) # → 0.364 781round((1 - eps) ** 200, 3) # → 0.018 782``` 783 784 785 786**Reading it:** the x-axis is the length of the chain; the y-axis is the 787chance it contains no slip at all. Every curve falls, and the only thing 788that flattens one is a lower slip rate per step. That is why length alone 789is not reasoning: the trained model in section 5 got better by spending 790its extra tokens on *rechecking*, which lowers the effective slip rate, and 791why verifiers that check each step matter. The same law governs agents 792taking many actions; see `primer.agents.planning`. 793 794**In code:** `steps_all_right` is the formula. 795 796### What thinking costs 797 798**Tiny example.** Eight sampled chains of 2,000 tokens each, at an 799illustrative \$10 per million output tokens, cost 16 cents per question. Run 800side by side at 50 tokens per second, the user waits 40 seconds whether you 801sample one chain or eight. 802 803$$ 804\text{dollars} = \frac{n \cdot T \cdot c}{10^{6}}, \qquad \text{seconds} = \frac{T}{v} 805$$ 806 807**Symbols** 808 809| Symbol | Meaning here | In the example | 810|---|---|---| 811| $n$ | chains sampled for one question | 8 | 812| $T$ | tokens in each chain | 2,000 | 813| $c$ | price per million output tokens | \$10 | 814| $10^{6}$ | one million: turns a per-million price into a per-token one | | 815| $v$ | generation speed, tokens per second | 50 | 816 817**In words:** "money grows with every token of every chain; waiting time, 818with chains run side by side, grows only with the length of one chain." 819 820**With the numbers:** 8 × 2,000 × 10 / 1,000,000 = **\$0.16**, and 8212,000 / 50 = **40 seconds**. One chain instead of eight costs \$0.02 and 822still takes 40 seconds. 823 824**In Python:** 825 826```python 827n, T, c, v = 8, 2000, 10.0, 50 828round(n * T * c / 10**6, 2) # → 0.16 829T / v # → 40.0 830# one chain instead of eight: an eighth of the money, the same wait 831round(1 * T * c / 10**6, 2) # → 0.02 832``` 833 834Thinking wider costs money; thinking longer costs money *and* time. See 835`primer.agents.cost` for pricing requests and measuring cost per successful 836task. 837 838**In code:** `reasoning_cost` returns both numbers for one question. 839 840### Other ways reasoning fails 841 842- **The chain is not a transcript.** The written reasoning need not be the 843 real cause of the answer. Turpin et al. (2023) nudged models towards an 844 answer with a hidden bias in the prompt; the answers followed the bias, 845 and the written explanations never mentioned it. Treat a chain of thought 846 as evidence about the model's reasoning, not a faithful record of it. 847- **Overthinking.** Long chains on easy questions waste tokens and time, 848 and a model can talk itself out of a right first answer. 849- **No checker, less progress.** Reinforcement learning with verifiable 850 rewards works where a program can check the answer. Open-ended writing, 851 judgement calls and long projects have no cheap checker, and gains there 852 are smaller. 853- **Checkers get gamed.** A verifier with gaps is a target: code that 854 special-cases the unit tests passes them. This is reward hacking; see 855 `primer.ml.reinforcement`. 856 857### Budgeting it 858 859**Everyday picture.** You wouldn't convene a committee to decide what to 860have for lunch, and you wouldn't let one person sign off a bridge design 861alone. Match the effort to the stakes and to whether the result can be 862checked. 863 864```mermaid 865flowchart TD 866 Q[Incoming task] --> E{"Easy, or<br/>latency-critical?"} 867 E -->|yes| N["No thinking,<br/>or a small budget"] 868 E -->|no| C{"Can a program<br/>check the answer?"} 869 C -->|"yes: tests, math"| BV["Think, sample n,<br/>keep what passes the check"] 870 C -->|no| B["Think with a capped budget;<br/>vote if answers compare exactly"] 871 N & BV & B --> M["Measure accuracy and cost per task;<br/>raise budgets only where it pays"] 872``` 873 874**Reading it:** start at the top with each incoming task. The first 875question routes easy or time-critical work away from thinking entirely, 876because that is where overthinking wastes the most. The second asks whether 877a program can check the answer: if so, sampling several chains and keeping 878the one that passes turns spare compute into accuracy almost for free 879(section 4). If not, cap the budget and vote only when answers can be 880compared. Everything ends in the same box: measure, because the right 881budget is an empirical question per task, not a constant. See 882`primer.agents.planning` for decomposing long tasks into checkable steps. 883 884## In 20 seconds 885 886- **Chain of thought:** every written token is another forward pass, so 887 writing steps gives a fixed-depth model more serial computation, and the 888 text is its working memory. 889- **Test-time compute:** spend more tokens at answer time, either one 890 longer chain or many chains; accuracy often grows with the logarithm of 891 the budget until the problems run out. 892- **Self-consistency:** sample several chains and vote on the final 893 answer; it helps when the right answer is the most common one and mistakes 894 are independent, and stalls when they are shared. 895- **Verifiers:** outcome checks score the answer, process checks score each 896 step; best-of-n with a good checker approaches pass@n = 1 − (1 − p)ⁿ. 897- **RL with verifiable rewards:** reward right final answers, compare each 898 attempt with its group (GRPO), and longer, self-checking reasoning emerges 899 because it pays; a length penalty keeps it from overthinking. 900- **Cost and limits:** thinking is paid per token, slips compound over long 901 chains, and the written chain is not a guaranteed account of why the 902 model answered. 903 904## Self-test questions 905 906**How can writing its reasoning out make a model more accurate, when its 907weights don't change?** 908Each token is one forward pass with a fixed amount of serial computation. 909A problem that needs more sequential steps than one pass can do can't be 910solved in a single token. Writing intermediate results spreads the work over 911many passes, and the written text carries each result to the next pass: it 912is the model's working memory. 913 914**What is test-time compute, and what are the two basic ways to spend it?** 915Computation spent while answering rather than while training. Think longer 916(one chain with a bigger thinking budget, which costs time and money) or 917think wider (many chains in parallel, then vote or verify, which costs money 918but no extra waiting when run side by side). 919 920**Why does majority voting over samples help, and when does it stop 921helping?** 922If each sample is independently right more often than any single wrong 923answer appears, the right answer becomes the biggest pile as samples grow. 924It stops helping when samples share the same mistakes (they are one opinion 925repeated), when the model is usually wrong in one consistent way, or when 926answers can't be compared exactly. 927 928**What is the difference between an outcome reward model and a process 929reward model?** 930An outcome reward model scores only the final answer, so it can't tell a 931lucky answer from a sound one and can't say where a wrong chain went wrong. 932A process reward model scores each step, catching errors that cancel out and 933pointing at the first bad step. 934 935**What is pass@n, and why is it a ceiling for best-of-n?** 936The chance that at least one of n samples is right, 1 − (1 − p)ⁿ. A 937best-of-n picker can only choose among the samples it has, so even a 938perfect verifier can't do better than "some sample was right". 939 940**How does reinforcement learning with verifiable rewards teach a model to 941reason, if nobody grades the reasoning?** 942Each problem gets a group of sampled attempts, each rewarded 1 or 0 by a 943program that checks the final answer. Attempts better than their group's 944average are made more likely, worse ones less likely. Whatever the right 945attempts had in common, including longer chains and rechecking, is 946reinforced. A group where every attempt scores the same teaches nothing, so 947the useful problems are ones the model solves only sometimes. 948 949**Why do reasoning models sometimes spend thousands of tokens on easy 950questions, and what reins that in?** 951A reward for correctness alone never says stop: any extra length that helps 952slightly is reinforced. A per-token penalty in training, a thinking budget 953at answer time, and routing easy tasks to little or no thinking all rein it 954in. The exact shape of the reward matters: dividing by the group's spread, 955as GRPO does, can magnify a tiny length penalty. 956 957**Is a model's chain of thought an accurate account of why it answered?** 958Not necessarily. Experiments that plant a hidden bias in a prompt show 959answers following the bias while the written reasoning never mentions it. 960The chain is useful evidence and often helpful, but it is not a guaranteed 961record of the computation. 962 963## The papers behind this lesson 964 965- **Wei et al., *Chain-of-Thought Prompting Elicits Reasoning in Large 966 Language Models* (2022)**: https://arxiv.org/abs/2201.11903. Showed that 967 prompting with a few worked examples that include intermediate steps makes 968 large models far better at arithmetic, commonsense and symbolic 969 reasoning. 970 [Annotated companion](../../papers/chain-of-thought-prompting.html) 971- **Wang et al., *Self-Consistency Improves Chain of Thought Reasoning in 972 Language Models* (2022)**: https://arxiv.org/abs/2203.11171. Introduced 973 sampling many chains of thought and taking a majority vote over their 974 final answers. 975 [Annotated companion](../../papers/self-consistency.html) 976- **Cobbe et al., *Training Verifiers to Solve Math Word Problems* 977 (2021)**: https://arxiv.org/abs/2110.14168. Introduced the GSM8K dataset 978 and showed that a trained verifier picking the best of many sampled 979 solutions beats fine-tuning alone. 980 [Annotated companion](../../papers/training-verifiers.html) 981- **Lightman et al., *Let's Verify Step by Step* (2023)**: 982 https://arxiv.org/abs/2305.20050. Showed that process supervision, a 983 reward model trained on labels for each step, selects correct solutions 984 more reliably than outcome supervision. 985 [Annotated companion](../../papers/lets-verify-step-by-step.html) 986- **Snell et al., *Scaling LLM Test-Time Compute Optimally can be More 987 Effective than Scaling Model Parameters* (2024)**: 988 https://arxiv.org/abs/2408.03314. Measured how to spend test-time compute 989 (longer revisions or verifier-guided search) and showed that adapting it 990 to each prompt's difficulty pays. 991- **Shao et al., *DeepSeekMath: Pushing the Limits of Mathematical 992 Reasoning in Open Language Models* (2024)**: 993 https://arxiv.org/abs/2402.03300. Introduced GRPO, reinforcement 994 learning that uses a group of sampled answers as its own baseline. 995 [Annotated companion](../../papers/deepseekmath-grpo.html) 996- **DeepSeek-AI, *DeepSeek-R1: Incentivizing Reasoning Capability in LLMs 997 via Reinforcement Learning* (2025)**: https://arxiv.org/abs/2501.12948. 998 Showed that reinforcement learning with verifiable rewards alone makes 999 long, self-checking chains of thought emerge. 1000 [Annotated companion](../../papers/deepseek-r1.html) 1001 1002## Further reading 1003 1004- Wei et al., *Chain-of-Thought Prompting* (2022): https://arxiv.org/abs/2201.11903 1005- Kojima et al., *Large Language Models are Zero-Shot Reasoners* ("let's think step by step", 2022): https://arxiv.org/abs/2205.11916 1006- Nye et al., *Show Your Work: Scratchpads for Intermediate Computation with Language Models* (2021): https://arxiv.org/abs/2112.00114 1007- Li et al., *Chain of Thought Empowers Transformers to Solve Inherently Serial Problems* (2024): https://arxiv.org/abs/2402.12875 1008- Wang et al., *Self-Consistency* (2022): https://arxiv.org/abs/2203.11171 1009- Brown et al., *Large Language Monkeys: Scaling Inference Compute with Repeated Sampling* (2024): https://arxiv.org/abs/2407.21787 1010- Cobbe et al., *Training Verifiers to Solve Math Word Problems* (2021): https://arxiv.org/abs/2110.14168 1011- Uesato et al., *Solving math word problems with process- and outcome-based feedback* (2022): https://arxiv.org/abs/2211.14275 1012- Lightman et al., *Let's Verify Step by Step* (2023): https://arxiv.org/abs/2305.20050 1013- Snell et al., *Scaling LLM Test-Time Compute Optimally* (2024): https://arxiv.org/abs/2408.03314 1014- Muennighoff et al., *s1: Simple test-time scaling* (2025): https://arxiv.org/abs/2501.19393 1015- Zelikman et al., *STaR: Bootstrapping Reasoning With Reasoning* (2022): https://arxiv.org/abs/2203.14465 1016- Shao et al., *DeepSeekMath* (GRPO, 2024): https://arxiv.org/abs/2402.03300 1017- DeepSeek-AI, *DeepSeek-R1* (2025): https://arxiv.org/abs/2501.12948 1018- Liu et al., *Understanding R1-Zero-Like Training: A Critical Perspective* (2025): https://arxiv.org/abs/2503.20783 1019- Chen et al., *Do NOT Think That Much for 2+3=? On the Overthinking of o1-Like LLMs* (2024): https://arxiv.org/abs/2412.21187 1020- Turpin et al., *Language Models Don't Always Say What They Think* (2023): https://arxiv.org/abs/2305.04388 1021""" 1022 1023from __future__ import annotations 1024 1025import math 1026import re 1027from collections import Counter 1028from dataclasses import dataclass, field 1029 1030import numpy as np 1031 1032from primer._show import banner, say, table, takeaway 1033 1034# --------------------------------------------------------------------------- 1035# 1. Thinking out loud: a toy model that does one addition per forward pass 1036# --------------------------------------------------------------------------- 1037 1038# The toy model's rough sense of "a digit I didn't have time to add": the average of 0..9. 1039DIGIT_MEAN = 4.5 1040 1041 1042@dataclass 1043class Attempt: 1044 """What the toy model wrote for one problem. 1045 1046 steps: the intermediate lines it wrote before answering, e.g. ["3+5=8", "8+8=16"] 1047 answer: the number it finally gave 1048 tokens: forward passes used, one per written step plus one for the answer 1049 """ 1050 1051 steps: list[str] = field(default_factory=list) 1052 answer: int = 0 1053 tokens: int = 1 1054 1055 1056def solve(numbers: list[int], max_tokens: int) -> Attempt: 1057 """Add up `numbers` with a model that can do ONE addition per forward pass. 1058 1059 Every token the model writes is one forward pass. A step token ("8+8=16") 1060 spends its pass on one addition and leaves the running total in the text, 1061 where the next pass can read it. The answer token also gets one addition. 1062 If the budget runs out before the sum is done, the model estimates the 1063 numbers it never reached at `DIGIT_MEAN` each: a guess, right only by luck. 1064 """ 1065 total = numbers[0] 1066 if len(numbers) == 1: 1067 return Attempt([], total, 1) 1068 steps: list[str] = [] 1069 i = 1 1070 # Thinking tokens: all but the last addition, as long as the budget leaves a token for the answer. 1071 while i < len(numbers) - 1 and len(steps) < max_tokens - 1: 1072 new = total + numbers[i] 1073 steps.append(f"{total}+{numbers[i]}={new}") 1074 total, i = new, i + 1 1075 # The answer token's own forward pass: one more addition. 1076 total += numbers[i] 1077 unreached = len(numbers) - (i + 1) 1078 # Round half up, so one unreached number is estimated as 5 (4.5 rounded). 1079 answer = total + math.floor(DIGIT_MEAN * unreached + 0.5) 1080 return Attempt(steps, int(answer), len(steps) + 1) 1081 1082 1083# --------------------------------------------------------------------------- 1084# 2. Test-time compute: accuracy as a function of the thinking budget 1085# --------------------------------------------------------------------------- 1086 1087# Problems need 1, 2, 4, ..., 32 additions: difficulty spaced by doubling, as real benchmarks roughly are. 1088NEEDS = (1, 2, 4, 8, 16, 32) 1089 1090 1091def budget_accuracy(budget: int, needs: tuple[int, ...] = NEEDS, lucky: float = 0.1) -> float: 1092 """acc(B) = F(B) + (1 - F(B)) * g. 1093 1094 F(B) is the share of problems that fit in a budget of B tokens (need at 1095 most B additions); the rest are guessed, and a guess is right with 1096 probability g (`lucky`). 1097 """ 1098 fits = sum(d <= budget for d in needs) / len(needs) 1099 return fits + (1 - fits) * lucky 1100 1101 1102def budget_sweep(budgets=(1, 2, 4, 8, 16, 32, 64), n_problems: int = 600, seed: int = 0) -> list[dict]: 1103 """Run `solve` on real random sums at each budget: accuracy and tokens actually spent. 1104 1105 Each problem needs d additions, d drawn evenly from `NEEDS`, so it has d + 1 1106 random digits. The same problems are reused at every budget. 1107 """ 1108 rng = np.random.default_rng(seed) 1109 problems = [rng.integers(0, 10, size=int(rng.choice(NEEDS)) + 1).tolist() for _ in range(n_problems)] 1110 rows = [] 1111 for b in budgets: 1112 attempts = [solve(p, b) for p in problems] 1113 right = [a.answer == sum(p) for a, p in zip(attempts, problems)] 1114 rows.append(dict(budget=b, accuracy=float(np.mean(right)), mean_tokens=float(np.mean([a.tokens for a in attempts])))) 1115 return rows 1116 1117 1118# --------------------------------------------------------------------------- 1119# 3. Sampling many and voting: self-consistency 1120# --------------------------------------------------------------------------- 1121 1122 1123def majority_vote(answers: list) -> object: 1124 """The most common answer. Ties go to the answer seen first.""" 1125 return Counter(answers).most_common(1)[0][0] 1126 1127 1128def majority_accuracy(p: float, n: int) -> float: 1129 """Chance that more than half of n independent votes are right, each right with probability p. 1130 1131 A yes/no question: every wrong vote lands on the same wrong answer. With 1132 an even n, a tie is broken by a coin flip, so it counts half. 1133 """ 1134 win = sum(math.comb(n, k) * p**k * (1 - p) ** (n - k) for k in range(n // 2 + 1, n + 1)) 1135 tie = math.comb(n, n // 2) * p ** (n // 2) * (1 - p) ** (n // 2) if n % 2 == 0 else 0.0 1136 return win + tie / 2 1137 1138 1139def correlated_vote_accuracy( 1140 p: float, n: int, rho: float = 0.0, n_wrong: int = 5, trials: int = 2000, seed: int = 0 1141) -> float: 1142 """Plurality-vote accuracy over n samples whose mistakes may be shared. 1143 1144 Answer 0 is right; answers 1..n_wrong are the wrong answers. Each problem 1145 has one shared draw (right with probability p, else one "trap" answer). 1146 Each sample copies the shared draw with probability `rho`, otherwise draws 1147 independently (right with probability p, else a random wrong answer). 1148 Every single sample is right with probability p either way: only the 1149 correlation between samples changes. 1150 """ 1151 rng = np.random.default_rng(seed) 1152 # (trials,): the shared draw per problem. 1153 shared = np.where(rng.random(trials) < p, 0, rng.integers(1, n_wrong + 1, trials)) 1154 # (trials, n): each sample's own independent draw. 1155 own = np.where(rng.random((trials, n)) < p, 0, rng.integers(1, n_wrong + 1, (trials, n))) 1156 votes = np.where(rng.random((trials, n)) < rho, shared[:, None], own) 1157 counts = np.zeros((trials, n_wrong + 1)) 1158 np.add.at(counts, (np.arange(trials)[:, None], votes), 1) 1159 # A random tiebreak below 1 vote, so ties don't systematically favour answer 0. 1160 winner = np.argmax(counts + rng.random(counts.shape) * 0.5, axis=1) 1161 return float(np.mean(winner == 0)) 1162 1163 1164# --------------------------------------------------------------------------- 1165# 4. Verifiers: checking the answer, or checking every step 1166# --------------------------------------------------------------------------- 1167 1168STEP = re.compile(r"^\s*(-?\d+)\s*\+\s*(-?\d+)\s*=\s*(-?\d+)\s*$") 1169 1170 1171def check_step(step: str) -> bool: 1172 """A process verifier for the toy: does this one line "a+b=c" hold?""" 1173 m = STEP.match(step) 1174 return bool(m) and int(m[1]) + int(m[2]) == int(m[3]) 1175 1176 1177def first_bad_step(steps: list[str]) -> int | None: 1178 """Index of the first step the process verifier rejects, or None if every step checks out.""" 1179 return next((i for i, s in enumerate(steps) if not check_step(s)), None) 1180 1181 1182def final_answer(steps: list[str]) -> int: 1183 """The number after the last "=": the chain's answer.""" 1184 return int(steps[-1].rsplit("=", 1)[1]) 1185 1186 1187def outcome_reward(steps: list[str], reference: int) -> float: 1188 """An outcome check: 1 if the final answer matches the reference, else 0. It never reads the steps.""" 1189 return 1.0 if final_answer(steps) == reference else 0.0 1190 1191 1192def noisy_chain(numbers: list[int], step_error: float, rng: np.random.Generator) -> list[str]: 1193 """Write every addition as a step; each slips by ±1 or ±2 with probability `step_error`. 1194 1195 A slip is carried forward: the next step adds to the wrong running total, 1196 which is how one early mistake poisons the final answer. 1197 """ 1198 total, steps = numbers[0], [] 1199 for x in numbers[1:]: 1200 new = total + x 1201 if rng.random() < step_error: 1202 new += int(rng.choice([-2, -1, 1, 2])) 1203 steps.append(f"{total}+{x}={new}") 1204 total = new 1205 return steps 1206 1207 1208def pass_at_n(p: float, n: int) -> float: 1209 """Chance that at least one of n independent samples is right: 1 - (1 - p)^n.""" 1210 return 1 - (1 - p) ** n 1211 1212 1213def verifier_experiment( 1214 ns=(1, 2, 4, 8, 16, 32), n_terms: int = 6, step_error: float = 0.2, trials: int = 400, seed: int = 0 1215) -> list[dict]: 1216 """Compare four ways to turn n sampled chains into one answer. 1217 1218 - single: the first sample, as if n were 1. 1219 - vote: the most common final answer (self-consistency). 1220 - verifier: the first chain whose every step checks out (best-of-n with a 1221 process verifier), falling back to the vote if none does. 1222 - pass_at_n: whether ANY chain's final answer is right, the ceiling a 1223 perfect answer-picker could reach. 1224 """ 1225 rng = np.random.default_rng(seed) 1226 hits = {n: dict(single=0, vote=0, verifier=0, pass_at_n=0) for n in ns} 1227 for _ in range(trials): 1228 numbers = rng.integers(0, 10, n_terms).tolist() 1229 truth = sum(numbers) 1230 chains = [noisy_chain(numbers, step_error, rng) for _ in range(max(ns))] 1231 for n in ns: 1232 pool = chains[:n] 1233 answers = [final_answer(c) for c in pool] 1234 clean = [c for c in pool if first_bad_step(c) is None] 1235 picked = final_answer(clean[0]) if clean else majority_vote(answers) 1236 h = hits[n] 1237 h["single"] += answers[0] == truth 1238 h["vote"] += majority_vote(answers) == truth 1239 h["verifier"] += picked == truth 1240 h["pass_at_n"] += truth in answers 1241 return [dict(n=n, **{k: v / trials for k, v in hits[n].items()}) for n in ns] 1242 1243 1244# --------------------------------------------------------------------------- 1245# 5. Learning to reason with RL: verifiable rewards and group-relative advantages 1246# --------------------------------------------------------------------------- 1247 1248DIFFICULTIES = (1, 2, 4, 8) # additions a problem needs 1249LENGTHS = (1, 2, 4, 8, 16, 32) # thinking lengths the policy can choose 1250INIT_TILT = 1.0 # how strongly the untrained policy prefers short answers 1251 1252 1253def group_advantages(rewards: np.ndarray, divide_by_spread: bool = True) -> np.ndarray: 1254 """GRPO's advantage: each reward compared with its own group, A_i = (r_i - mean) / std. 1255 1256 A group where every sample scored the same carries no information about 1257 which choice was better, so every advantage is 0. With 1258 `divide_by_spread=False` the advantage is just r_i - mean, which keeps a 1259 tiny reward difference tiny instead of blowing it up to ±1. 1260 """ 1261 rewards = np.asarray(rewards, dtype=float) 1262 centred = rewards - rewards.mean() 1263 spread = rewards.std() 1264 if not divide_by_spread: 1265 return centred 1266 if spread == 0: 1267 return np.zeros_like(rewards) 1268 return centred / spread 1269 1270 1271def success_probability(length: int, needs: int, step_error: float = 0.1, lucky: float = 0.1) -> float: 1272 """How often a chain of `length` tokens solves a problem needing `needs` steps. 1273 1274 Too short (length < needs): it must guess, right with probability `lucky`. 1275 Otherwise every step gets `length // needs` tries: the first attempt plus 1276 rechecks that can catch and fix a slip. A step fails only if every try 1277 slips, so the chain succeeds with probability (1 - e^tries)^needs. More 1278 length always helps a little, which is why a reward for correctness alone 1279 never tells the model to stop. 1280 """ 1281 if length < needs: 1282 return lucky 1283 tries = length // needs 1284 return (1 - step_error**tries) ** needs 1285 1286 1287def train_reasoner( 1288 iterations: int = 600, 1289 group_size: int = 16, 1290 lr: float = 0.2, 1291 length_penalty: float = 0.0, 1292 step_error: float = 0.1, 1293 lucky: float = 0.1, 1294 divide_by_spread: bool = True, 1295 seed: int = 0, 1296) -> dict: 1297 """Train a tiny policy to choose how long to think, from outcome rewards only. 1298 1299 The policy is a table of logits: one row per difficulty, one column per 1300 thinking length in `LENGTHS`. It starts out preferring short answers, like 1301 a model that was never rewarded for thinking. Each iteration, for each 1302 difficulty, it samples a group of lengths, scores each one with a 1303 verifiable reward (1 if right, 0 if wrong, minus `length_penalty` per 1304 token), turns the rewards into group-relative advantages and nudges the 1305 logits along the policy gradient: logits += lr * mean(A_i * (onehot_i - pi)). 1306 1307 Returns the expected mean length and accuracy after every iteration, and 1308 the length each difficulty ends up preferring. 1309 """ 1310 rng = np.random.default_rng(seed) 1311 lengths = np.array(LENGTHS) 1312 # Short answers favoured at the start: pi(1 token) ≈ 0.63, pi(32 tokens) ≈ 0.004. 1313 logits = np.tile(-INIT_TILT * np.arange(len(LENGTHS)), (len(DIFFICULTIES), 1)) 1314 # success[d, a]: how often length a solves difficulty d (d rows, a columns). 1315 success = np.array([[success_probability(L, d, step_error, lucky) for L in LENGTHS] for d in DIFFICULTIES]) 1316 1317 def policy() -> np.ndarray: 1318 e = np.exp(logits - logits.max(axis=1, keepdims=True)) 1319 return e / e.sum(axis=1, keepdims=True) 1320 1321 def record(history: dict) -> None: 1322 pi = policy() 1323 # Expected values under the current policy, so the curves aren't sampling noise. 1324 history["mean_length"].append(float((pi @ lengths).mean())) 1325 history["accuracy"].append(float((pi * success).sum(axis=1).mean())) 1326 1327 history: dict = dict(mean_length=[], accuracy=[]) 1328 record(history) 1329 for _ in range(iterations): 1330 pi = policy() 1331 for row in range(len(DIFFICULTIES)): 1332 # A group of G attempts at the same problem, each thinking for a sampled length. 1333 actions = rng.choice(len(LENGTHS), size=group_size, p=pi[row]) 1334 solved = rng.random(group_size) < success[row, actions] 1335 rewards = solved.astype(float) - length_penalty * lengths[actions] 1336 adv = group_advantages(rewards, divide_by_spread) 1337 # d log pi(a) / d logits = onehot(a) - pi: raise the chosen length in proportion to its advantage. 1338 onehot = np.eye(len(LENGTHS))[actions] 1339 logits[row] += lr * (adv[:, None] * (onehot - pi[row])).mean(axis=0) 1340 record(history) 1341 pi = policy() 1342 history["policy"] = pi 1343 history["preferred_length"] = {d: int(lengths[np.argmax(pi[i])]) for i, d in enumerate(DIFFICULTIES)} 1344 history["expected_length"] = {d: float(pi[i] @ lengths) for i, d in enumerate(DIFFICULTIES)} 1345 return history 1346 1347 1348# --------------------------------------------------------------------------- 1349# 6. What reasoning costs, and where it still fails 1350# --------------------------------------------------------------------------- 1351 1352 1353def steps_all_right(step_error: float, steps: int) -> float: 1354 """Chance a chain of `steps` steps has no slip at all: (1 - e)^k.""" 1355 return (1 - step_error) ** steps 1356 1357 1358def reasoning_cost(samples: int, tokens: int, dollars_per_million: float, tokens_per_second: float) -> dict: 1359 """Money and waiting time for one question. 1360 1361 Money grows with every token of every sample. Waiting time, when the 1362 samples run side by side, is set by the length of one chain. 1363 """ 1364 return dict( 1365 dollars=samples * tokens * dollars_per_million / 1_000_000, 1366 seconds=tokens / tokens_per_second, 1367 ) 1368 1369 1370# --------------------------------------------------------------------------- 1371# 7. Figures (rendered into the HTML docs by `make figures`) 1372# --------------------------------------------------------------------------- 1373 1374 1375def _direct_vs_cot(max_additions: int = 10, per_size: int = 300, seed: int = 0) -> list[dict]: 1376 """Accuracy and tokens for answering at once vs. writing every step, by problem size.""" 1377 rng = np.random.default_rng(seed) 1378 rows = [] 1379 for d in range(1, max_additions + 1): 1380 problems = [rng.integers(0, 10, d + 1).tolist() for _ in range(per_size)] 1381 direct = [solve(p, 1) for p in problems] 1382 cot = [solve(p, 64) for p in problems] 1383 rows.append(dict( 1384 additions=d, 1385 direct=float(np.mean([a.answer == sum(p) for a, p in zip(direct, problems)])), 1386 cot=float(np.mean([a.answer == sum(p) for a, p in zip(cot, problems)])), 1387 cot_tokens=float(np.mean([a.tokens for a in cot])), 1388 )) 1389 return rows 1390 1391 1392def figures() -> dict: 1393 """Plot this lesson's data. matplotlib is imported here, and only here, 1394 so the lesson itself needs nothing beyond NumPy.""" 1395 import matplotlib 1396 1397 matplotlib.use("Agg") 1398 import matplotlib.pyplot as plt 1399 1400 BLUE, RED, GREEN, AMBER, PURPLE, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#7c3aed", "#9ca3af" 1401 figs = {} 1402 1403 # --- 1. Answering at once vs. thinking out loud -------------------------- 1404 rows = _direct_vs_cot() 1405 ds = [r["additions"] for r in rows] 1406 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1407 a1.plot(ds, [r["direct"] for r in rows], "o-", color=RED, label="answer in 1 token") 1408 a1.plot(ds, [r["cot"] for r in rows], "o-", color=BLUE, label="write steps, then answer") 1409 a1.set_xlabel("additions the problem needs") 1410 a1.set_ylabel("accuracy") 1411 a1.set_ylim(-0.03, 1.05) 1412 a1.set_title("One addition per pass: steps are the only way") 1413 a1.legend(frameon=False) 1414 a2.plot(ds, [1] * len(ds), "o-", color=RED, label="answer in 1 token") 1415 a2.plot(ds, [r["cot_tokens"] for r in rows], "o-", color=BLUE, label="write steps, then answer") 1416 a2.set_xlabel("additions the problem needs") 1417 a2.set_ylabel("tokens (forward passes) used") 1418 a2.set_title("The price: one token per step") 1419 fig.tight_layout() 1420 figs["direct_vs_cot"] = fig 1421 1422 # --- 2. Accuracy vs. thinking budget -------------------------------------- 1423 budgets = (1, 2, 4, 8, 16, 32, 64) 1424 sweep = budget_sweep(budgets, n_problems=600, seed=0) 1425 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1426 a1.plot(budgets, [budget_accuracy(b) for b in budgets], "-", color=MUTED, label="formula, g = 0.1") 1427 a1.plot(budgets, [r["accuracy"] for r in sweep], "o", color=BLUE, label="simulated sums") 1428 a1.set_xscale("log", base=2) 1429 a1.set_xlabel("thinking budget B (tokens, log scale)") 1430 a1.set_ylabel("accuracy") 1431 a1.set_ylim(0, 1.05) 1432 a1.set_title("Each doubling of budget buys the same step") 1433 a1.legend(frameon=False) 1434 a2.plot(budgets, [r["mean_tokens"] for r in sweep], "o-", color=AMBER) 1435 a2.plot(budgets, budgets, ":", color=MUTED, label="budget allowed") 1436 a2.set_xscale("log", base=2) 1437 a2.set_yscale("log", base=2) 1438 a2.set_xlabel("thinking budget B (tokens, log scale)") 1439 a2.set_ylabel("tokens actually spent (mean)") 1440 a2.set_title("Easy problems stop early; past 32, nothing changes") 1441 a2.legend(frameon=False) 1442 fig.tight_layout() 1443 figs["budget_scaling"] = fig 1444 1445 # --- 3. Voting: the binomial formula and scattered wrong answers ---------- 1446 ns = [1, 3, 5, 7, 9, 15, 21, 31] 1447 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1448 for p, color in ((0.8, GREEN), (0.6, BLUE), (0.4, RED)): 1449 ax.plot(ns, [majority_accuracy(p, n) for n in ns], "o-", color=color, label=f"yes/no question, p = {p}") 1450 ax.plot(ns, [correlated_vote_accuracy(0.4, n, rho=0.0, n_wrong=10, trials=1500) for n in ns], "s--", color=RED, 1451 label="p = 0.4, wrong answers scattered over 10 values") 1452 ax.set_xlabel("samples voting, n") 1453 ax.set_ylabel("accuracy of the vote") 1454 ax.set_ylim(0, 1.05) 1455 ax.set_title("Voting amplifies whatever answer is most common") 1456 ax.legend(frameon=False, fontsize=8) 1457 figs["voting"] = fig 1458 1459 # --- 4. Correlated mistakes put a ceiling on voting ----------------------- 1460 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1461 for rho, color in ((0.0, BLUE), (0.3, GREEN), (0.6, AMBER), (0.9, RED)): 1462 ax.plot(ns, [correlated_vote_accuracy(0.6, n, rho=rho, n_wrong=5, trials=1500) for n in ns], "o-", color=color, 1463 label=f"shared-mistake rate rho = {rho}") 1464 ax.axhline(0.6, color=MUTED, ls=":") 1465 ax.text(12, 0.56, "one sample alone: 0.6", color="#4b5563", fontsize=8) 1466 ax.set_xlabel("samples voting, n") 1467 ax.set_ylabel("accuracy of the vote") 1468 ax.set_ylim(0.2, 1.03) 1469 ax.set_title("Same per-sample accuracy, very different votes") 1470 ax.legend(frameon=False, fontsize=8, loc="lower right") 1471 figs["correlated_votes"] = fig 1472 1473 # --- 5. Best-of-n with a verifier ------------------------------------------ 1474 vns = (1, 2, 4, 8, 16, 32) 1475 vrows = verifier_experiment(vns, trials=400, seed=0) 1476 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1477 ax.plot(vns, [r["pass_at_n"] for r in vrows], ":", color=MUTED, lw=2, label="pass@n: any sample right (ceiling)") 1478 ax.plot(vns, [r["verifier"] for r in vrows], "o-", color=GREEN, label="best-of-n, step checker picks") 1479 ax.plot(vns, [r["vote"] for r in vrows], "o-", color=BLUE, label="majority vote") 1480 ax.plot(vns, [r["single"] for r in vrows], "o-", color=RED, label="one sample") 1481 ax.set_xscale("log", base=2) 1482 ax.set_xlabel("chains sampled, n (log scale)") 1483 ax.set_ylabel("accuracy") 1484 ax.set_ylim(0, 1.05) 1485 ax.set_title("A good checker turns many tries into a right answer") 1486 ax.legend(frameon=False, fontsize=8, loc="lower right") 1487 figs["verifier"] = fig 1488 1489 # --- 6. RL: learning how long to think ------------------------------------ 1490 runs = ( 1491 ("correctness only", dict(), BLUE), 1492 ("penalty 0.01 per token, ÷ spread (GRPO)", dict(length_penalty=0.01), RED), 1493 ("penalty 0.01 per token, no ÷ spread", dict(length_penalty=0.01, divide_by_spread=False), GREEN), 1494 ) 1495 histories = [(label, train_reasoner(seed=0, **kw), color) for label, kw, color in runs] 1496 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1497 for label, h, color in histories: 1498 a1.plot(h["mean_length"], color=color, label=label) 1499 a2.plot(h["accuracy"], color=color, label=label) 1500 a1.set_xlabel("training iteration") 1501 a1.set_ylabel("mean thinking length (tokens)") 1502 a1.set_title("It learns to write longer") 1503 a2.set_xlabel("training iteration") 1504 a2.set_ylabel("accuracy") 1505 a2.set_ylim(0.35, 1.0) 1506 a2.set_title("...and gets more right") 1507 a2.legend(frameon=False, fontsize=8, loc="lower right") 1508 fig.tight_layout() 1509 figs["rl_training"] = fig 1510 1511 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1512 x = np.arange(len(DIFFICULTIES)) 1513 width = 0.26 1514 for i, (label, h, color) in enumerate(histories): 1515 ax.bar(x + (i - 1) * width, [h["expected_length"][d] for d in DIFFICULTIES], width, color=color, label=label) 1516 ax.plot(x, DIFFICULTIES, "k_", ms=28, mew=2, label="steps needed d") 1517 ax.plot(x, [2 * d for d in DIFFICULTIES], "_", color="#4b5563", ms=28, mew=1.5, ls="none", label="2d: room for one recheck") 1518 ax.set_xticks(x, [f"d = {d}" for d in DIFFICULTIES]) 1519 ax.set_ylim(0, 18) 1520 ax.set_xlabel("difficulty: additions the problem needs") 1521 ax.set_ylabel("tokens the trained model spends") 1522 ax.set_title("Hard problems get more thinking, with room to recheck") 1523 ax.legend(frameon=False, fontsize=8, loc="upper left") 1524 figs["rl_lengths"] = fig 1525 1526 # --- 7. Compounding slips over long chains -------------------------------- 1527 ks = np.arange(1, 201) 1528 fig, ax = plt.subplots(figsize=(6.4, 3.4)) 1529 for e, color in ((0.005, GREEN), (0.02, BLUE), (0.05, RED)): 1530 ax.plot(ks, [steps_all_right(e, k) for k in ks], color=color, label=f"slip rate per step = {e}") 1531 ax.plot([50], [steps_all_right(0.02, 50)], "o", color=BLUE) 1532 ax.annotate("50 steps at 2%: 0.36", (50, steps_all_right(0.02, 50)), xytext=(70, 0.55), arrowprops=dict(arrowstyle="->", color="#4b5563"), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1533 ax.set_xlabel("steps in the chain, k") 1534 ax.set_ylabel("chance every step is right") 1535 ax.set_ylim(0, 1.03) 1536 ax.set_title("Long chains need checking, not just length") 1537 ax.legend(frameon=False) 1538 figs["compounding"] = fig 1539 1540 return figs 1541 1542 1543# --------------------------------------------------------------------------- 1544# 8. Narrated walkthrough 1545# --------------------------------------------------------------------------- 1546 1547 1548def demo() -> None: 1549 banner("1. Thinking out loud: one addition per forward pass") 1550 say( 1551 """ 1552 Our toy model can do exactly one addition per forward pass, and every 1553 token it writes is one forward pass. Ask it for 3 + 5 + 8 + 2. 1554 """ 1555 ) 1556 direct = solve([3, 5, 8, 2], max_tokens=1) 1557 cot = solve([3, 5, 8, 2], max_tokens=10) 1558 table( 1559 ["how it answers", "what it wrote", "answer", "tokens"], 1560 [ 1561 ("at once", "(nothing)", direct.answer, direct.tokens), 1562 ("thinking first", ", ".join(cot.steps), cot.answer, cot.tokens), 1563 ], 1564 ) 1565 say( 1566 """ 1567 At once, it adds 3 + 5 = 8 and has no pass left for 8 and 2, so it 1568 estimates them at 4.5 each: 17, wrong. Thinking first, it writes each 1569 running total where the next pass can read it: 18, right, for 3 tokens. 1570 """ 1571 ) 1572 takeaway("Each written token is another forward pass: writing steps buys serial computation the weights alone don't have.") 1573 1574 banner("2. Test-time compute: accuracy vs. thinking budget") 1575 rows = budget_sweep() 1576 table( 1577 ["budget B", "formula acc(B)", "simulated accuracy", "tokens actually spent"], 1578 [(r["budget"], budget_accuracy(r["budget"]), r["accuracy"], r["mean_tokens"]) for r in rows], 1579 floatfmt=".3f", 1580 ) 1581 say( 1582 """ 1583 Problems need 1, 2, 4, 8, 16 or 32 additions. Each doubling of the 1584 budget lets one more difficulty level fit, so accuracy climbs by a 1585 steady step per doubling, then stops once the hardest problem fits. 1586 """ 1587 ) 1588 takeaway("With difficulty spread over many scales, accuracy grows with the log of the thinking budget, until it doesn't.") 1589 1590 banner("3. Sample many, vote: self-consistency") 1591 say("Three sampled answers to 3 + 5 + 8 + 2: 18, 17, 18. The vote picks " + str(majority_vote([18, 17, 18])) + ".") 1592 table( 1593 ["samples n", "yes/no, p = 0.6", "yes/no, p = 0.4", "p = 0.6, rho = 0.9 (shared mistakes)"], 1594 [(n, majority_accuracy(0.6, n), majority_accuracy(0.4, n), correlated_vote_accuracy(0.6, n, rho=0.9)) for n in (1, 3, 5, 15, 31)], 1595 floatfmt=".3f", 1596 ) 1597 say( 1598 """ 1599 Independent 60% votes climb towards certainty; independent 40% votes on 1600 a yes/no question sink. When samples share their mistakes (rho = 0.9), 1601 31 votes are barely better than one: they are one opinion, repeated. 1602 """ 1603 ) 1604 takeaway("Voting helps only when the right answer is the most common one and the mistakes are independent.") 1605 1606 banner("4. Verifiers: check the answer, or check every step") 1607 slipped = ["3+5=8", "8+8=17", "17+2=19"] 1608 lucky = ["3+5=9", "9+8=16", "16+2=18"] 1609 table( 1610 ["chain", "outcome reward (answer 18?)", "first bad step (process check)"], 1611 [(", ".join(c), int(outcome_reward(c, 18)), first_bad_step(c)) for c in (slipped, lucky)], 1612 ) 1613 say( 1614 """ 1615 The outcome check only sees the last number: it gives the lucky chain, 1616 whose two slips cancelled, full marks. The step check reads every line 1617 and points at the exact slip. 1618 """ 1619 ) 1620 table( 1621 ["chains n", "one sample", "vote", "best-of-n, step checker", "pass@n ceiling"], 1622 [(r["n"], r["single"], r["vote"], r["verifier"], r["pass_at_n"]) for r in verifier_experiment()], 1623 floatfmt=".3f", 1624 ) 1625 takeaway("A reliable checker turns 'right sometimes' into 'right almost always' with enough samples; voting gets there far more slowly.") 1626 1627 banner("5. Learning to reason with RL: verifiable rewards, group advantages") 1628 say(f"Group rewards [1, 0, 0, 1] give advantages {group_advantages(np.array([1.0, 0.0, 0.0, 1.0])).tolist()}.") 1629 for label, kw in ( 1630 ("correctness only", dict()), 1631 ("penalty 0.01/token, GRPO", dict(length_penalty=0.01)), 1632 ("penalty 0.01/token, no ÷ spread", dict(length_penalty=0.01, divide_by_spread=False)), 1633 ): 1634 h = train_reasoner(seed=0, **kw) 1635 print( 1636 f"{label:32s} length {h['mean_length'][0]:.1f} -> {h['mean_length'][-1]:.1f} tokens, " 1637 f"accuracy {h['accuracy'][0]:.2f} -> {h['accuracy'][-1]:.2f}, " 1638 f"prefers {h['preferred_length']}" 1639 ) 1640 print() 1641 say( 1642 """ 1643 Rewarded only for right answers, the policy learns to think longer 1644 (about 2 tokens to about 7.6) and to leave room for a recheck on hard 1645 problems (16 tokens for 8 additions). A per-token penalty trims easy 1646 problems; dividing by the group's spread magnifies that penalty until 1647 1-addition problems get 1 token and accept an occasional slip. 1648 """ 1649 ) 1650 takeaway("Reward the outcome and let the model search: longer, self-checking reasoning emerges because it pays.") 1651 1652 banner("6. What it costs, and where it fails") 1653 table( 1654 ["slip rate", "steps", "chance all right"], 1655 [(e, k, steps_all_right(e, k)) for e, k in ((0.02, 10), (0.02, 50), (0.02, 200), (0.005, 200))], 1656 floatfmt=".3f", 1657 ) 1658 c1 = reasoning_cost(1, 2000, 10.0, 50) 1659 c8 = reasoning_cost(8, 2000, 10.0, 50) 1660 say( 1661 f""" 1662 One 2,000-token chain at $10 per million tokens costs ${c1['dollars']:.2f} 1663 and takes {c1['seconds']:.0f} s at 50 tokens per second. Eight in parallel 1664 cost ${c8['dollars']:.2f} and still take {c8['seconds']:.0f} s. 1665 """ 1666 ) 1667 takeaway("Thinking is paid per token: spend it where a check shows it helps, and cap it everywhere else.") 1668 1669 1670if __name__ == "__main__": 1671 demo()
1043@dataclass 1044class Attempt: 1045 """What the toy model wrote for one problem. 1046 1047 steps: the intermediate lines it wrote before answering, e.g. ["3+5=8", "8+8=16"] 1048 answer: the number it finally gave 1049 tokens: forward passes used, one per written step plus one for the answer 1050 """ 1051 1052 steps: list[str] = field(default_factory=list) 1053 answer: int = 0 1054 tokens: int = 1
What the toy model wrote for one problem.
steps: the intermediate lines it wrote before answering, e.g. ["3+5=8", "8+8=16"] answer: the number it finally gave tokens: forward passes used, one per written step plus one for the answer
1057def solve(numbers: list[int], max_tokens: int) -> Attempt: 1058 """Add up `numbers` with a model that can do ONE addition per forward pass. 1059 1060 Every token the model writes is one forward pass. A step token ("8+8=16") 1061 spends its pass on one addition and leaves the running total in the text, 1062 where the next pass can read it. The answer token also gets one addition. 1063 If the budget runs out before the sum is done, the model estimates the 1064 numbers it never reached at `DIGIT_MEAN` each: a guess, right only by luck. 1065 """ 1066 total = numbers[0] 1067 if len(numbers) == 1: 1068 return Attempt([], total, 1) 1069 steps: list[str] = [] 1070 i = 1 1071 # Thinking tokens: all but the last addition, as long as the budget leaves a token for the answer. 1072 while i < len(numbers) - 1 and len(steps) < max_tokens - 1: 1073 new = total + numbers[i] 1074 steps.append(f"{total}+{numbers[i]}={new}") 1075 total, i = new, i + 1 1076 # The answer token's own forward pass: one more addition. 1077 total += numbers[i] 1078 unreached = len(numbers) - (i + 1) 1079 # Round half up, so one unreached number is estimated as 5 (4.5 rounded). 1080 answer = total + math.floor(DIGIT_MEAN * unreached + 0.5) 1081 return Attempt(steps, int(answer), len(steps) + 1)
Add up numbers with a model that can do ONE addition per forward pass.
Every token the model writes is one forward pass. A step token ("8+8=16")
spends its pass on one addition and leaves the running total in the text,
where the next pass can read it. The answer token also gets one addition.
If the budget runs out before the sum is done, the model estimates the
numbers it never reached at DIGIT_MEAN each: a guess, right only by luck.
1092def budget_accuracy(budget: int, needs: tuple[int, ...] = NEEDS, lucky: float = 0.1) -> float: 1093 """acc(B) = F(B) + (1 - F(B)) * g. 1094 1095 F(B) is the share of problems that fit in a budget of B tokens (need at 1096 most B additions); the rest are guessed, and a guess is right with 1097 probability g (`lucky`). 1098 """ 1099 fits = sum(d <= budget for d in needs) / len(needs) 1100 return fits + (1 - fits) * lucky
acc(B) = F(B) + (1 - F(B)) * g.
F(B) is the share of problems that fit in a budget of B tokens (need at
most B additions); the rest are guessed, and a guess is right with
probability g (lucky).
1103def budget_sweep(budgets=(1, 2, 4, 8, 16, 32, 64), n_problems: int = 600, seed: int = 0) -> list[dict]: 1104 """Run `solve` on real random sums at each budget: accuracy and tokens actually spent. 1105 1106 Each problem needs d additions, d drawn evenly from `NEEDS`, so it has d + 1 1107 random digits. The same problems are reused at every budget. 1108 """ 1109 rng = np.random.default_rng(seed) 1110 problems = [rng.integers(0, 10, size=int(rng.choice(NEEDS)) + 1).tolist() for _ in range(n_problems)] 1111 rows = [] 1112 for b in budgets: 1113 attempts = [solve(p, b) for p in problems] 1114 right = [a.answer == sum(p) for a, p in zip(attempts, problems)] 1115 rows.append(dict(budget=b, accuracy=float(np.mean(right)), mean_tokens=float(np.mean([a.tokens for a in attempts])))) 1116 return rows
1124def majority_vote(answers: list) -> object: 1125 """The most common answer. Ties go to the answer seen first.""" 1126 return Counter(answers).most_common(1)[0][0]
The most common answer. Ties go to the answer seen first.
1129def majority_accuracy(p: float, n: int) -> float: 1130 """Chance that more than half of n independent votes are right, each right with probability p. 1131 1132 A yes/no question: every wrong vote lands on the same wrong answer. With 1133 an even n, a tie is broken by a coin flip, so it counts half. 1134 """ 1135 win = sum(math.comb(n, k) * p**k * (1 - p) ** (n - k) for k in range(n // 2 + 1, n + 1)) 1136 tie = math.comb(n, n // 2) * p ** (n // 2) * (1 - p) ** (n // 2) if n % 2 == 0 else 0.0 1137 return win + tie / 2
Chance that more than half of n independent votes are right, each right with probability p.
A yes/no question: every wrong vote lands on the same wrong answer. With an even n, a tie is broken by a coin flip, so it counts half.
1172def check_step(step: str) -> bool: 1173 """A process verifier for the toy: does this one line "a+b=c" hold?""" 1174 m = STEP.match(step) 1175 return bool(m) and int(m[1]) + int(m[2]) == int(m[3])
A process verifier for the toy: does this one line "a+b=c" hold?
1178def first_bad_step(steps: list[str]) -> int | None: 1179 """Index of the first step the process verifier rejects, or None if every step checks out.""" 1180 return next((i for i, s in enumerate(steps) if not check_step(s)), None)
Index of the first step the process verifier rejects, or None if every step checks out.
1183def final_answer(steps: list[str]) -> int: 1184 """The number after the last "=": the chain's answer.""" 1185 return int(steps[-1].rsplit("=", 1)[1])
The number after the last "=": the chain's answer.
1188def outcome_reward(steps: list[str], reference: int) -> float: 1189 """An outcome check: 1 if the final answer matches the reference, else 0. It never reads the steps.""" 1190 return 1.0 if final_answer(steps) == reference else 0.0
An outcome check: 1 if the final answer matches the reference, else 0. It never reads the steps.
1193def noisy_chain(numbers: list[int], step_error: float, rng: np.random.Generator) -> list[str]: 1194 """Write every addition as a step; each slips by ±1 or ±2 with probability `step_error`. 1195 1196 A slip is carried forward: the next step adds to the wrong running total, 1197 which is how one early mistake poisons the final answer. 1198 """ 1199 total, steps = numbers[0], [] 1200 for x in numbers[1:]: 1201 new = total + x 1202 if rng.random() < step_error: 1203 new += int(rng.choice([-2, -1, 1, 2])) 1204 steps.append(f"{total}+{x}={new}") 1205 total = new 1206 return steps
Write every addition as a step; each slips by ±1 or ±2 with probability step_error.
A slip is carried forward: the next step adds to the wrong running total, which is how one early mistake poisons the final answer.
1209def pass_at_n(p: float, n: int) -> float: 1210 """Chance that at least one of n independent samples is right: 1 - (1 - p)^n.""" 1211 return 1 - (1 - p) ** n
Chance that at least one of n independent samples is right: 1 - (1 - p)^n.
1214def verifier_experiment( 1215 ns=(1, 2, 4, 8, 16, 32), n_terms: int = 6, step_error: float = 0.2, trials: int = 400, seed: int = 0 1216) -> list[dict]: 1217 """Compare four ways to turn n sampled chains into one answer. 1218 1219 - single: the first sample, as if n were 1. 1220 - vote: the most common final answer (self-consistency). 1221 - verifier: the first chain whose every step checks out (best-of-n with a 1222 process verifier), falling back to the vote if none does. 1223 - pass_at_n: whether ANY chain's final answer is right, the ceiling a 1224 perfect answer-picker could reach. 1225 """ 1226 rng = np.random.default_rng(seed) 1227 hits = {n: dict(single=0, vote=0, verifier=0, pass_at_n=0) for n in ns} 1228 for _ in range(trials): 1229 numbers = rng.integers(0, 10, n_terms).tolist() 1230 truth = sum(numbers) 1231 chains = [noisy_chain(numbers, step_error, rng) for _ in range(max(ns))] 1232 for n in ns: 1233 pool = chains[:n] 1234 answers = [final_answer(c) for c in pool] 1235 clean = [c for c in pool if first_bad_step(c) is None] 1236 picked = final_answer(clean[0]) if clean else majority_vote(answers) 1237 h = hits[n] 1238 h["single"] += answers[0] == truth 1239 h["vote"] += majority_vote(answers) == truth 1240 h["verifier"] += picked == truth 1241 h["pass_at_n"] += truth in answers 1242 return [dict(n=n, **{k: v / trials for k, v in hits[n].items()}) for n in ns]
Compare four ways to turn n sampled chains into one answer.
- single: the first sample, as if n were 1.
- vote: the most common final answer (self-consistency).
- verifier: the first chain whose every step checks out (best-of-n with a process verifier), falling back to the vote if none does.
- pass_at_n: whether ANY chain's final answer is right, the ceiling a perfect answer-picker could reach.
1254def group_advantages(rewards: np.ndarray, divide_by_spread: bool = True) -> np.ndarray: 1255 """GRPO's advantage: each reward compared with its own group, A_i = (r_i - mean) / std. 1256 1257 A group where every sample scored the same carries no information about 1258 which choice was better, so every advantage is 0. With 1259 `divide_by_spread=False` the advantage is just r_i - mean, which keeps a 1260 tiny reward difference tiny instead of blowing it up to ±1. 1261 """ 1262 rewards = np.asarray(rewards, dtype=float) 1263 centred = rewards - rewards.mean() 1264 spread = rewards.std() 1265 if not divide_by_spread: 1266 return centred 1267 if spread == 0: 1268 return np.zeros_like(rewards) 1269 return centred / spread
GRPO's advantage: each reward compared with its own group, A_i = (r_i - mean) / std.
A group where every sample scored the same carries no information about
which choice was better, so every advantage is 0. With
divide_by_spread=False the advantage is just r_i - mean, which keeps a
tiny reward difference tiny instead of blowing it up to ±1.
1272def success_probability(length: int, needs: int, step_error: float = 0.1, lucky: float = 0.1) -> float: 1273 """How often a chain of `length` tokens solves a problem needing `needs` steps. 1274 1275 Too short (length < needs): it must guess, right with probability `lucky`. 1276 Otherwise every step gets `length // needs` tries: the first attempt plus 1277 rechecks that can catch and fix a slip. A step fails only if every try 1278 slips, so the chain succeeds with probability (1 - e^tries)^needs. More 1279 length always helps a little, which is why a reward for correctness alone 1280 never tells the model to stop. 1281 """ 1282 if length < needs: 1283 return lucky 1284 tries = length // needs 1285 return (1 - step_error**tries) ** needs
How often a chain of length tokens solves a problem needing needs steps.
Too short (length < needs): it must guess, right with probability lucky.
Otherwise every step gets length // needs tries: the first attempt plus
rechecks that can catch and fix a slip. A step fails only if every try
slips, so the chain succeeds with probability (1 - e^tries)^needs. More
length always helps a little, which is why a reward for correctness alone
never tells the model to stop.
1288def train_reasoner( 1289 iterations: int = 600, 1290 group_size: int = 16, 1291 lr: float = 0.2, 1292 length_penalty: float = 0.0, 1293 step_error: float = 0.1, 1294 lucky: float = 0.1, 1295 divide_by_spread: bool = True, 1296 seed: int = 0, 1297) -> dict: 1298 """Train a tiny policy to choose how long to think, from outcome rewards only. 1299 1300 The policy is a table of logits: one row per difficulty, one column per 1301 thinking length in `LENGTHS`. It starts out preferring short answers, like 1302 a model that was never rewarded for thinking. Each iteration, for each 1303 difficulty, it samples a group of lengths, scores each one with a 1304 verifiable reward (1 if right, 0 if wrong, minus `length_penalty` per 1305 token), turns the rewards into group-relative advantages and nudges the 1306 logits along the policy gradient: logits += lr * mean(A_i * (onehot_i - pi)). 1307 1308 Returns the expected mean length and accuracy after every iteration, and 1309 the length each difficulty ends up preferring. 1310 """ 1311 rng = np.random.default_rng(seed) 1312 lengths = np.array(LENGTHS) 1313 # Short answers favoured at the start: pi(1 token) ≈ 0.63, pi(32 tokens) ≈ 0.004. 1314 logits = np.tile(-INIT_TILT * np.arange(len(LENGTHS)), (len(DIFFICULTIES), 1)) 1315 # success[d, a]: how often length a solves difficulty d (d rows, a columns). 1316 success = np.array([[success_probability(L, d, step_error, lucky) for L in LENGTHS] for d in DIFFICULTIES]) 1317 1318 def policy() -> np.ndarray: 1319 e = np.exp(logits - logits.max(axis=1, keepdims=True)) 1320 return e / e.sum(axis=1, keepdims=True) 1321 1322 def record(history: dict) -> None: 1323 pi = policy() 1324 # Expected values under the current policy, so the curves aren't sampling noise. 1325 history["mean_length"].append(float((pi @ lengths).mean())) 1326 history["accuracy"].append(float((pi * success).sum(axis=1).mean())) 1327 1328 history: dict = dict(mean_length=[], accuracy=[]) 1329 record(history) 1330 for _ in range(iterations): 1331 pi = policy() 1332 for row in range(len(DIFFICULTIES)): 1333 # A group of G attempts at the same problem, each thinking for a sampled length. 1334 actions = rng.choice(len(LENGTHS), size=group_size, p=pi[row]) 1335 solved = rng.random(group_size) < success[row, actions] 1336 rewards = solved.astype(float) - length_penalty * lengths[actions] 1337 adv = group_advantages(rewards, divide_by_spread) 1338 # d log pi(a) / d logits = onehot(a) - pi: raise the chosen length in proportion to its advantage. 1339 onehot = np.eye(len(LENGTHS))[actions] 1340 logits[row] += lr * (adv[:, None] * (onehot - pi[row])).mean(axis=0) 1341 record(history) 1342 pi = policy() 1343 history["policy"] = pi 1344 history["preferred_length"] = {d: int(lengths[np.argmax(pi[i])]) for i, d in enumerate(DIFFICULTIES)} 1345 history["expected_length"] = {d: float(pi[i] @ lengths) for i, d in enumerate(DIFFICULTIES)} 1346 return history
Train a tiny policy to choose how long to think, from outcome rewards only.
The policy is a table of logits: one row per difficulty, one column per
thinking length in LENGTHS. It starts out preferring short answers, like
a model that was never rewarded for thinking. Each iteration, for each
difficulty, it samples a group of lengths, scores each one with a
verifiable reward (1 if right, 0 if wrong, minus length_penalty per
token), turns the rewards into group-relative advantages and nudges the
logits along the policy gradient: logits += lr * mean(A_i * (onehot_i - pi)).
Returns the expected mean length and accuracy after every iteration, and the length each difficulty ends up preferring.
1354def steps_all_right(step_error: float, steps: int) -> float: 1355 """Chance a chain of `steps` steps has no slip at all: (1 - e)^k.""" 1356 return (1 - step_error) ** steps
Chance a chain of steps steps has no slip at all: (1 - e)^k.
1359def reasoning_cost(samples: int, tokens: int, dollars_per_million: float, tokens_per_second: float) -> dict: 1360 """Money and waiting time for one question. 1361 1362 Money grows with every token of every sample. Waiting time, when the 1363 samples run side by side, is set by the length of one chain. 1364 """ 1365 return dict( 1366 dollars=samples * tokens * dollars_per_million / 1_000_000, 1367 seconds=tokens / tokens_per_second, 1368 )
Money and waiting time for one question.
Money grows with every token of every sample. Waiting time, when the samples run side by side, is set by the length of one chain.
1393def figures() -> dict: 1394 """Plot this lesson's data. matplotlib is imported here, and only here, 1395 so the lesson itself needs nothing beyond NumPy.""" 1396 import matplotlib 1397 1398 matplotlib.use("Agg") 1399 import matplotlib.pyplot as plt 1400 1401 BLUE, RED, GREEN, AMBER, PURPLE, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#7c3aed", "#9ca3af" 1402 figs = {} 1403 1404 # --- 1. Answering at once vs. thinking out loud -------------------------- 1405 rows = _direct_vs_cot() 1406 ds = [r["additions"] for r in rows] 1407 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1408 a1.plot(ds, [r["direct"] for r in rows], "o-", color=RED, label="answer in 1 token") 1409 a1.plot(ds, [r["cot"] for r in rows], "o-", color=BLUE, label="write steps, then answer") 1410 a1.set_xlabel("additions the problem needs") 1411 a1.set_ylabel("accuracy") 1412 a1.set_ylim(-0.03, 1.05) 1413 a1.set_title("One addition per pass: steps are the only way") 1414 a1.legend(frameon=False) 1415 a2.plot(ds, [1] * len(ds), "o-", color=RED, label="answer in 1 token") 1416 a2.plot(ds, [r["cot_tokens"] for r in rows], "o-", color=BLUE, label="write steps, then answer") 1417 a2.set_xlabel("additions the problem needs") 1418 a2.set_ylabel("tokens (forward passes) used") 1419 a2.set_title("The price: one token per step") 1420 fig.tight_layout() 1421 figs["direct_vs_cot"] = fig 1422 1423 # --- 2. Accuracy vs. thinking budget -------------------------------------- 1424 budgets = (1, 2, 4, 8, 16, 32, 64) 1425 sweep = budget_sweep(budgets, n_problems=600, seed=0) 1426 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1427 a1.plot(budgets, [budget_accuracy(b) for b in budgets], "-", color=MUTED, label="formula, g = 0.1") 1428 a1.plot(budgets, [r["accuracy"] for r in sweep], "o", color=BLUE, label="simulated sums") 1429 a1.set_xscale("log", base=2) 1430 a1.set_xlabel("thinking budget B (tokens, log scale)") 1431 a1.set_ylabel("accuracy") 1432 a1.set_ylim(0, 1.05) 1433 a1.set_title("Each doubling of budget buys the same step") 1434 a1.legend(frameon=False) 1435 a2.plot(budgets, [r["mean_tokens"] for r in sweep], "o-", color=AMBER) 1436 a2.plot(budgets, budgets, ":", color=MUTED, label="budget allowed") 1437 a2.set_xscale("log", base=2) 1438 a2.set_yscale("log", base=2) 1439 a2.set_xlabel("thinking budget B (tokens, log scale)") 1440 a2.set_ylabel("tokens actually spent (mean)") 1441 a2.set_title("Easy problems stop early; past 32, nothing changes") 1442 a2.legend(frameon=False) 1443 fig.tight_layout() 1444 figs["budget_scaling"] = fig 1445 1446 # --- 3. Voting: the binomial formula and scattered wrong answers ---------- 1447 ns = [1, 3, 5, 7, 9, 15, 21, 31] 1448 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1449 for p, color in ((0.8, GREEN), (0.6, BLUE), (0.4, RED)): 1450 ax.plot(ns, [majority_accuracy(p, n) for n in ns], "o-", color=color, label=f"yes/no question, p = {p}") 1451 ax.plot(ns, [correlated_vote_accuracy(0.4, n, rho=0.0, n_wrong=10, trials=1500) for n in ns], "s--", color=RED, 1452 label="p = 0.4, wrong answers scattered over 10 values") 1453 ax.set_xlabel("samples voting, n") 1454 ax.set_ylabel("accuracy of the vote") 1455 ax.set_ylim(0, 1.05) 1456 ax.set_title("Voting amplifies whatever answer is most common") 1457 ax.legend(frameon=False, fontsize=8) 1458 figs["voting"] = fig 1459 1460 # --- 4. Correlated mistakes put a ceiling on voting ----------------------- 1461 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1462 for rho, color in ((0.0, BLUE), (0.3, GREEN), (0.6, AMBER), (0.9, RED)): 1463 ax.plot(ns, [correlated_vote_accuracy(0.6, n, rho=rho, n_wrong=5, trials=1500) for n in ns], "o-", color=color, 1464 label=f"shared-mistake rate rho = {rho}") 1465 ax.axhline(0.6, color=MUTED, ls=":") 1466 ax.text(12, 0.56, "one sample alone: 0.6", color="#4b5563", fontsize=8) 1467 ax.set_xlabel("samples voting, n") 1468 ax.set_ylabel("accuracy of the vote") 1469 ax.set_ylim(0.2, 1.03) 1470 ax.set_title("Same per-sample accuracy, very different votes") 1471 ax.legend(frameon=False, fontsize=8, loc="lower right") 1472 figs["correlated_votes"] = fig 1473 1474 # --- 5. Best-of-n with a verifier ------------------------------------------ 1475 vns = (1, 2, 4, 8, 16, 32) 1476 vrows = verifier_experiment(vns, trials=400, seed=0) 1477 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1478 ax.plot(vns, [r["pass_at_n"] for r in vrows], ":", color=MUTED, lw=2, label="pass@n: any sample right (ceiling)") 1479 ax.plot(vns, [r["verifier"] for r in vrows], "o-", color=GREEN, label="best-of-n, step checker picks") 1480 ax.plot(vns, [r["vote"] for r in vrows], "o-", color=BLUE, label="majority vote") 1481 ax.plot(vns, [r["single"] for r in vrows], "o-", color=RED, label="one sample") 1482 ax.set_xscale("log", base=2) 1483 ax.set_xlabel("chains sampled, n (log scale)") 1484 ax.set_ylabel("accuracy") 1485 ax.set_ylim(0, 1.05) 1486 ax.set_title("A good checker turns many tries into a right answer") 1487 ax.legend(frameon=False, fontsize=8, loc="lower right") 1488 figs["verifier"] = fig 1489 1490 # --- 6. RL: learning how long to think ------------------------------------ 1491 runs = ( 1492 ("correctness only", dict(), BLUE), 1493 ("penalty 0.01 per token, ÷ spread (GRPO)", dict(length_penalty=0.01), RED), 1494 ("penalty 0.01 per token, no ÷ spread", dict(length_penalty=0.01, divide_by_spread=False), GREEN), 1495 ) 1496 histories = [(label, train_reasoner(seed=0, **kw), color) for label, kw, color in runs] 1497 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1498 for label, h, color in histories: 1499 a1.plot(h["mean_length"], color=color, label=label) 1500 a2.plot(h["accuracy"], color=color, label=label) 1501 a1.set_xlabel("training iteration") 1502 a1.set_ylabel("mean thinking length (tokens)") 1503 a1.set_title("It learns to write longer") 1504 a2.set_xlabel("training iteration") 1505 a2.set_ylabel("accuracy") 1506 a2.set_ylim(0.35, 1.0) 1507 a2.set_title("...and gets more right") 1508 a2.legend(frameon=False, fontsize=8, loc="lower right") 1509 fig.tight_layout() 1510 figs["rl_training"] = fig 1511 1512 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 1513 x = np.arange(len(DIFFICULTIES)) 1514 width = 0.26 1515 for i, (label, h, color) in enumerate(histories): 1516 ax.bar(x + (i - 1) * width, [h["expected_length"][d] for d in DIFFICULTIES], width, color=color, label=label) 1517 ax.plot(x, DIFFICULTIES, "k_", ms=28, mew=2, label="steps needed d") 1518 ax.plot(x, [2 * d for d in DIFFICULTIES], "_", color="#4b5563", ms=28, mew=1.5, ls="none", label="2d: room for one recheck") 1519 ax.set_xticks(x, [f"d = {d}" for d in DIFFICULTIES]) 1520 ax.set_ylim(0, 18) 1521 ax.set_xlabel("difficulty: additions the problem needs") 1522 ax.set_ylabel("tokens the trained model spends") 1523 ax.set_title("Hard problems get more thinking, with room to recheck") 1524 ax.legend(frameon=False, fontsize=8, loc="upper left") 1525 figs["rl_lengths"] = fig 1526 1527 # --- 7. Compounding slips over long chains -------------------------------- 1528 ks = np.arange(1, 201) 1529 fig, ax = plt.subplots(figsize=(6.4, 3.4)) 1530 for e, color in ((0.005, GREEN), (0.02, BLUE), (0.05, RED)): 1531 ax.plot(ks, [steps_all_right(e, k) for k in ks], color=color, label=f"slip rate per step = {e}") 1532 ax.plot([50], [steps_all_right(0.02, 50)], "o", color=BLUE) 1533 ax.annotate("50 steps at 2%: 0.36", (50, steps_all_right(0.02, 50)), xytext=(70, 0.55), arrowprops=dict(arrowstyle="->", color="#4b5563"), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1534 ax.set_xlabel("steps in the chain, k") 1535 ax.set_ylabel("chance every step is right") 1536 ax.set_ylim(0, 1.03) 1537 ax.set_title("Long chains need checking, not just length") 1538 ax.legend(frameon=False) 1539 figs["compounding"] = fig 1540 1541 return figs
Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.
1549def demo() -> None: 1550 banner("1. Thinking out loud: one addition per forward pass") 1551 say( 1552 """ 1553 Our toy model can do exactly one addition per forward pass, and every 1554 token it writes is one forward pass. Ask it for 3 + 5 + 8 + 2. 1555 """ 1556 ) 1557 direct = solve([3, 5, 8, 2], max_tokens=1) 1558 cot = solve([3, 5, 8, 2], max_tokens=10) 1559 table( 1560 ["how it answers", "what it wrote", "answer", "tokens"], 1561 [ 1562 ("at once", "(nothing)", direct.answer, direct.tokens), 1563 ("thinking first", ", ".join(cot.steps), cot.answer, cot.tokens), 1564 ], 1565 ) 1566 say( 1567 """ 1568 At once, it adds 3 + 5 = 8 and has no pass left for 8 and 2, so it 1569 estimates them at 4.5 each: 17, wrong. Thinking first, it writes each 1570 running total where the next pass can read it: 18, right, for 3 tokens. 1571 """ 1572 ) 1573 takeaway("Each written token is another forward pass: writing steps buys serial computation the weights alone don't have.") 1574 1575 banner("2. Test-time compute: accuracy vs. thinking budget") 1576 rows = budget_sweep() 1577 table( 1578 ["budget B", "formula acc(B)", "simulated accuracy", "tokens actually spent"], 1579 [(r["budget"], budget_accuracy(r["budget"]), r["accuracy"], r["mean_tokens"]) for r in rows], 1580 floatfmt=".3f", 1581 ) 1582 say( 1583 """ 1584 Problems need 1, 2, 4, 8, 16 or 32 additions. Each doubling of the 1585 budget lets one more difficulty level fit, so accuracy climbs by a 1586 steady step per doubling, then stops once the hardest problem fits. 1587 """ 1588 ) 1589 takeaway("With difficulty spread over many scales, accuracy grows with the log of the thinking budget, until it doesn't.") 1590 1591 banner("3. Sample many, vote: self-consistency") 1592 say("Three sampled answers to 3 + 5 + 8 + 2: 18, 17, 18. The vote picks " + str(majority_vote([18, 17, 18])) + ".") 1593 table( 1594 ["samples n", "yes/no, p = 0.6", "yes/no, p = 0.4", "p = 0.6, rho = 0.9 (shared mistakes)"], 1595 [(n, majority_accuracy(0.6, n), majority_accuracy(0.4, n), correlated_vote_accuracy(0.6, n, rho=0.9)) for n in (1, 3, 5, 15, 31)], 1596 floatfmt=".3f", 1597 ) 1598 say( 1599 """ 1600 Independent 60% votes climb towards certainty; independent 40% votes on 1601 a yes/no question sink. When samples share their mistakes (rho = 0.9), 1602 31 votes are barely better than one: they are one opinion, repeated. 1603 """ 1604 ) 1605 takeaway("Voting helps only when the right answer is the most common one and the mistakes are independent.") 1606 1607 banner("4. Verifiers: check the answer, or check every step") 1608 slipped = ["3+5=8", "8+8=17", "17+2=19"] 1609 lucky = ["3+5=9", "9+8=16", "16+2=18"] 1610 table( 1611 ["chain", "outcome reward (answer 18?)", "first bad step (process check)"], 1612 [(", ".join(c), int(outcome_reward(c, 18)), first_bad_step(c)) for c in (slipped, lucky)], 1613 ) 1614 say( 1615 """ 1616 The outcome check only sees the last number: it gives the lucky chain, 1617 whose two slips cancelled, full marks. The step check reads every line 1618 and points at the exact slip. 1619 """ 1620 ) 1621 table( 1622 ["chains n", "one sample", "vote", "best-of-n, step checker", "pass@n ceiling"], 1623 [(r["n"], r["single"], r["vote"], r["verifier"], r["pass_at_n"]) for r in verifier_experiment()], 1624 floatfmt=".3f", 1625 ) 1626 takeaway("A reliable checker turns 'right sometimes' into 'right almost always' with enough samples; voting gets there far more slowly.") 1627 1628 banner("5. Learning to reason with RL: verifiable rewards, group advantages") 1629 say(f"Group rewards [1, 0, 0, 1] give advantages {group_advantages(np.array([1.0, 0.0, 0.0, 1.0])).tolist()}.") 1630 for label, kw in ( 1631 ("correctness only", dict()), 1632 ("penalty 0.01/token, GRPO", dict(length_penalty=0.01)), 1633 ("penalty 0.01/token, no ÷ spread", dict(length_penalty=0.01, divide_by_spread=False)), 1634 ): 1635 h = train_reasoner(seed=0, **kw) 1636 print( 1637 f"{label:32s} length {h['mean_length'][0]:.1f} -> {h['mean_length'][-1]:.1f} tokens, " 1638 f"accuracy {h['accuracy'][0]:.2f} -> {h['accuracy'][-1]:.2f}, " 1639 f"prefers {h['preferred_length']}" 1640 ) 1641 print() 1642 say( 1643 """ 1644 Rewarded only for right answers, the policy learns to think longer 1645 (about 2 tokens to about 7.6) and to leave room for a recheck on hard 1646 problems (16 tokens for 8 additions). A per-token penalty trims easy 1647 problems; dividing by the group's spread magnifies that penalty until 1648 1-addition problems get 1 token and accept an occasional slip. 1649 """ 1650 ) 1651 takeaway("Reward the outcome and let the model search: longer, self-checking reasoning emerges because it pays.") 1652 1653 banner("6. What it costs, and where it fails") 1654 table( 1655 ["slip rate", "steps", "chance all right"], 1656 [(e, k, steps_all_right(e, k)) for e, k in ((0.02, 10), (0.02, 50), (0.02, 200), (0.005, 200))], 1657 floatfmt=".3f", 1658 ) 1659 c1 = reasoning_cost(1, 2000, 10.0, 50) 1660 c8 = reasoning_cost(8, 2000, 10.0, 50) 1661 say( 1662 f""" 1663 One 2,000-token chain at $10 per million tokens costs ${c1['dollars']:.2f} 1664 and takes {c1['seconds']:.0f} s at 50 tokens per second. Eight in parallel 1665 cost ${c8['dollars']:.2f} and still take {c8['seconds']:.0f} s. 1666 """ 1667 ) 1668 takeaway("Thinking is paid per token: spend it where a check shows it helps, and cap it everywhere else.")