primer.ml.structured_output

Structured output: answers that fit a shape, every time

Run: python -m primer.ml.structured_output

New to the notation? primer.notation explains every symbol used here from zero. This lesson builds on sampling from primer.ml.inference and on tool arguments from primer.agents.tools.

Level 1: The practitioner's guide

In one sentence. Structured output means making a model's answer fit an exact shape (a JSON object, a label from a fixed list, a date) every time, so that a program can read it with no person in the loop.

When you need it. The moment something other than a person reads the answer: a tool call's arguments, a row for a database, a classification label, a form to fill in. A person shrugs off a stray word in front of the JSON; a parser rejects the whole answer. You don't need it for prose a person will read, and you don't need it when the reader is another model that copes with loose text. The tell: if you are writing code to strip "Sure, here you go!" from the front of responses, you need it.

How often does asking nicely fail? This lesson's toy model, asked 200 times for {"age": 42} with only the prompt to guide it, gets it exactly right 133 times (66.5%). Real models do far better than a toy, but not perfectly, and the failures grow with length: in an answer of 100 tokens where each token is right 98% of the time, the whole answer comes out valid only about 13% of the time (Level 2 shows why). At a million calls a day, a 1% failure rate is ten thousand broken answers a day.

Your options. Five ways, from the cheapest to the most certain:

Option What it does What it guarantees What it costs Where it lives
Prompting and examples Ask for the format and show an example or two Nothing; it raises the odds A few extra input tokens Your prompt
Post-processing Repair the common slips: strip chatter, close a bracket Nothing, but it catches the frequent cases A small parser you maintain Your code
Validate and retry Parse; on failure, ask again, quoting the error Valid eventually, if the model is usually right A whole extra call per retry, and latency Your code
Constrained decoding Forbid, at every token, anything that cannot lead to a valid answer A valid shape, by construction A grammar compiled once, a small check per token, some drift in what the model says The model server: JSON mode, strict schemas, grammar engines
Fine-tuning on the format Train the model on thousands of examples in the shape Far more reliable; still not certain Data, a training run, a model to host Training

How to choose. Start from what reads the answer and how often it may be wrong.

  • One field, a label from a list, a yes or no: prompt for it and validate. Retries are cheap because the answer is short.
  • A JSON object your code depends on, at volume: use the hosted API's strict schema, or a grammar engine in front of a model you run yourself. It removes the parsing code, the type checks and the retry loop in one move.
  • A format no engine supports (a custom mini-language, a legacy fixed-width record): post-process what you can, validate, retry, and consider fine-tuning once the volume justifies it.
  • Whatever you pick, validate the values afterwards. A schema proves shape, not truth: {"amount": 0, "currency": "USD"} fits a payment schema exactly and is still a bad payment.

What it costs. Prompting costs tokens. Retries cost whole calls and double the slowest requests. Constrained decoding costs a one-time compile of the schema (a noticeable pause on the first request, cached after that) and a small check per token; nested schemas cost more than flat ones. It can also cost quality: forcing a model off the path it wanted can make it invent a value to satisfy a required field, or wander, legally, until the token budget runs out. In this lesson's toy, a strict schema with a required age the model has no answer for produces "age": 9700. Fine-tuning costs the most up front and the least per call.

What breaks.

  • A required field the model can't fill becomes a made-up value that parses. Make such fields optional, or allow null.
  • Valid JSON, wrong content. JSON mode alone guarantees something parseable, not your keys or your types. That needs a schema.
  • Truncation. A valid shape cut off by the output limit is invalid. Set the limit with the schema's size in mind.
  • Token boundaries. A token can straddle a boundary in the grammar, so real engines check character by character; a home-made masker that judges whole tokens rejects valid answers.
  • Drift. A constrained model can sound different, because the mask changes which continuations it is allowed.
  • Business rules. Shape is checked; meaning is not. Keep the validator.

In the wild. Hosted APIs expose the mechanism at three strengths: JSON mode (any valid JSON), strict structured outputs (your schema, guaranteed) and strict tool use (a tool's arguments must fit its input schema). Claude's structured outputs and strict tool use are one example, linked in Further reading. For models you run yourself, Outlines and XGrammar turn a JSON Schema or a regular expression into token masks, and llama.cpp accepts a grammar file (GBNF). Libraries such as instructor (validate and retry against a Python type) and guidance (constrain a generation to a pattern or a set of options) wrap these steps. Every agent framework leans on all of this: a tool call is a structured output.

Go deeper. Level 2 builds constrained decoding from nothing: why failures compound with length (a one-line formula), how a mask is applied before each token is drawn, how a pattern becomes a finite-state machine and a JSON Schema becomes a stack, and what each pitfall above looks like in numbers you can rerun. If you only needed to choose, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

A language model writes one token at a time, and each token is a draw from a probability distribution (see primer.ml.inference). Most of the time you want free text. Sometimes a program is going to read the answer: a tool call's arguments, a row for a database, a label from a fixed list. Then the answer must have an exact shape, such as a JSON object with an integer field called age, and one stray character breaks it.

Structured output is the family of techniques that make the shape certain. The main one, constrained decoding, is surprisingly small: before each token is drawn, find every token that could not possibly continue a valid answer, and forbid it. This lesson builds that from scratch: first for simple patterns, with a finite-state machine, then for nested JSON, with a stack.

1. Asking nicely is not enough

Everyday picture. You read a form aloud over the phone to a friend and ask them to type it in exactly. They are careful and mostly right. But now and then they add "Sure, here you go!" before the form, or they forget the last bracket, or a finger slips. A person reading the result shrugs it off. A program reading it stops dead: one wrong character and the whole thing is rejected.

Tiny worked example. This lesson's toy model, ToyModel, has mostly learned to answer {"age": 42}. It is a pretend model, a few lines of NumPy, that copies that answer with some noise on every choice and a habit of opening with "Sure". Asked 200 times, with nothing but the request to guide it, it writes the exact answer 133 times (66.5%). The other 67 look like this:

What it wrote What went wrong
Sure{"age": 42} a friendly word before the JSON (22 times)
{"age": 42w a slip where the closing brace should be
{"age": 428 an extra digit, then it stopped
{"age"6 42} a digit where the colon belongs
Sure{"age": 42}EUR chatter at both ends

Nothing here is a big mistake. Each one is a single bad token.

flowchart LR P["Prompt: reply with JSON"] --> M[Model draws one token] M --> T[Append it to the answer] T -->|not finished| M T -->|finished| J{Parser} J -->|every token right| OK[Valid JSON] J -->|one token wrong| ERR[Parse error:<br/>the whole answer is lost]

Reading it: follow the loop on the left: the model adds one token at a time and nothing checks the answer until it is finished. Only then does the parser look at it, and it has two exits. One bad token anywhere in the loop sends the whole answer to the error exit, however good the rest was.

That is why failures pile up with length. If each token is right with chance $p$, and the chances are roughly independent, the whole answer is right only when every token is:

Level 3: the formula and its symbols

$$ P(\text{valid}) = p^{\,n} $$

Symbols

Symbol Meaning here In the example
$p$ the chance that one token is right, between 0 and 1 0.98
$n$ how many tokens the answer has 10
$p^{\,n}$ $p$ multiplied by itself $n$ times: the chance that all $n$ are right $0.98^{10}$
$P(\text{valid})$ the chance that the whole answer parses 0.817

In words: "the chance that a whole answer is valid is the chance that one token is right, multiplied together once for every token."

With the numbers: a model that gets 98% of tokens right, writing a 10-token answer, produces valid output $0.98^{10} = 0.817$ of the time: nearly one answer in five is broken. Make the answer 100 tokens long and it drops to $0.98^{100} = 0.133$.

Level 3: in Python

In Python:

# p: chance one token is right; n: tokens in the answer
p, n = 0.98, 10
# p^n: all n tokens right
round(p ** n, 3)  # → 0.817
# a ten times longer answer
round(p ** 100, 3)  # → 0.133

Valid-output rate falls as the answer grows: at 98% per token, 10 tokens are valid 82% of the time and 100 tokens only 13%

Reading it: the x-axis is the length of the answer in tokens, the y-axis the share of answers that come out valid. Each curve is one per-token accuracy. Even the top curve, 99.9% per token, sags over a long answer, and the 95% curve is near zero by 100 tokens. The dot marks the worked example. Better prompting moves you to a higher curve, but no curve stays at 100%.

In code: ToyModel is the pretend model, sample_answer draws one answer token by token, and chance_all_valid is the formula.

Why it matters in practice. A program that calls a model a million times a day and fails 1% of the time fails ten thousand times a day. Long answers, nested objects and small models make it worse. Prompting and examples help, but they only move you to a better curve. To get to 100% you have to change how tokens are chosen.

2. Constrained decoding: mask, then sample

Everyday picture. Picture a keyboard whose keys lock and unlock as you type. You still decide what to write. But at every keystroke, any key that would make the text break the form is locked for that one keystroke. You cannot make the mistake, because the key isn't there.

Tiny worked example. At the very first step the model scores four tokens (these scores are logits: raw preferences, before softmax turns them into probabilities; see primer.ml.attention for softmax from zero). Only { and [ can start a JSON value, so the other two are locked:

Token Logit $z$ Allowed? $m$ Share before the mask Share after the mask
Sure 2.0 0 0.579 0
Here 1.0 0 0.213 0
{ 0.5 1 0.129 0.622
[ 0.0 1 0.078 0.378

Before the mask the model put 79% of its belief on chatter. After the mask, the chatter has a chance of exactly zero, and the two allowed tokens share everything. They keep their odds against each other: { was 1.65 times as likely as [ before, and it still is.

flowchart LR A[Answer so far] --> M[Model: a logit<br/>for every token] A --> C[Constraint: which tokens<br/>can still lead to a valid answer?] M --> K[Set forbidden logits to minus infinity] C --> K K --> S[Softmax over what is left<br/>then sample] S --> N{End token?} N -->|no| A2[Append the token] --> A N -->|yes| D[Done: valid by construction]

Reading it: two arrows leave the answer so far. The top path is the ordinary model, which scores every token as it always does. The bottom path is the new part: a checker that knows the shape and says which tokens are still possible. They meet in the mask box, and from there it is ordinary sampling again. The end token is itself just a token, so it is allowed only when the answer is complete.

The mask as a formula is softmax with a 0-or-1 switch on every term:

Level 3: the formula and its symbols

$$ p_i = \frac{m_i \, e^{z_i}}{\sum_{j=1}^{|V|} m_j \, e^{z_j}} $$

Symbols

Symbol Meaning here In the example
$V$ the vocabulary: every token the model can write the 4 tokens in the table
$\lvert V \rvert$ how many tokens that is 4
$i$ the token whose probability we are computing 3, the {
$z_i$ token $i$'s logit $z_3 = 0.5$
$m_i$ the mask: 1 if token $i$ can continue a valid answer, 0 if not $m = (0, 0, 1, 1)$
$e^{z_i}$ $e \approx 2.718$ raised to the logit: always positive $e^{0.5} = 1.649$
$\sum_{j=1}^{\lvert V \rvert}$ add up the following for every token $j$ $0 + 0 + 1.649 + 1$
$p_i$ the probability that token $i$ is drawn 0.622

In words: "a token's probability is its usual softmax share, except that forbidden tokens count as zero on top and bottom, so the allowed tokens share all the probability between them."

With the numbers: $p_{{} = \dfrac{1 \cdot e^{0.5}}{0 + 0 + 1 \cdot e^{0.5} + 1 \cdot e^{0}} = \dfrac{1.649}{2.649} = 0.622$.

Level 3: in Python

In Python:

import math
# Sure, " Here", "{", "["
z = [2.0, 1.0, 0.5, 0.0]
# m_i: 1 if the token may come next
m = [0, 0, 1, 1]
# m_i e^(z_i): forbidden tokens contribute nothing
kept = [m_i * math.exp(z_i) for m_i, z_i in zip(m, z)]
[round(k, 3) for k in kept]  # → [0.0, 0.0, 1.649, 1.0]
# Σ_j m_j e^(z_j)
total = sum(kept)
round(total, 3)  # → 2.649
[round(k / total, 3) for k in kept]  # → [0.0, 0.0, 0.622, 0.378]
# without the mask, most of the belief went to chatter
e = [math.exp(z_i) for z_i in z]
[round(x / sum(e), 3) for x in e]  # → [0.579, 0.213, 0.129, 0.078]

Setting a logit to minus infinity does the same thing, because $e^{-\infty} = 0$: it is the trick the causal mask uses in primer.ml.attention. After masking, temperature, top-k and top-p from primer.ml.inference work exactly as before, on the tokens that are left.

At the toy model's first step, 19% of its belief sits on Sure; after the mask, the brace gets all of it

Reading it: these are the toy model's real first-step probabilities for the {"age": 42} task, top four tokens only. Grey bars are what the model wanted; blue bars are what it may choose from after the mask. "Sure" had nearly a fifth of the belief and drops to zero. The brace, the only legal way to begin, takes everything.

Why 100% by construction. The checker keeps one promise: every answer so far can still be finished validly. It holds at the start (the empty answer can be finished). Each step only allows a token that keeps it. The end token is allowed only when the answer is already complete. So every answer that ends is valid. No luck is involved.

Unconstrained, the age answer is valid 66.5% of the time and the person object 18%; constrained, every finished answer is valid

Reading it: each bar is 200 answers from the toy model, split by what happened. Green is valid. The unconstrained bars show all three failures: chatter before the JSON, a broken character inside, and (rarely) running out of token budget. The person object is longer, so it breaks far more often, just as $p^n$ predicts. The constrained bars have no chatter and nothing broken. The one sliver left, 5 of the 200 person answers, is cut off: the 40-token budget ran out mid-object. The guarantee covers every prefix, but only finishing makes a whole answer, so leave room in the budget.

In code: masked_softmax applies the mask, sample_answer takes an optional constraint and masks every step, and validity_experiment produces the bars above.

Why it matters in practice. Constrained decoding changes nothing about the model: no retraining, same weights. It only changes which tokens may be drawn, which is why it can be added to any open model at serving time. The hard part is the checker: answering "which of 100,000 tokens could still lead to a valid answer?" fast, at every step. The next two sections build it.

3. From a pattern to a state machine

Everyday picture. A subway map. You stand at a station. Each line leaving it is labelled with one character. To write a character, you ride the line with that label; if no line from your station has it, that character is impossible here. Some stations are marked "you may stop here". That map is a finite-state machine: a fixed set of states (stations) and one move per character. A regular expression, the pattern language behind [0-9]+ and cat|car|dog, can always be drawn as one.

Tiny worked example. The pattern cat|car|dog (one of three words) becomes this machine:

stateDiagram-v2 [*] --> S0 S0 --> S1: c S0 --> S2: d S1 --> S3: a S2 --> S4: o S3 --> S5: r S3 --> S6: t S4 --> S7: g S5 --> [*] S6 --> [*] S7 --> [*]

Reading it: start at S0 and read a word one character at a time. "cat" goes S0, S1, S3, S6, and S6 has an exit arrow, so "cat" is accepted. "cow" gets stuck at S1, which has no "o" line. After "ca" (state S3) only "r" and "t" are possible: the machine has turned "what may come next?" into "which lines leave this station?".

Tokens are not characters. A model writes tokens, and one token can hold several characters. The rule: a token is allowed if the machine can swallow all of its characters, one after another, without getting stuck. With a vocabulary of 13 tokens:

Token From S0 From S3 (after "ca")
c, d allowed stuck
ca, cat, do, dog allowed: every character has a line stuck
a, at, o, og, g stuck at the first character stuck
t, r stuck allowed

cat is allowed at the start even though no single line reads "cat": the machine rides c, then a, then t. at is part of a real word and still never allowed at the start, because S0 has no "a" line.

Level 3: the formula and its symbols

$$ \delta^(s, t) = \delta\big(\cdots\delta(\delta(s, c_1), c_2)\cdots, c_k\big) \qquad A(s) = {\, t \in V : \delta^(s, t) \text{ is defined} \,} $$

Symbols

Symbol Meaning here In the example
$s$ a state of the machine S0
$\delta(s, c)$ the transition function: the state you reach from $s$ by reading character $c$, or undefined if there is no such line $\delta(0, \text{c}) = 1$
$t$ a token cat
$c_1, \ldots, c_k$ the characters of $t$, in order; $k$ is how many c, a, t; $k = 3$
$\delta^*(s, t)$ read every character of $t$ in turn, starting from $s$ $\delta^*(0, \texttt{cat}) = 6$
$V$ the vocabulary the 13 tokens above
${\, t \in V : \ldots \,}$ "the set of every token $t$ in $V$ for which ... holds"
$A(s)$ the allowed tokens at state $s$ $A(0) = {$c, d, ca, cat, do, dog$}$

The end token joins $A(s)$ only when $s$ is an accepting state (S5, S6 or S7 here), because only there is the text complete.

In words: "to see whether a token fits, walk its characters through the machine one at a time; the allowed tokens are the ones that never get stuck."

With the numbers: $\delta^*(0, \texttt{cat}) = \delta(\delta(\delta(0, \text{c}), \text{a}), \text{t}) = \delta(\delta(1, \text{a}), \text{t}) = \delta(3, \text{t}) = 6$, so cat is in $A(0)$. $\delta(0, \text{a})$ is undefined, so at is not.

Level 3: in Python

In Python:

# the cat|car|dog machine: state -> {character: next state}
delta = {0: {"c": 1, "d": 2}, 1: {"a": 3}, 2: {"o": 4}, 3: {"r": 5, "t": 6}, 4: {"g": 7}, 5: {}, 6: {}, 7: {}}
accepting = {5, 6, 7}
def walk(s, token):
    # δ*: one δ per character; None means stuck
    for ch in token:
        s = delta[s].get(ch) if s is not None else None
    return s
walk(0, "cat")  # → 6
walk(0, "at")  # → None
V = ["c", "a", "t", "r", "d", "o", "g", "ca", "cat", "do", "dog", "at", "og"]
# A(0): every token the machine can swallow whole from the start
[t for t in V if walk(0, t) is not None]  # → ['c', 'd', 'ca', 'cat', 'do', 'dog']
# A(3), after "ca"
[t for t in V if walk(3, t) is not None]  # → ['t', 'r']
# the end token only where the text is complete
walk(0, "cat") in accepting  # → True

How a pattern becomes a machine. Two classic steps, both in the code:

flowchart LR P["Pattern<br/>cat|car|dog"] --> T[Thompson's construction:<br/>one small machine per piece,<br/>glued with free jumps] T --> N[NFA: may be in<br/>several states at once] N --> D[Subset construction:<br/>each new state is a set<br/>of NFA states] D --> F[DFA: exactly one<br/>state at a time] F --> TB[Table: for every state,<br/>which tokens are allowed]

Reading it: left to right, the pattern gets more mechanical. Thompson's construction reads the pattern like a sentence and builds a tiny machine for each piece: one line for a character, a fork for |, a loop for +. Glued together they make an NFA (a non-deterministic machine), which can be in several states at once: after "c" it is both "inside cat" and "inside car". The subset construction turns each set of NFA states into one state of a DFA (a deterministic machine), which is always in exactly one state, so following it is a dictionary lookup. The last box is the payoff: since the DFA has a fixed number of states, the allowed tokens can be worked out for every state before generation starts.

The pattern -?[0-9]+ (an optional minus sign, then digits) compiles to just three states: the start, "saw a minus", and "saw at least one digit", the only accepting one. The lesson's running example, \{"age": [0-9]+\}, compiles to 11.

Each row is one state of the age machine, each column a token; only a handful of cells are lit, and the digit states allow many tokens at once

Reading it: rows are the 11 states, labelled by the text that reaches them; columns are tokens; a dark cell means "allowed here". Most rows have one or two dark cells: the structure is fixed, so only one next character is legal, written alone or as the start of a longer token such as "age" or ":. The two digit rows are where the model has real choice: any digit, a two-digit token such as 42, or, once one digit is down, the closing brace. The last row allows only the end token. This whole table is computed once, before the first token.

In code: compile_pattern runs both constructions and returns a DFA; DFA.walk follows text through it; allowed_tokens is $A(s)$, and mask_table precomputes it for every state. RegexConstraint plugs the table into sample_answer.

Why it matters in practice. Integers, dates, enums, phone numbers and fixed-key objects are all patterns. Reframing generation as moving between the states of a machine, with the allowed tokens indexed per state, is the idea of Willard and Louf (2023) behind the open-source Outlines library, and it makes each step's mask a single lookup.

4. JSON Schema: nesting needs a stack

Everyday picture. A stack of plates. Each time you open something, a bracket, a brace or a quote, you put a plate on the stack with a note: "an array is open", "an object is open". To close something, you may only take the top plate: close the most recent thing first. When the stack is empty, you're done.

A subway map can't do this job. JSON can nest as deep as you like: [[[[...]]]]. A machine with, say, 50 states can't tell 50 open brackets from 51, so it can't know how many closers it still owes. Counting without limit needs memory without limit. A finite-state machine plus a stack is called a pushdown automaton, and it is exactly enough for nested formats such as JSON, SQL and most programming languages.

Tiny worked example. Read {"a": [1, {"b": one character at a time.

flowchart LR S1["after {<br/>stack: object"] --> S2["after [<br/>stack: object, array"] S2 --> S3["after the inner {<br/>stack: object, array, object"] S3 --> S4["after 2}<br/>stack: object, array"] S4 --> S5["after ]}<br/>stack: empty, done"]

Reading it: each box is a moment in the reading, with the stack written bottom first. Every opener adds a frame on the right; every closer removes the rightmost one. At the third box the innermost open thing is an object, so } is the only closer allowed, and ] would be refused even though an array is open further down. The last box is empty: the value is complete, and only now is the end token allowed.

The checker answers one question about any prefix: can it still be finished?

Prefix Status Why
{"a": [1, {"b": open a value for "b" may come next
{"a": [1} dead the array is on top, so } can't close it
{"a": [1, {"b": 2}]} complete the stack is empty

The schema steers every character. A JSON Schema (see primer.agents.tools) names the fields and their types. This lesson's person schema allows name (a string), age (an integer) and pets (an array of "cat" or "dog"); name and age are required and nothing else is allowed. The checker enforces each rule at the first character that breaks it:

Prefix Status The rule that decides
{"na open "na" can still become "name"
{"nx dead no field starts with "nx"
{"name": "Ada"} dead age is required, so the object may not close yet
{"age": 3. dead an integer has no decimal point
{"age": 01 dead JSON numbers never start with a zero followed by digits
{"pets": ["x dead only "cat" and "dog" are listed
{"name": "Ada", "age": 36} complete every required field is present

Inside one frame, a small state machine does the work. Here is an object's:

stateDiagram-v2 [*] --> open: { open --> key: quote, if a field may be added key --> key: a letter that still spells an unused field key --> colon: closing quote, if the key is complete colon --> value: colon, then at most one space value --> after_value: the value's own frame is pushed and later popped after_value --> key: comma, space, quote after_value --> [*]: }, only if every required field is present open --> [*]: }, only if nothing is required

Reading it: these are the phases one object frame moves through, from opening brace to closing brace. The labels carry the schema's rules: a key letter is allowed only if it still spells a field not yet used, and the closing brace only if every required field is in. The value arrow is where nesting happens: the value gets its own frame pushed on top of the stack, and this frame waits in "value" until that frame is popped.

The leaves of JSON (numbers, true, false, null) are regular patterns, so the checker reads numbers with the machines from section 3. Real engines split the work the same way: patterns for the small pieces, a stack for the nesting.

One rule here is deliberate: at most one space after : and ,. JSON allows any amount of whitespace, and a constrained model that has lost its way can pad with spaces until the budget runs out. Real engines limit whitespace for the same reason.

In code: SchemaChecker.step reads one character and returns the new stack (a tuple of Frame entries), or nothing when the prefix is dead; SchemaChecker.status answers open, dead or complete for a prefix; SchemaChecker.open_containers shows the stack; SchemaConstraint tries every token against the stack at each step.

Why it matters in practice. This is how a schema becomes a guarantee: compile the schema into grammar rules, run them as a pushdown automaton, and mask every token that would kill it. Geng et al. (2023) showed the approach works for structured tasks without any fine-tuning, and PICARD (Scholak et al., 2021) did the same for SQL by parsing incrementally. The stack has a cost: it can grow without limit, so the masks can't all be tabled in advance as they were for a pattern. Section 5 comes back to that.

5. Costs and pitfalls

Constrained decoding guarantees the shape. It does not guarantee a good answer, and it isn't free.

The mask changes what the model says

Everyday picture. A satellite navigation system that only forbids illegal turns, one junction at a time, and never looks ahead. It happily takes the motorway because the motorway looks fastest right now, and only at the end discovers that the one legal exit is a long detour. A driver who could see the whole map would have taken the side road from the start.

Tiny worked example 1: forced to invent. The model is asked for a person's age, but the text never says it. The model's honest belief for the next token:

Token Model's belief Allowed by "type": "integer"? After the mask
null 0.80 no 0
3 0.12 yes 0.12 / 0.20 = 0.60
5 0.08 yes 0.08 / 0.20 = 0.40

The mask throws away 80% of the model's belief, and the answer is a made-up age, delivered as confidently as a real one. The toy model does this too: told to write the person schema while it "wants" to write only a name, it is forced to add an age, and invents numbers such as 9700. The fix is in the schema, not the decoder: allow null ("type": ["integer", "null"]), or add a field for "not stated". A related finding (Tam et al., 2024): forcing a strict format from the first token can hurt a model's reasoning, so let the model reason in free text first, or put a reasoning field before the answer field.

Tiny worked example 2: locally tempting, globally wrong. A two-token model. Valid answers must end in "y".

flowchart LR S[start] -->|A 0.9| A[A] S -->|B 0.1| B[B] A -->|x 0.99| AX["Ax: invalid"] A -->|y 0.01| AY["Ay: valid, 0.9 × 0.01 = 0.009"] B -->|x 0.1| BX["Bx: invalid"] B -->|y 0.9| BY["By: valid, 0.1 × 0.9 = 0.09"]

Reading it: each arrow is one token with the model's probability for it, and each leaf is a whole answer with the product of the probabilities along its path. Of the two valid answers, the model much prefers By (0.09 against 0.009). But masking decides one token at a time. At the first step both A and B can still end in "y", so nothing is masked, and A wins 90% of the time. At the second step the mask forces "y". The result: Ay, the answer the model itself thought ten times less likely, comes out 90% of the time.

Level 3: the formula and its symbols

$$ P(y \mid \text{valid}) = \frac{P(y)}{\sum_{y' \in \mathcal{L}} P(y')} \qquad P_{\text{mask}}(y) = \prod_{i=1}^{n} p'(y_i \mid y_{

Symbols

Symbol Meaning here In the example
$y$ one whole answer Ay
$y_i$ its $i$-th token; $y_{ $y_2 =$ y, $y_{<2} =$ A
$n$ the number of tokens in the answer 2
$P(y)$ the model's own chance of writing $y$: its token chances multiplied together $0.9 \times 0.01 = 0.009$
$\mathcal{L}$ the set of valid answers (the "language" the grammar allows) {Ay, By}
$\sum_{y' \in \mathcal{L}}$ add up over every valid answer $y'$ $0.009 + 0.09$
$P(y \mid \text{valid})$ "the chance of $y$ given that the answer is valid": the model's own odds, among valid answers only 0.091
$p'(y_i \mid y_{ the masked, renormalised chance of token $y_i$ at that step $p'(\text{y} \mid \text{A}) = 1$
$\prod_{i=1}^{n}$ multiply together over every step $0.9 \times 1$
$P_{\text{mask}}(y)$ the chance that token-by-token masking produces $y$ 0.9

In words: "what the model believes, restricted to valid answers, is each valid answer's probability divided by the total of all valid ones; what masking actually produces is the product of the step-by-step masked probabilities, and the two need not agree."

With the numbers: $P(\text{Ay} \mid \text{valid}) = 0.009 / 0.099 = 0.091$, but $P_{\text{mask}}(\text{Ay}) = 0.9 \times 1 = 0.9$: ten times the model's own preference.

Level 3: in Python

In Python:

first = {"A": 0.9, "B": 0.1}
second = {"A": {"x": 0.99, "y": 0.01}, "B": {"x": 0.1, "y": 0.9}}
# P(y) for each valid answer: the model's token chances multiplied
P = {a + "y": first[a] * second[a]["y"] for a in first}
{k: round(v, 3) for k, v in P.items()}  # → {'Ay': 0.009, 'By': 0.09}
# P(y | valid): divide by the total over valid answers
total = sum(P.values())
{k: round(v / total, 3) for k, v in P.items()}  # → {'Ay': 0.091, 'By': 0.909}
# P_mask: step 1 keeps both, step 2 renormalises "y" to 1
{a + "y": first[a] * (second[a]["y"] / second[a]["y"]) for a in first}  # → {'Ay': 0.9, 'By': 0.1}

The model's own odds among valid answers favour By 91 to 9; token-by-token masking produces Ay 90% of the time

Reading it: two answers, two bars each. Grey is the model's own preference among valid answers; blue is what masked decoding actually produces. The bars point in opposite directions. Nothing invalid comes out, yet the distribution is badly bent. Park et al. (2024) name this problem and propose a correction; in practice it is milder when the model already writes the format well, because then the mask rarely has to overrule it.

The toy model shows the everyday version: once noise knocks it off its answer, the mask keeps it legal but not sensible, and it writes digits until the closing brace happens to win, as in {"age": 4174209}.

In code: forced_choice renormalises a belief over the allowed tokens and reports the share thrown away; NULLABLE_AGE_SCHEMA is the fix; distortion_example computes both distributions above.

Token boundaries

Everyday picture. You can say "forty-two", or spell it "four, two". Both arrive at the same text, but only one is how you would naturally say it.

Tiny worked example. The toy vocabulary has a 42 token and single digits, so "42" can be written two ways: 42, or 4 then 2. The whole answer {"age": 42} can be spelled 16 ways, and the person answer 1,536 ways. The mask allows every one of them, but a trained model has almost only ever seen the first, its tokenizer's usual spelling (see primer.ml.tokenization). When the mask forces it onto an unusual spelling, it is in unfamiliar territory and its next choices get worse.

flowchart LR S["after the space"] -->|42, the usual token| E[before the brace] S -->|4| M[after the 4] M -->|2| E

Reading it: two paths lead from the same place to the same place and write the same text. The top path is the one the model saw thousands of times in training; the bottom one it rarely saw. The checker walks characters, so it cannot tell them apart; only the model can.

Two more boundary effects follow from the same fact. A single token can cross several structural boundaries at once, such as "} closing a string and an object together; walking the token character by character, as allowed_tokens does, handles that. And if the prompt ends in the middle of what would normally be one token (a prompt ending in {"age": when the model would usually write ": as one token), the natural token is no longer available. Some engines back up one token and let the model rewrite it, which is called token healing.

In code: tokenizations lists every way a vocabulary can spell a text.

The speed of computing masks

Everyday picture. Before every keystroke, a proofreader checks every word in the dictionary against the rules: slow. Or: a card for every station, printed once, listing which words fit there: fast, once the cards exist.

Tiny worked example. A real vocabulary has around 128,000 tokens. Say they average 4 characters, the answer is 200 tokens long, and the pattern's machine has 50 states. Checking every token at every step walks 200 × 128,000 × 4 = 102.4 million characters for one answer. Building the table walks 50 × 128,000 × 4 = 25.6 million characters once; after that, each step only reads one row of 128,000 yes-or-no entries.

Level 3: the formula and its symbols

$$ W_{\text{naive}} = n \cdot \lvert V \rvert \cdot \bar{L} \qquad W_{\text{table}} = S \cdot \lvert V \rvert \cdot \bar{L} + n \cdot \lvert V \rvert $$

Symbols

Symbol Meaning here In the example
$W$ work: characters walked or table entries read
$n$ tokens generated 200
$\lvert V \rvert$ vocabulary size 128,000
$\bar{L}$ the average token length in characters (the bar means "average") 4
$S$ states in the machine 50
$\cdot$ multiply

In words: "the naive way walks every token at every step; the table walks every token once per state, up front, then reads one row per step."

With the numbers: $W_{\text{naive}} = 200 \cdot 128{,}000 \cdot 4 = 102{,}400{,}000$ for every answer. $W_{\text{table}} = 50 \cdot 128{,}000 \cdot 4 + 200 \cdot 128{,}000 = 51{,}200{,}000$ for the first answer, and only the second term, 25.6 million cheap reads, for each answer after that.

Level 3: in Python

In Python:

n, V, L_bar, S = 200, 128_000, 4, 50
# W_naive = n · |V| · L̄
n * V * L_bar  # → 102400000
# W_table = S · |V| · L̄ + n · |V|
S * V * L_bar + n * V  # → 51200000
# ten answers: the table's build cost is paid once
10 * n * V * L_bar, S * V * L_bar + 10 * n * V  # → (1024000000, 281600000)

Checking every token at every step costs the same for every answer; the table pays once, then grows slowly

Reading it: the x-axis counts answers generated with one schema, the y-axis the total work so far. The naive line climbs steeply and steadily. The table line starts above zero (the build) and then climbs gently (one row read per step). They cross partway through the very first answer. That is why hosted APIs compile a schema once and cache it: the first request with a new schema is slower, and the ones after are fast.

For a schema with nesting, the stack can grow without limit, so no table can cover every situation. XGrammar (Dong et al., 2024) splits the vocabulary: most tokens are context-independent (whether they fit depends only on the current grammar position, not on what is deeper in the stack), and those are prechecked into tables; only the few context-dependent ones are checked against the stack at run time.

In code: mask_cost is the formula; mask_table builds the table once and RegexConstraint reads a row per step, while SchemaConstraint walks every token against the stack at every step, the naive way.

When validating and retrying is enough

Everyday picture. Instead of a form that can't be filled in wrong, you let people fill in a blank sheet, check it, and hand it back with a note when it is wrong. Fine if most people get it right first time; miserable if most don't.

Tiny worked example. The toy model writes a valid age answer 66.5% of the time. Retrying until it succeeds takes 1 / 0.665 = 1.50 attempts on average, and three tries in a row all fail only 0.335³ = 3.8% of the time. For the longer person object, valid 18% of the time, it takes 5.56 attempts on average: slow and expensive.

flowchart LR A[Ask the model] --> V{Validate} V -->|valid| D[Use it] V -->|invalid| E[Send back the<br/>validator's message] --> A E -.->|too many tries| F[Give up or<br/>fall back]

Reading it: the happy path goes straight through. Every failure costs a full extra model call, round the loop. Sending back the validator's exact message (as primer.agents.tools.validate produces) makes the second try much more likely to succeed. The dotted exit is the budget: a loop needs a limit.

Level 3: the formula and its symbols

$$ \mathbb{E}[\text{attempts}] = \frac{1}{p} \qquad P(\text{all } k \text{ tries fail}) = (1 - p)^k $$

Symbols

Symbol Meaning here In the example
$p$ the chance one try is valid 0.665
$\mathbb{E}[\ldots]$ the expected value: the long-run average over many repeats
$\frac{1}{p}$ the average number of tries to the first success 1 / 0.665 = 1.50
$k$ a number of tries 3
$1 - p$ the chance one try fails 0.335
$(1 - p)^k$ the chance that $k$ independent tries all fail $0.335^3 = 0.0376$

In words: "if each try succeeds with chance p, you need one over p tries on average, and the chance that k tries all fail is the chance of one failure, multiplied together k times."

With the numbers: $1 / 0.665 = 1.50$; $(1 - 0.665)^3 = 0.0376$; for the person object, $1 / 0.18 = 5.56$.

Level 3: in Python

In Python:

# the toy model's unconstrained rate on the age answer
p = 0.665
# E[attempts] = 1 / p
round(1 / p, 2)  # → 1.5
# (1 - p)^k: three tries, all invalid
round((1 - p) ** 3, 4)  # → 0.0376
# the longer person answer, valid 18% of the time
round(1 / 0.18, 2)  # → 5.56

Expected attempts are close to 1 for a model that is usually right, and shoot up as the valid rate falls

Reading it: the x-axis is the chance a single try is valid; the y-axis is the average number of tries until one is. The curve is flat on the right and steep on the left. The dots mark three cases: a model that is 95% right barely notices retries (1.05 tries); the toy age answer needs half a try extra; the toy person answer needs more than five calls.

Validate and retry when: you can't change the decoder (a hosted model without a structured-output option), the model is already right nearly every time, or the rule can't be expressed as a grammar. Constrain when answers are long or nested, the model is small, or latency matters. Either way, keep the validator: a grammar enforces shape, not rules such as minimum, and never truth.

In code: expected_attempts and chance_still_failing are the two formulas.

6. JSON mode, strict schemas and tool calls

Everyday picture. Two paper forms. One only insists that you write in block capitals: whatever you write is readable, but nothing says what goes where. The other has a labelled box for each field, and tick boxes where only certain answers are allowed. The first is JSON mode; the second is a strict schema.

Tiny worked example. Give the toy model a reference answer that leaves out the age, {"name": "Ada"}, and ask 50 times. With JSON mode (the empty schema {}, which allows any JSON value) every finished answer parses: 44 copies of {"name": "Ada"} with the required age missing, one {"name": 42628} whose name is a number, and one bare false. All valid JSON; not one valid person. (The other 3 ran out of budget.) With the person schema, every one of the 37 answers that finish has an age, because the closing brace stays locked until one is written. Since the model had no age to give, it invents one, such as {"name": "Ada","age": 9700,"pets": []}: section 5's warning in action. The other 13 show a second cost of forcing a model off its path: it loses its way and wanders, legally, until the token budget runs out.

A tool call is the same thing with a name attached: the model's arguments are a JSON object that must fit the tool's input_schema (see primer.agents.tools and primer.agents.llm).

sequenceDiagram participant App as Your code participant API as Model server participant G as Grammar engine App->>API: messages + tool with a strict JSON Schema API->>G: compile the schema (cached for next time) loop every token G-->>API: mask of allowed tokens API->>API: mask, softmax, sample end API-->>App: tool_use with arguments that fit the schema App->>App: business rules: does customer C-999 exist? is the amount above the minimum?

Reading it: the grammar engine lives inside the model server, next to sampling. It compiles the schema once (the slow first request from section 5) and then hands a mask to every step of the loop. What reaches your code is guaranteed to have the right shape. The last arrow is the part no grammar covers: whether the values are true and allowed. The payment tool in primer.agents.tools makes this concrete: {"amount": 0, "currency": "USD"} fits the schema's shape exactly, and primer.agents.tools.validate still rejects it, because "minimum": 0.01 is a rule about the value, not the shape.

In code: SchemaConstraint with the schema {} is JSON mode, and with PERSON_SCHEMA it is a strict schema; primer.agents.tools.ToolRegistry.definitions shows how a strict tool definition is sent.

Why it matters in practice. You now have three tools and know what each buys. JSON mode guarantees something parseable. A strict schema guarantees the shape your code expects, so the parsing and type-checking code disappears. Validation after the fact still catches what only your program knows. Hosted APIs (for example Claude's structured outputs and strict tool use) and open engines (Outlines, llama.cpp grammars, XGrammar) all run the mechanism built in this lesson: a checker that masks the logits before every draw.

In 20 seconds

  • The problem: every token is a chance to break the format, so an answer of n tokens is valid about $p^n$ of the time. Prompting raises p but never to 1.
  • Constrained decoding: before each draw, set the logit of every token that can't lead to a valid answer to minus infinity, then sample as usual. Allow the end token only when the answer is complete. Valid by construction, if the answer is allowed to finish.
  • Patterns: compile to a finite-state machine; a token is allowed if the machine can read all its characters; precompute the allowed tokens for every state, so each step is a lookup.
  • JSON Schema: nesting needs a stack (a pushdown automaton); the schema decides which keys, types and closers are legal at each character.
  • Pitfalls: the mask can force a made-up value or bend the model's preferences; tokens don't align with grammar pieces; nested grammars cost more to check. Retrying is fine when the model is usually right. Always validate business rules afterwards.

Self-test questions

Why does asking a model for JSON in the prompt fail some of the time, and why do longer outputs fail more? Each token is a separate draw with some small chance of being wrong, and a single wrong token breaks the parse. The chance that all n tokens are right is about $p^n$, which shrinks as n grows: 98% per token gives 82% at 10 tokens and 13% at 100.

What exactly does constrained decoding change in the model? Nothing in the weights. At each step it sets the logits of forbidden tokens to minus infinity (probability zero) before sampling. The allowed tokens keep their relative odds, and temperature and top-p still apply to them.

Why is the output valid "by construction", and what can still go wrong with the shape? Every prefix is kept completable, and the end token is allowed only when the answer is complete, so any answer that ends is valid. It can still be cut off by the token budget, leaving a valid but unfinished prefix.

How do you decide whether a multi-character token is allowed in a given state? Walk its characters through the state machine one at a time from the current state. It is allowed if the walk never gets stuck. cat is allowed at the start of cat|car|dog; at is not, because the start has no "a" line.

Why can't a finite-state machine check arbitrary JSON? JSON nests without limit, and closing correctly requires remembering every open bracket in order. A machine with a fixed number of states can't count without limit. A stack, which a pushdown automaton adds, can.

Why can masks be precomputed for a pattern but not fully for a JSON Schema? A pattern's machine has a fixed number of states, so the allowed tokens for each can be tabled once. With nesting the stack can take unboundedly many forms. Engines precompute the tokens whose fate depends only on the current position and check the rest against the stack at run time.

How can constrained decoding make answers worse? It forces the model's choices into the allowed set even when the model believed something else: an integer field makes it invent a number when the honest answer was null, and token-by-token masking can commit early to a path the model thought unlikely overall. Allowing null, or letting the model reason before the structured part, helps.

When is validating and retrying good enough? When the model is valid almost every time (95% needs about 1.05 calls on average), when you can't change the decoder, or when the rule can't be written as a grammar. It gets expensive fast as the valid rate falls: 1/p calls on average.

What does a strict schema guarantee about tool arguments, and what doesn't it? It guarantees the shape: field names, types, required fields, enum values. It does not guarantee the values are true or allowed: a customer ID can be well formed and not exist, and an amount can fit the type while breaking a minimum. Your code still validates business rules.

The papers behind this lesson

  • Willard and Louf, Efficient Guided Generation for Large Language Models (2023): https://arxiv.org/abs/2307.09702. Recast constrained generation as moving between the states of a finite-state machine, with the allowed tokens indexed per state in advance, the design behind Outlines.
  • Geng, Josifoski, Peyrard and West, Grammar-Constrained Decoding for Structured NLP Tasks without Finetuning (2023): https://arxiv.org/abs/2305.13971. Showed that constraining decoding with a formal grammar lets an off-the-shelf model produce complex structured outputs reliably, with no task-specific training.
  • Scholak, Schucher and Bahdanau, PICARD: Parsing Incrementally for Constrained Auto-Regressive Decoding from Language Models (2021): https://arxiv.org/abs/2109.05093. Rejected tokens that an incremental parser could not accept, making generated SQL valid as it was written.
  • Dong et al., XGrammar: Flexible and Efficient Structured Generation Engine for Large Language Models (2024): https://arxiv.org/abs/2411.15100. Made grammar-constrained decoding fast by prechecking context-independent tokens and checking only the context-dependent ones against the stack.
  • Park et al., Grammar-Aligned Decoding (2024): https://arxiv.org/abs/2405.21047. Showed that token-by-token masking distorts the model's distribution over valid outputs, and proposed a way to sample closer to the model's own conditional odds.
  • Tam et al., Let Me Speak Freely? A Study on the Impact of Format Restrictions on Performance of Large Language Models (2024): https://arxiv.org/abs/2408.02442. Measured how strict output formats can lower a model's reasoning performance.

Further reading

on GitHub
   1r"""
   2# Structured output: answers that fit a shape, every time
   3
   4Run: `python -m primer.ml.structured_output`
   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 tool
   8arguments from `primer.agents.tools`.
   9
  10## Level 1: The practitioner's guide
  11
  12**In one sentence.** Structured output means making a model's answer fit an
  13exact shape (a JSON object, a label from a fixed list, a date) every time, so
  14that a program can read it with no person in the loop.
  15
  16**When you need it.** The moment something other than a person reads the
  17answer: a tool call's arguments, a row for a database, a classification
  18label, a form to fill in. A person shrugs off a stray word in front of the
  19JSON; a parser rejects the whole answer. You don't need it for prose a
  20person will read, and you don't need it when the reader is another model that
  21copes with loose text. The tell: if you are writing code to strip "Sure, here
  22you go!" from the front of responses, you need it.
  23
  24How often does asking nicely fail? This lesson's toy model, asked 200 times
  25for `{"age": 42}` with only the prompt to guide it, gets it exactly right 133
  26times (66.5%). Real models do far better than a toy, but not perfectly, and
  27the failures grow with length: in an answer of 100 tokens where each token is
  28right 98% of the time, the whole answer comes out valid only about 13% of the
  29time (Level 2 shows why). At a million calls a day, a 1% failure rate is ten
  30thousand broken answers a day.
  31
  32**Your options.** Five ways, from the cheapest to the most certain:
  33
  34| Option | What it does | What it guarantees | What it costs | Where it lives |
  35|---|---|---|---|---|
  36| Prompting and examples | Ask for the format and show an example or two | Nothing; it raises the odds | A few extra input tokens | Your prompt |
  37| Post-processing | Repair the common slips: strip chatter, close a bracket | Nothing, but it catches the frequent cases | A small parser you maintain | Your code |
  38| Validate and retry | Parse; on failure, ask again, quoting the error | Valid eventually, if the model is usually right | A whole extra call per retry, and latency | Your code |
  39| Constrained decoding | Forbid, at every token, anything that cannot lead to a valid answer | A valid shape, by construction | A grammar compiled once, a small check per token, some drift in what the model says | The model server: JSON mode, strict schemas, grammar engines |
  40| Fine-tuning on the format | Train the model on thousands of examples in the shape | Far more reliable; still not certain | Data, a training run, a model to host | Training |
  41
  42**How to choose.** Start from what reads the answer and how often it may be
  43wrong.
  44
  45- One field, a label from a list, a yes or no: prompt for it and validate.
  46  Retries are cheap because the answer is short.
  47- A JSON object your code depends on, at volume: use the hosted API's strict
  48  schema, or a grammar engine in front of a model you run yourself. It
  49  removes the parsing code, the type checks and the retry loop in one move.
  50- A format no engine supports (a custom mini-language, a legacy fixed-width
  51  record): post-process what you can, validate, retry, and consider
  52  fine-tuning once the volume justifies it.
  53- Whatever you pick, validate the values afterwards. A schema proves shape,
  54  not truth: `{"amount": 0, "currency": "USD"}` fits a payment schema exactly
  55  and is still a bad payment.
  56
  57**What it costs.** Prompting costs tokens. Retries cost whole calls and
  58double the slowest requests. Constrained decoding costs a one-time compile of
  59the schema (a noticeable pause on the first request, cached after that) and a
  60small check per token; nested schemas cost more than flat ones. It can also
  61cost quality: forcing a model off the path it wanted can make it invent a
  62value to satisfy a required field, or wander, legally, until the token budget
  63runs out. In this lesson's toy, a strict schema with a required `age` the
  64model has no answer for produces `"age": 9700`. Fine-tuning costs the most up
  65front and the least per call.
  66
  67**What breaks.**
  68
  69- **A required field the model can't fill** becomes a made-up value that
  70  parses. Make such fields optional, or allow null.
  71- **Valid JSON, wrong content.** JSON mode alone guarantees something
  72  parseable, not your keys or your types. That needs a schema.
  73- **Truncation.** A valid shape cut off by the output limit is invalid. Set
  74  the limit with the schema's size in mind.
  75- **Token boundaries.** A token can straddle a boundary in the grammar, so
  76  real engines check character by character; a home-made masker that judges
  77  whole tokens rejects valid answers.
  78- **Drift.** A constrained model can sound different, because the mask
  79  changes which continuations it is allowed.
  80- **Business rules.** Shape is checked; meaning is not. Keep the validator.
  81
  82**In the wild.** Hosted APIs expose the mechanism at three strengths: JSON
  83mode (any valid JSON), strict structured outputs (your schema, guaranteed)
  84and strict tool use (a tool's arguments must fit its input schema). Claude's
  85structured outputs and strict tool use are one example, linked in Further
  86reading. For models you run yourself, Outlines and XGrammar turn a JSON
  87Schema or a regular expression into token masks, and llama.cpp accepts a
  88grammar file (GBNF). Libraries such as instructor (validate and retry against
  89a Python type) and guidance (constrain a generation to a pattern or a set of
  90options) wrap these steps. Every agent framework leans on all of this: a
  91tool call is a structured output.
  92
  93**Go deeper.** Level 2 builds constrained decoding from nothing: why failures
  94compound with length (a one-line formula), how a mask is applied before each
  95token is drawn, how a pattern becomes a finite-state machine and a JSON
  96Schema becomes a stack, and what each pitfall above looks like in numbers you
  97can rerun. If you only needed to choose, you are done.
  98
  99## Level 2: How it works, from scratch
 100
 101A language model writes one token at a time, and each token is a draw from a
 102probability distribution (see `primer.ml.inference`). Most of the time you
 103want free text. Sometimes a program is going to read the answer: a tool call's
 104arguments, a row for a database, a label from a fixed list. Then the answer
 105must have an exact **shape**, such as a JSON object with an integer field
 106called `age`, and one stray character breaks it.
 107
 108**Structured output** is the family of techniques that make the shape
 109certain. The main one, **constrained decoding**, is surprisingly small: before
 110each token is drawn, find every token that could not possibly continue a valid
 111answer, and forbid it. This lesson builds that from scratch: first for simple
 112patterns, with a finite-state machine, then for nested JSON, with a stack.
 113
 114## 1. Asking nicely is not enough
 115
 116**Everyday picture.** You read a form aloud over the phone to a friend and ask
 117them to type it in exactly. They are careful and mostly right. But now and
 118then they add "Sure, here you go!" before the form, or they forget the last
 119bracket, or a finger slips. A person reading the result shrugs it off. A
 120program reading it stops dead: one wrong character and the whole thing is
 121rejected.
 122
 123**Tiny worked example.** This lesson's toy model, `ToyModel`, has mostly
 124learned to answer `{"age": 42}`. It is a pretend model, a few lines of NumPy,
 125that copies that answer with some noise on every choice and a habit of
 126opening with "Sure". Asked 200 times, with nothing but the request to guide
 127it, it writes the exact answer 133 times (66.5%). The other 67 look like this:
 128
 129| What it wrote | What went wrong |
 130|---|---|
 131| `Sure{"age": 42}` | a friendly word before the JSON (22 times) |
 132| `{"age": 42w` | a slip where the closing brace should be |
 133| `{"age": 428` | an extra digit, then it stopped |
 134| `{"age"6 42}` | a digit where the colon belongs |
 135| `Sure{"age": 42}EUR` | chatter at both ends |
 136
 137Nothing here is a big mistake. Each one is a single bad token.
 138
 139```mermaid
 140flowchart LR
 141  P["Prompt: reply with JSON"] --> M[Model draws one token]
 142  M --> T[Append it to the answer]
 143  T -->|not finished| M
 144  T -->|finished| J{Parser}
 145  J -->|every token right| OK[Valid JSON]
 146  J -->|one token wrong| ERR[Parse error:<br/>the whole answer is lost]
 147```
 148
 149**Reading it:** follow the loop on the left: the model adds one token at a
 150time and nothing checks the answer until it is finished. Only then does the
 151parser look at it, and it has two exits. One bad token anywhere in the loop
 152sends the whole answer to the error exit, however good the rest was.
 153
 154That is why failures pile up with length. If each token is right with chance
 155$p$, and the chances are roughly independent, the whole answer is right only
 156when every token is:
 157
 158$$
 159P(\text{valid}) = p^{\,n}
 160$$
 161
 162**Symbols**
 163
 164| Symbol | Meaning here | In the example |
 165|---|---|---|
 166| $p$ | the chance that one token is right, between 0 and 1 | 0.98 |
 167| $n$ | how many tokens the answer has | 10 |
 168| $p^{\,n}$ | $p$ multiplied by itself $n$ times: the chance that all $n$ are right | $0.98^{10}$ |
 169| $P(\text{valid})$ | the chance that the whole answer parses | 0.817 |
 170
 171**In words:** "the chance that a whole answer is valid is the chance that one
 172token is right, multiplied together once for every token."
 173
 174**With the numbers:** a model that gets 98% of tokens right, writing a
 17510-token answer, produces valid output $0.98^{10} = 0.817$ of the time: nearly
 176one answer in five is broken. Make the answer 100 tokens long and it drops to
 177$0.98^{100} = 0.133$.
 178
 179**In Python:**
 180
 181```python
 182# p: chance one token is right; n: tokens in the answer
 183p, n = 0.98, 10
 184# p^n: all n tokens right
 185round(p ** n, 3)  # → 0.817
 186# a ten times longer answer
 187round(p ** 100, 3)  # → 0.133
 188```
 189
 190![Valid-output rate falls as the answer grows: at 98% per token, 10 tokens are valid 82% of the time and 100 tokens only 13%](figures/primer.ml.structured_output.validity_vs_length.svg)
 191
 192**Reading it:** the x-axis is the length of the answer in tokens, the y-axis
 193the share of answers that come out valid. Each curve is one per-token
 194accuracy. Even the top curve, 99.9% per token, sags over a long answer, and
 195the 95% curve is near zero by 100 tokens. The dot marks the worked example.
 196Better prompting moves you to a higher curve, but no curve stays at 100%.
 197
 198**In code:** `ToyModel` is the pretend model, `sample_answer` draws one answer token by token, and `chance_all_valid` is the formula.
 199
 200**Why it matters in practice.** A program that calls a model a million times
 201a day and fails 1% of the time fails ten thousand times a day. Long answers,
 202nested objects and small models make it worse. Prompting and examples help,
 203but they only move you to a better curve. To get to 100% you have to change
 204how tokens are chosen.
 205
 206## 2. Constrained decoding: mask, then sample
 207
 208**Everyday picture.** Picture a keyboard whose keys lock and unlock as you
 209type. You still decide what to write. But at every keystroke, any key that
 210would make the text break the form is locked for that one keystroke. You
 211cannot make the mistake, because the key isn't there.
 212
 213**Tiny worked example.** At the very first step the model scores four tokens
 214(these scores are **logits**: raw preferences, before softmax turns them
 215into probabilities; see `primer.ml.attention` for softmax from zero). Only
 216`{` and `[` can start a JSON value, so the other two are locked:
 217
 218| Token | Logit $z$ | Allowed? $m$ | Share before the mask | Share after the mask |
 219|---|---|---|---|---|
 220| `Sure` | 2.0 | 0 | 0.579 | **0** |
 221| ` Here` | 1.0 | 0 | 0.213 | **0** |
 222| `{` | 0.5 | 1 | 0.129 | **0.622** |
 223| `[` | 0.0 | 1 | 0.078 | **0.378** |
 224
 225Before the mask the model put 79% of its belief on chatter. After the mask,
 226the chatter has a chance of exactly zero, and the two allowed tokens share
 227everything. They keep their odds against each other: `{` was 1.65 times as
 228likely as `[` before, and it still is.
 229
 230```mermaid
 231flowchart LR
 232  A[Answer so far] --> M[Model: a logit<br/>for every token]
 233  A --> C[Constraint: which tokens<br/>can still lead to a valid answer?]
 234  M --> K[Set forbidden logits to minus infinity]
 235  C --> K
 236  K --> S[Softmax over what is left<br/>then sample]
 237  S --> N{End token?}
 238  N -->|no| A2[Append the token] --> A
 239  N -->|yes| D[Done: valid by construction]
 240```
 241
 242**Reading it:** two arrows leave the answer so far. The top path is the
 243ordinary model, which scores every token as it always does. The bottom path
 244is the new part: a checker that knows the shape and says which tokens are
 245still possible. They meet in the mask box, and from there it is ordinary
 246sampling again. The end token is itself just a token, so it is allowed only
 247when the answer is complete.
 248
 249The mask as a formula is softmax with a 0-or-1 switch on every term:
 250
 251$$
 252p_i = \frac{m_i \, e^{z_i}}{\sum_{j=1}^{|V|} m_j \, e^{z_j}}
 253$$
 254
 255**Symbols**
 256
 257| Symbol | Meaning here | In the example |
 258|---|---|---|
 259| $V$ | the vocabulary: every token the model can write | the 4 tokens in the table |
 260| $\lvert V \rvert$ | how many tokens that is | 4 |
 261| $i$ | the token whose probability we are computing | 3, the `{` |
 262| $z_i$ | token $i$'s logit | $z_3 = 0.5$ |
 263| $m_i$ | the mask: 1 if token $i$ can continue a valid answer, 0 if not | $m = (0, 0, 1, 1)$ |
 264| $e^{z_i}$ | $e \approx 2.718$ raised to the logit: always positive | $e^{0.5} = 1.649$ |
 265| $\sum_{j=1}^{\lvert V \rvert}$ | add up the following for every token $j$ | $0 + 0 + 1.649 + 1$ |
 266| $p_i$ | the probability that token $i$ is drawn | 0.622 |
 267
 268**In words:** "a token's probability is its usual softmax share, except that
 269forbidden tokens count as zero on top and bottom, so the allowed tokens share
 270all the probability between them."
 271
 272**With the numbers:** $p_{\{} = \dfrac{1 \cdot e^{0.5}}{0 + 0 + 1 \cdot e^{0.5} + 1 \cdot e^{0}} = \dfrac{1.649}{2.649} = 0.622$.
 273
 274**In Python:**
 275
 276```python
 277import math
 278# Sure, " Here", "{", "["
 279z = [2.0, 1.0, 0.5, 0.0]
 280# m_i: 1 if the token may come next
 281m = [0, 0, 1, 1]
 282# m_i e^(z_i): forbidden tokens contribute nothing
 283kept = [m_i * math.exp(z_i) for m_i, z_i in zip(m, z)]
 284[round(k, 3) for k in kept]  # → [0.0, 0.0, 1.649, 1.0]
 285# Σ_j m_j e^(z_j)
 286total = sum(kept)
 287round(total, 3)  # → 2.649
 288[round(k / total, 3) for k in kept]  # → [0.0, 0.0, 0.622, 0.378]
 289# without the mask, most of the belief went to chatter
 290e = [math.exp(z_i) for z_i in z]
 291[round(x / sum(e), 3) for x in e]  # → [0.579, 0.213, 0.129, 0.078]
 292```
 293
 294Setting a logit to minus infinity does the same thing, because
 295$e^{-\infty} = 0$: it is the trick the causal mask uses in
 296`primer.ml.attention`. After masking, temperature, top-k and top-p from
 297`primer.ml.inference` work exactly as before, on the tokens that are left.
 298
 299![At the toy model's first step, 19% of its belief sits on Sure; after the mask, the brace gets all of it](figures/primer.ml.structured_output.masked_step.svg)
 300
 301**Reading it:** these are the toy model's real first-step probabilities for
 302the `{"age": 42}` task, top four tokens only. Grey bars are what the model
 303wanted; blue bars are what it may choose from after the mask. "Sure" had
 304nearly a fifth of the belief and drops to zero. The brace, the only legal way
 305to begin, takes everything.
 306
 307**Why 100% by construction.** The checker keeps one promise: *every answer so
 308far can still be finished validly*. It holds at the start (the empty answer
 309can be finished). Each step only allows a token that keeps it. The end token
 310is allowed only when the answer is already complete. So every answer that
 311ends is valid. No luck is involved.
 312
 313![Unconstrained, the age answer is valid 66.5% of the time and the person object 18%; constrained, every finished answer is valid](figures/primer.ml.structured_output.validity_bars.svg)
 314
 315**Reading it:** each bar is 200 answers from the toy model, split by what
 316happened. Green is valid. The unconstrained bars show all three failures:
 317chatter before the JSON, a broken character inside, and (rarely) running out
 318of token budget. The person object is longer, so it breaks far more often,
 319just as $p^n$ predicts. The constrained bars have no chatter and nothing
 320broken. The one sliver left, 5 of the 200 person answers, is **cut off**: the
 32140-token budget ran out mid-object. The guarantee covers every prefix, but
 322only finishing makes a whole answer, so leave room in the budget.
 323
 324**In code:** `masked_softmax` applies the mask, `sample_answer` takes an optional constraint and masks every step, and `validity_experiment` produces the bars above.
 325
 326**Why it matters in practice.** Constrained decoding changes nothing about the
 327model: no retraining, same weights. It only changes which tokens may be
 328drawn, which is why it can be added to any open model at serving time. The
 329hard part is the checker: answering "which of 100,000 tokens could still lead
 330to a valid answer?" fast, at every step. The next two sections build it.
 331
 332## 3. From a pattern to a state machine
 333
 334**Everyday picture.** A subway map. You stand at a station. Each line leaving
 335it is labelled with one character. To write a character, you ride the line
 336with that label; if no line from your station has it, that character is
 337impossible here. Some stations are marked "you may stop here". That map is a
 338**finite-state machine**: a fixed set of states (stations) and one move per
 339character. A **regular expression**, the pattern language behind
 340`[0-9]+` and `cat|car|dog`, can always be drawn as one.
 341
 342**Tiny worked example.** The pattern `cat|car|dog` (one of three words)
 343becomes this machine:
 344
 345```mermaid
 346stateDiagram-v2
 347  [*] --> S0
 348  S0 --> S1: c
 349  S0 --> S2: d
 350  S1 --> S3: a
 351  S2 --> S4: o
 352  S3 --> S5: r
 353  S3 --> S6: t
 354  S4 --> S7: g
 355  S5 --> [*]
 356  S6 --> [*]
 357  S7 --> [*]
 358```
 359
 360**Reading it:** start at S0 and read a word one character at a time. "cat"
 361goes S0, S1, S3, S6, and S6 has an exit arrow, so "cat" is accepted. "cow"
 362gets stuck at S1, which has no "o" line. After "ca" (state S3) only "r" and
 363"t" are possible: the machine has turned "what may come next?" into "which
 364lines leave this station?".
 365
 366**Tokens are not characters.** A model writes tokens, and one token can hold
 367several characters. The rule: **a token is allowed if the machine can swallow
 368all of its characters, one after another, without getting stuck.** With a
 369vocabulary of 13 tokens:
 370
 371| Token | From S0 | From S3 (after "ca") |
 372|---|---|---|
 373| `c`, `d` | allowed | stuck |
 374| `ca`, `cat`, `do`, `dog` | allowed: every character has a line | stuck |
 375| `a`, `at`, `o`, `og`, `g` | stuck at the first character | stuck |
 376| `t`, `r` | stuck | allowed |
 377
 378`cat` is allowed at the start even though no single line reads "cat": the
 379machine rides c, then a, then t. `at` is part of a real word and still never
 380allowed at the start, because S0 has no "a" line.
 381
 382$$
 383\delta^*(s, t) = \delta\big(\cdots\delta(\delta(s, c_1), c_2)\cdots, c_k\big)
 384\qquad
 385A(s) = \{\, t \in V : \delta^*(s, t) \text{ is defined} \,\}
 386$$
 387
 388**Symbols**
 389
 390| Symbol | Meaning here | In the example |
 391|---|---|---|
 392| $s$ | a state of the machine | S0 |
 393| $\delta(s, c)$ | the **transition function**: the state you reach from $s$ by reading character $c$, or undefined if there is no such line | $\delta(0, \text{c}) = 1$ |
 394| $t$ | a token | `cat` |
 395| $c_1, \ldots, c_k$ | the characters of $t$, in order; $k$ is how many | c, a, t; $k = 3$ |
 396| $\delta^*(s, t)$ | read every character of $t$ in turn, starting from $s$ | $\delta^*(0, \texttt{cat}) = 6$ |
 397| $V$ | the vocabulary | the 13 tokens above |
 398| $\{\, t \in V : \ldots \,\}$ | "the set of every token $t$ in $V$ for which ... holds" | |
 399| $A(s)$ | the allowed tokens at state $s$ | $A(0) = \{$c, d, ca, cat, do, dog$\}$ |
 400
 401The end token joins $A(s)$ only when $s$ is an **accepting** state (S5, S6
 402or S7 here), because only there is the text complete.
 403
 404**In words:** "to see whether a token fits, walk its characters through the
 405machine one at a time; the allowed tokens are the ones that never get stuck."
 406
 407**With the numbers:** $\delta^*(0, \texttt{cat}) = \delta(\delta(\delta(0, \text{c}), \text{a}), \text{t}) = \delta(\delta(1, \text{a}), \text{t}) = \delta(3, \text{t}) = 6$,
 408so `cat` is in $A(0)$. $\delta(0, \text{a})$ is undefined, so `at` is not.
 409
 410**In Python:**
 411
 412```python
 413# the cat|car|dog machine: state -> {character: next state}
 414delta = {0: {"c": 1, "d": 2}, 1: {"a": 3}, 2: {"o": 4}, 3: {"r": 5, "t": 6}, 4: {"g": 7}, 5: {}, 6: {}, 7: {}}
 415accepting = {5, 6, 7}
 416def walk(s, token):
 417    # δ*: one δ per character; None means stuck
 418    for ch in token:
 419        s = delta[s].get(ch) if s is not None else None
 420    return s
 421walk(0, "cat")  # → 6
 422walk(0, "at")  # → None
 423V = ["c", "a", "t", "r", "d", "o", "g", "ca", "cat", "do", "dog", "at", "og"]
 424# A(0): every token the machine can swallow whole from the start
 425[t for t in V if walk(0, t) is not None]  # → ['c', 'd', 'ca', 'cat', 'do', 'dog']
 426# A(3), after "ca"
 427[t for t in V if walk(3, t) is not None]  # → ['t', 'r']
 428# the end token only where the text is complete
 429walk(0, "cat") in accepting  # → True
 430```
 431
 432**How a pattern becomes a machine.** Two classic steps, both in the code:
 433
 434```mermaid
 435flowchart LR
 436  P["Pattern<br/>cat|car|dog"] --> T[Thompson's construction:<br/>one small machine per piece,<br/>glued with free jumps]
 437  T --> N[NFA: may be in<br/>several states at once]
 438  N --> D[Subset construction:<br/>each new state is a set<br/>of NFA states]
 439  D --> F[DFA: exactly one<br/>state at a time]
 440  F --> TB[Table: for every state,<br/>which tokens are allowed]
 441```
 442
 443**Reading it:** left to right, the pattern gets more mechanical. Thompson's
 444construction reads the pattern like a sentence and builds a tiny machine for
 445each piece: one line for a character, a fork for `|`, a loop for `+`. Glued
 446together they make an **NFA** (a non-deterministic machine), which can be in
 447several states at once: after "c" it is both "inside cat" and "inside car".
 448The subset construction turns each *set* of NFA states into one state of a
 449**DFA** (a deterministic machine), which is always in exactly one state, so
 450following it is a dictionary lookup. The last box is the payoff: since the
 451DFA has a fixed number of states, the allowed tokens can be worked out for
 452every state before generation starts.
 453
 454The pattern `-?[0-9]+` (an optional minus sign, then digits) compiles to just
 455three states: the start, "saw a minus", and "saw at least one digit", the only
 456accepting one. The lesson's running example, `\{"age": [0-9]+\}`, compiles to
 45711.
 458
 459![Each row is one state of the age machine, each column a token; only a handful of cells are lit, and the digit states allow many tokens at once](figures/primer.ml.structured_output.mask_table.svg)
 460
 461**Reading it:** rows are the 11 states, labelled by the text that reaches
 462them; columns are tokens; a dark cell means "allowed here". Most rows have one
 463or two dark cells: the structure is fixed, so only one next character is
 464legal, written alone or as the start of a longer token such as `"age"` or
 465`": `. The two digit rows are
 466where the model has real choice: any digit, a two-digit token such as `42`,
 467or, once one digit is down, the closing brace. The last row allows only the
 468end token. This whole table is computed once, before the first token.
 469
 470**In code:** `compile_pattern` runs both constructions and returns a `DFA`; `DFA.walk` follows text through it; `allowed_tokens` is $A(s)$, and `mask_table` precomputes it for every state. `RegexConstraint` plugs the table into `sample_answer`.
 471
 472**Why it matters in practice.** Integers, dates, enums, phone numbers and
 473fixed-key objects are all patterns. Reframing generation as moving between
 474the states of a machine, with the allowed tokens indexed per state, is the
 475idea of Willard and Louf (2023) behind the open-source Outlines library, and
 476it makes each step's mask a single lookup.
 477
 478## 4. JSON Schema: nesting needs a stack
 479
 480**Everyday picture.** A stack of plates. Each time you open something, a
 481bracket, a brace or a quote, you put a plate on the stack with a note: "an
 482array is open", "an object is open". To close something, you may only take the
 483**top** plate: close the most recent thing first. When the stack is empty,
 484you're done.
 485
 486A subway map can't do this job. JSON can nest as deep as you like: `[[[[...]]]]`.
 487A machine with, say, 50 states can't tell 50 open brackets from 51, so it
 488can't know how many closers it still owes. Counting without limit needs
 489memory without limit. A finite-state machine plus a stack is called a
 490**pushdown automaton**, and it is exactly enough for nested formats such as
 491JSON, SQL and most programming languages.
 492
 493**Tiny worked example.** Read `{"a": [1, {"b": ` one character at a time.
 494
 495```mermaid
 496flowchart LR
 497  S1["after {<br/>stack: object"] --> S2["after [<br/>stack: object, array"]
 498  S2 --> S3["after the inner {<br/>stack: object, array, object"]
 499  S3 --> S4["after 2}<br/>stack: object, array"]
 500  S4 --> S5["after ]}<br/>stack: empty, done"]
 501```
 502
 503**Reading it:** each box is a moment in the reading, with the stack written
 504bottom first. Every opener adds a frame on the right; every closer removes
 505the rightmost one. At the third box the innermost open thing is an object,
 506so `}` is the only closer allowed, and `]` would be refused even though an
 507array is open further down. The last box is empty: the value is complete, and
 508only now is the end token allowed.
 509
 510The checker answers one question about any prefix: can it still be finished?
 511
 512| Prefix | Status | Why |
 513|---|---|---|
 514| `{"a": [1, {"b": ` | open | a value for "b" may come next |
 515| `{"a": [1}` | dead | the array is on top, so `}` can't close it |
 516| `{"a": [1, {"b": 2}]}` | complete | the stack is empty |
 517
 518**The schema steers every character.** A **JSON Schema** (see
 519`primer.agents.tools`) names the fields and their types. This lesson's person
 520schema allows `name` (a string), `age` (an integer) and `pets` (an array of
 521"cat" or "dog"); `name` and `age` are required and nothing else is allowed.
 522The checker enforces each rule at the first character that breaks it:
 523
 524| Prefix | Status | The rule that decides |
 525|---|---|---|
 526| `{"na` | open | "na" can still become "name" |
 527| `{"nx` | dead | no field starts with "nx" |
 528| `{"name": "Ada"}` | dead | `age` is required, so the object may not close yet |
 529| `{"age": 3.` | dead | an integer has no decimal point |
 530| `{"age": 01` | dead | JSON numbers never start with a zero followed by digits |
 531| `{"pets": ["x` | dead | only "cat" and "dog" are listed |
 532| `{"name": "Ada", "age": 36}` | complete | every required field is present |
 533
 534Inside one frame, a small state machine does the work. Here is an object's:
 535
 536```mermaid
 537stateDiagram-v2
 538  [*] --> open: {
 539  open --> key: quote, if a field may be added
 540  key --> key: a letter that still spells an unused field
 541  key --> colon: closing quote, if the key is complete
 542  colon --> value: colon, then at most one space
 543  value --> after_value: the value's own frame is pushed and later popped
 544  after_value --> key: comma, space, quote
 545  after_value --> [*]: }, only if every required field is present
 546  open --> [*]: }, only if nothing is required
 547```
 548
 549**Reading it:** these are the phases one object frame moves through, from
 550opening brace to closing brace. The labels carry the schema's rules: a key
 551letter is allowed only if it still spells a field not yet used, and the
 552closing brace only if every required field is in. The value arrow is where
 553nesting happens: the value gets its own frame pushed on top of the stack, and
 554this frame waits in "value" until that frame is popped.
 555
 556The leaves of JSON (numbers, `true`, `false`, `null`) are regular patterns,
 557so the checker reads numbers with the machines from section 3. Real engines
 558split the work the same way: patterns for the small pieces, a stack for the
 559nesting.
 560
 561One rule here is deliberate: at most one space after `:` and `,`. JSON allows
 562any amount of whitespace, and a constrained model that has lost its way can
 563pad with spaces until the budget runs out. Real engines limit whitespace for
 564the same reason.
 565
 566**In code:** `SchemaChecker.step` reads one character and returns the new stack (a tuple of `Frame` entries), or nothing when the prefix is dead; `SchemaChecker.status` answers open, dead or complete for a prefix; `SchemaChecker.open_containers` shows the stack; `SchemaConstraint` tries every token against the stack at each step.
 567
 568**Why it matters in practice.** This is how a schema becomes a guarantee:
 569compile the schema into grammar rules, run them as a pushdown automaton, and
 570mask every token that would kill it. Geng et al. (2023) showed the approach
 571works for structured tasks without any fine-tuning, and PICARD (Scholak et
 572al., 2021) did the same for SQL by parsing incrementally. The stack has a
 573cost: it can grow without limit, so the masks can't all be tabled in advance
 574as they were for a pattern. Section 5 comes back to that.
 575
 576## 5. Costs and pitfalls
 577
 578Constrained decoding guarantees the shape. It does not guarantee a good
 579answer, and it isn't free.
 580
 581### The mask changes what the model says
 582
 583**Everyday picture.** A satellite navigation system that only forbids illegal
 584turns, one junction at a time, and never looks ahead. It happily takes the
 585motorway because the motorway looks fastest right now, and only at the end
 586discovers that the one legal exit is a long detour. A driver who could see the
 587whole map would have taken the side road from the start.
 588
 589**Tiny worked example 1: forced to invent.** The model is asked for a
 590person's age, but the text never says it. The model's honest belief for the
 591next token:
 592
 593| Token | Model's belief | Allowed by `"type": "integer"`? | After the mask |
 594|---|---|---|---|
 595| `null` | 0.80 | no | 0 |
 596| `3` | 0.12 | yes | 0.12 / 0.20 = **0.60** |
 597| `5` | 0.08 | yes | 0.08 / 0.20 = **0.40** |
 598
 599The mask throws away 80% of the model's belief, and the answer is a made-up
 600age, delivered as confidently as a real one. The toy model does this too:
 601told to write the person schema while it "wants" to write only a name, it is
 602forced to add an age, and invents numbers such as `9700`. The fix is in the
 603schema, not the decoder: allow `null` (`"type": ["integer", "null"]`), or add
 604a field for "not stated". A related finding (Tam et al., 2024): forcing a strict
 605format from the first token can hurt a model's reasoning, so let the model
 606reason in free text first, or put a reasoning field before the answer field.
 607
 608**Tiny worked example 2: locally tempting, globally wrong.** A two-token
 609model. Valid answers must end in "y".
 610
 611```mermaid
 612flowchart LR
 613  S[start] -->|A 0.9| A[A]
 614  S -->|B 0.1| B[B]
 615  A -->|x 0.99| AX["Ax: invalid"]
 616  A -->|y 0.01| AY["Ay: valid, 0.9 × 0.01 = 0.009"]
 617  B -->|x 0.1| BX["Bx: invalid"]
 618  B -->|y 0.9| BY["By: valid, 0.1 × 0.9 = 0.09"]
 619```
 620
 621**Reading it:** each arrow is one token with the model's probability for it,
 622and each leaf is a whole answer with the product of the probabilities along
 623its path. Of the two valid answers, the model much prefers By (0.09 against
 6240.009). But masking decides one token at a time. At the first step both A
 625and B can still end in "y", so nothing is masked, and A wins 90% of the time.
 626At the second step the mask forces "y". The result: Ay, the answer the model
 627itself thought ten times less likely, comes out 90% of the time.
 628
 629$$
 630P(y \mid \text{valid}) = \frac{P(y)}{\sum_{y' \in \mathcal{L}} P(y')}
 631\qquad
 632P_{\text{mask}}(y) = \prod_{i=1}^{n} p'(y_i \mid y_{<i})
 633$$
 634
 635**Symbols**
 636
 637| Symbol | Meaning here | In the example |
 638|---|---|---|
 639| $y$ | one whole answer | Ay |
 640| $y_i$ | its $i$-th token; $y_{<i}$ is every token before it | $y_2 =$ y, $y_{<2} =$ A |
 641| $n$ | the number of tokens in the answer | 2 |
 642| $P(y)$ | the model's own chance of writing $y$: its token chances multiplied together | $0.9 \times 0.01 = 0.009$ |
 643| $\mathcal{L}$ | the set of valid answers (the "language" the grammar allows) | {Ay, By} |
 644| $\sum_{y' \in \mathcal{L}}$ | add up over every valid answer $y'$ | $0.009 + 0.09$ |
 645| $P(y \mid \text{valid})$ | "the chance of $y$ given that the answer is valid": the model's own odds, among valid answers only | 0.091 |
 646| $p'(y_i \mid y_{<i})$ | the masked, renormalised chance of token $y_i$ at that step | $p'(\text{y} \mid \text{A}) = 1$ |
 647| $\prod_{i=1}^{n}$ | multiply together over every step | $0.9 \times 1$ |
 648| $P_{\text{mask}}(y)$ | the chance that token-by-token masking produces $y$ | 0.9 |
 649
 650**In words:** "what the model believes, restricted to valid answers, is each
 651valid answer's probability divided by the total of all valid ones; what
 652masking actually produces is the product of the step-by-step masked
 653probabilities, and the two need not agree."
 654
 655**With the numbers:** $P(\text{Ay} \mid \text{valid}) = 0.009 / 0.099 = 0.091$,
 656but $P_{\text{mask}}(\text{Ay}) = 0.9 \times 1 = 0.9$: ten times the model's
 657own preference.
 658
 659**In Python:**
 660
 661```python
 662first = {"A": 0.9, "B": 0.1}
 663second = {"A": {"x": 0.99, "y": 0.01}, "B": {"x": 0.1, "y": 0.9}}
 664# P(y) for each valid answer: the model's token chances multiplied
 665P = {a + "y": first[a] * second[a]["y"] for a in first}
 666{k: round(v, 3) for k, v in P.items()}  # → {'Ay': 0.009, 'By': 0.09}
 667# P(y | valid): divide by the total over valid answers
 668total = sum(P.values())
 669{k: round(v / total, 3) for k, v in P.items()}  # → {'Ay': 0.091, 'By': 0.909}
 670# P_mask: step 1 keeps both, step 2 renormalises "y" to 1
 671{a + "y": first[a] * (second[a]["y"] / second[a]["y"]) for a in first}  # → {'Ay': 0.9, 'By': 0.1}
 672```
 673
 674![The model's own odds among valid answers favour By 91 to 9; token-by-token masking produces Ay 90% of the time](figures/primer.ml.structured_output.distortion.svg)
 675
 676**Reading it:** two answers, two bars each. Grey is the model's own
 677preference among valid answers; blue is what masked decoding actually
 678produces. The bars point in opposite directions. Nothing invalid comes out,
 679yet the distribution is badly bent. Park et al. (2024) name this problem and
 680propose a correction; in practice it is milder when the model already writes
 681the format well, because then the mask rarely has to overrule it.
 682
 683The toy model shows the everyday version: once noise knocks it off its answer,
 684the mask keeps it legal but not sensible, and it writes digits until the
 685closing brace happens to win, as in `{"age": 4174209}`.
 686
 687**In code:** `forced_choice` renormalises a belief over the allowed tokens and reports the share thrown away; `NULLABLE_AGE_SCHEMA` is the fix; `distortion_example` computes both distributions above.
 688
 689### Token boundaries
 690
 691**Everyday picture.** You can say "forty-two", or spell it "four, two". Both
 692arrive at the same text, but only one is how you would naturally say it.
 693
 694**Tiny worked example.** The toy vocabulary has a `42` token and single
 695digits, so "42" can be written two ways: `42`, or `4` then `2`. The whole
 696answer `{"age": 42}` can be spelled 16 ways, and the person answer 1,536
 697ways. The mask allows every one of them, but a trained model has almost only
 698ever seen the first, its tokenizer's usual spelling (see
 699`primer.ml.tokenization`). When the mask forces it onto an unusual spelling,
 700it is in unfamiliar territory and its next choices get worse.
 701
 702```mermaid
 703flowchart LR
 704  S["after the space"] -->|42, the usual token| E[before the brace]
 705  S -->|4| M[after the 4]
 706  M -->|2| E
 707```
 708
 709**Reading it:** two paths lead from the same place to the same place and
 710write the same text. The top path is the one the model saw thousands of
 711times in training; the bottom one it rarely saw. The checker walks
 712characters, so it cannot tell them apart; only the model can.
 713
 714Two more boundary effects follow from the same fact. A single token can cross
 715several structural boundaries at once, such as `"}` closing a string and an
 716object together; walking the token character by character, as `allowed_tokens`
 717does, handles that. And if the prompt ends in the middle of what would
 718normally be one token (a prompt ending in `{"age":` when the model would
 719usually write `": ` as one token), the natural token is no longer available.
 720Some engines back up one token and let the model rewrite it, which is called
 721**token healing**.
 722
 723**In code:** `tokenizations` lists every way a vocabulary can spell a text.
 724
 725### The speed of computing masks
 726
 727**Everyday picture.** Before every keystroke, a proofreader checks every word
 728in the dictionary against the rules: slow. Or: a card for every station,
 729printed once, listing which words fit there: fast, once the cards exist.
 730
 731**Tiny worked example.** A real vocabulary has around 128,000 tokens. Say
 732they average 4 characters, the answer is 200 tokens long, and the pattern's
 733machine has 50 states. Checking every token at every step walks
 734200 × 128,000 × 4 = 102.4 million characters for one answer. Building the
 735table walks 50 × 128,000 × 4 = 25.6 million characters once; after that,
 736each step only reads one row of 128,000 yes-or-no entries.
 737
 738$$
 739W_{\text{naive}} = n \cdot \lvert V \rvert \cdot \bar{L}
 740\qquad
 741W_{\text{table}} = S \cdot \lvert V \rvert \cdot \bar{L} + n \cdot \lvert V \rvert
 742$$
 743
 744**Symbols**
 745
 746| Symbol | Meaning here | In the example |
 747|---|---|---|
 748| $W$ | work: characters walked or table entries read | |
 749| $n$ | tokens generated | 200 |
 750| $\lvert V \rvert$ | vocabulary size | 128,000 |
 751| $\bar{L}$ | the average token length in characters (the bar means "average") | 4 |
 752| $S$ | states in the machine | 50 |
 753| $\cdot$ | multiply | |
 754
 755**In words:** "the naive way walks every token at every step; the table walks
 756every token once per state, up front, then reads one row per step."
 757
 758**With the numbers:** $W_{\text{naive}} = 200 \cdot 128{,}000 \cdot 4 = 102{,}400{,}000$
 759for every answer. $W_{\text{table}} = 50 \cdot 128{,}000 \cdot 4 + 200 \cdot 128{,}000 = 51{,}200{,}000$
 760for the first answer, and only the second term, 25.6 million cheap reads,
 761for each answer after that.
 762
 763**In Python:**
 764
 765```python
 766n, V, L_bar, S = 200, 128_000, 4, 50
 767# W_naive = n · |V| · L̄
 768n * V * L_bar  # → 102400000
 769# W_table = S · |V| · L̄ + n · |V|
 770S * V * L_bar + n * V  # → 51200000
 771# ten answers: the table's build cost is paid once
 77210 * n * V * L_bar, S * V * L_bar + 10 * n * V  # → (1024000000, 281600000)
 773```
 774
 775![Checking every token at every step costs the same for every answer; the table pays once, then grows slowly](figures/primer.ml.structured_output.mask_cost.svg)
 776
 777**Reading it:** the x-axis counts answers generated with one schema, the
 778y-axis the total work so far. The naive line climbs steeply and steadily. The
 779table line starts above zero (the build) and then climbs gently (one
 780row read per step). They cross partway through the very first answer. That is why hosted
 781APIs compile a schema once and cache it: the first request with a new schema
 782is slower, and the ones after are fast.
 783
 784For a schema with nesting, the stack can grow without limit, so no table can
 785cover every situation. XGrammar (Dong et al., 2024) splits the vocabulary:
 786most tokens are **context-independent** (whether they fit depends only on
 787the current grammar position, not on what is deeper in the stack), and those
 788are prechecked into tables; only the few context-dependent ones are checked
 789against the stack at run time.
 790
 791**In code:** `mask_cost` is the formula; `mask_table` builds the table once and `RegexConstraint` reads a row per step, while `SchemaConstraint` walks every token against the stack at every step, the naive way.
 792
 793### When validating and retrying is enough
 794
 795**Everyday picture.** Instead of a form that can't be filled in wrong, you
 796let people fill in a blank sheet, check it, and hand it back with a note when
 797it is wrong. Fine if most people get it right first time; miserable if most
 798don't.
 799
 800**Tiny worked example.** The toy model writes a valid age answer 66.5% of the
 801time. Retrying until it succeeds takes 1 / 0.665 = 1.50 attempts on average,
 802and three tries in a row all fail only 0.335³ = 3.8% of the time. For the
 803longer person object, valid 18% of the time, it takes 5.56 attempts on
 804average: slow and expensive.
 805
 806```mermaid
 807flowchart LR
 808  A[Ask the model] --> V{Validate}
 809  V -->|valid| D[Use it]
 810  V -->|invalid| E[Send back the<br/>validator's message] --> A
 811  E -.->|too many tries| F[Give up or<br/>fall back]
 812```
 813
 814**Reading it:** the happy path goes straight through. Every failure costs a
 815full extra model call, round the loop. Sending back the validator's exact
 816message (as `primer.agents.tools.validate` produces) makes the second try
 817much more likely to succeed. The dotted exit is the budget: a loop needs a
 818limit.
 819
 820$$
 821\mathbb{E}[\text{attempts}] = \frac{1}{p}
 822\qquad
 823P(\text{all } k \text{ tries fail}) = (1 - p)^k
 824$$
 825
 826**Symbols**
 827
 828| Symbol | Meaning here | In the example |
 829|---|---|---|
 830| $p$ | the chance one try is valid | 0.665 |
 831| $\mathbb{E}[\ldots]$ | the **expected value**: the long-run average over many repeats | |
 832| $\frac{1}{p}$ | the average number of tries to the first success | 1 / 0.665 = 1.50 |
 833| $k$ | a number of tries | 3 |
 834| $1 - p$ | the chance one try fails | 0.335 |
 835| $(1 - p)^k$ | the chance that $k$ independent tries all fail | $0.335^3 = 0.0376$ |
 836
 837**In words:** "if each try succeeds with chance p, you need one over p tries
 838on average, and the chance that k tries all fail is the chance of one
 839failure, multiplied together k times."
 840
 841**With the numbers:** $1 / 0.665 = 1.50$; $(1 - 0.665)^3 = 0.0376$; for the
 842person object, $1 / 0.18 = 5.56$.
 843
 844**In Python:**
 845
 846```python
 847# the toy model's unconstrained rate on the age answer
 848p = 0.665
 849# E[attempts] = 1 / p
 850round(1 / p, 2)  # → 1.5
 851# (1 - p)^k: three tries, all invalid
 852round((1 - p) ** 3, 4)  # → 0.0376
 853# the longer person answer, valid 18% of the time
 854round(1 / 0.18, 2)  # → 5.56
 855```
 856
 857![Expected attempts are close to 1 for a model that is usually right, and shoot up as the valid rate falls](figures/primer.ml.structured_output.retry.svg)
 858
 859**Reading it:** the x-axis is the chance a single try is valid; the y-axis is
 860the average number of tries until one is. The curve is flat on the right
 861and steep on the left. The dots mark three cases: a model that is 95% right
 862barely notices retries (1.05 tries); the toy age answer needs half a try
 863extra; the toy person answer needs more than five calls.
 864
 865Validate and retry when: you can't change the decoder (a hosted model
 866without a structured-output option), the model is already right nearly every
 867time, or the rule can't be expressed as a grammar. Constrain when answers are
 868long or nested, the model is small, or latency matters. Either way, keep the
 869validator: a grammar enforces shape, not rules such as `minimum`, and never
 870truth.
 871
 872**In code:** `expected_attempts` and `chance_still_failing` are the two formulas.
 873
 874## 6. JSON mode, strict schemas and tool calls
 875
 876**Everyday picture.** Two paper forms. One only insists that you write in
 877block capitals: whatever you write is readable, but nothing says what goes
 878where. The other has a labelled box for each field, and tick boxes where
 879only certain answers are allowed. The first is **JSON mode**; the second is
 880a **strict schema**.
 881
 882**Tiny worked example.** Give the toy model a reference answer that leaves
 883out the age, `{"name": "Ada"}`, and ask 50 times. With JSON mode (the empty
 884schema `{}`, which allows any JSON value) every finished answer parses: 44
 885copies of `{"name": "Ada"}` with the required age missing, one
 886`{"name": 42628}` whose name is a number, and one bare `false`. All valid
 887JSON; not one valid person. (The other 3 ran out of budget.) With the person
 888schema, every one of the 37 answers that finish has an age, because the
 889closing brace stays locked until one is written. Since the model had no age
 890to give, it invents one, such as `{"name": "Ada","age": 9700,"pets": []}`:
 891section 5's warning in action. The other 13 show a second cost of forcing a
 892model off its path: it loses its way and wanders, legally, until the token
 893budget runs out.
 894
 895A tool call is the same thing with a name attached: the model's arguments are
 896a JSON object that must fit the tool's `input_schema` (see
 897`primer.agents.tools` and `primer.agents.llm`).
 898
 899```mermaid
 900sequenceDiagram
 901  participant App as Your code
 902  participant API as Model server
 903  participant G as Grammar engine
 904  App->>API: messages + tool with a strict JSON Schema
 905  API->>G: compile the schema (cached for next time)
 906  loop every token
 907    G-->>API: mask of allowed tokens
 908    API->>API: mask, softmax, sample
 909  end
 910  API-->>App: tool_use with arguments that fit the schema
 911  App->>App: business rules: does customer C-999 exist? is the amount above the minimum?
 912```
 913
 914**Reading it:** the grammar engine lives inside the model server, next to
 915sampling. It compiles the schema once (the slow first request from section
 9165) and then hands a mask to every step of the loop. What reaches your code is
 917guaranteed to have the right shape. The last arrow is the part no grammar
 918covers: whether the values are true and allowed. The payment tool in
 919`primer.agents.tools` makes this concrete: `{"amount": 0, "currency": "USD"}`
 920fits the schema's shape exactly, and `primer.agents.tools.validate` still
 921rejects it, because `"minimum": 0.01` is a rule about the value, not the
 922shape.
 923
 924**In code:** `SchemaConstraint` with the schema `{}` is JSON mode, and with `PERSON_SCHEMA` it is a strict schema; `primer.agents.tools.ToolRegistry.definitions` shows how a strict tool definition is sent.
 925
 926**Why it matters in practice.** You now have three tools and know what each
 927buys. JSON mode guarantees something parseable. A strict schema guarantees
 928the shape your code expects, so the parsing and type-checking code
 929disappears. Validation after the fact still catches what only your program
 930knows. Hosted APIs (for example Claude's structured outputs and strict tool
 931use) and open engines (Outlines, llama.cpp grammars, XGrammar) all run the
 932mechanism built in this lesson: a checker that masks the logits before every
 933draw.
 934
 935## In 20 seconds
 936
 937- **The problem:** every token is a chance to break the format, so an
 938  answer of n tokens is valid about $p^n$ of the time. Prompting raises p
 939  but never to 1.
 940- **Constrained decoding:** before each draw, set the logit of every token
 941  that can't lead to a valid answer to minus infinity, then sample as usual.
 942  Allow the end token only when the answer is complete. Valid by
 943  construction, if the answer is allowed to finish.
 944- **Patterns:** compile to a finite-state machine; a token is allowed if the
 945  machine can read all its characters; precompute the allowed tokens for
 946  every state, so each step is a lookup.
 947- **JSON Schema:** nesting needs a stack (a pushdown automaton); the schema
 948  decides which keys, types and closers are legal at each character.
 949- **Pitfalls:** the mask can force a made-up value or bend the model's
 950  preferences; tokens don't align with grammar pieces; nested grammars cost
 951  more to check. Retrying is fine when the model is usually right. Always
 952  validate business rules afterwards.
 953
 954## Self-test questions
 955
 956**Why does asking a model for JSON in the prompt fail some of the time, and why do longer outputs fail more?**
 957Each token is a separate draw with some small chance of being wrong, and a
 958single wrong token breaks the parse. The chance that all n tokens are right
 959is about $p^n$, which shrinks as n grows: 98% per token gives 82% at 10 tokens
 960and 13% at 100.
 961
 962**What exactly does constrained decoding change in the model?**
 963Nothing in the weights. At each step it sets the logits of forbidden tokens to
 964minus infinity (probability zero) before sampling. The allowed tokens keep
 965their relative odds, and temperature and top-p still apply to them.
 966
 967**Why is the output valid "by construction", and what can still go wrong with the shape?**
 968Every prefix is kept completable, and the end token is allowed only when the
 969answer is complete, so any answer that ends is valid. It can still be cut off
 970by the token budget, leaving a valid but unfinished prefix.
 971
 972**How do you decide whether a multi-character token is allowed in a given state?**
 973Walk its characters through the state machine one at a time from the current
 974state. It is allowed if the walk never gets stuck. `cat` is allowed at the
 975start of `cat|car|dog`; `at` is not, because the start has no "a" line.
 976
 977**Why can't a finite-state machine check arbitrary JSON?**
 978JSON nests without limit, and closing correctly requires remembering every
 979open bracket in order. A machine with a fixed number of states can't count
 980without limit. A stack, which a pushdown automaton adds, can.
 981
 982**Why can masks be precomputed for a pattern but not fully for a JSON Schema?**
 983A pattern's machine has a fixed number of states, so the allowed tokens for
 984each can be tabled once. With nesting the stack can take unboundedly many
 985forms. Engines precompute the tokens whose fate depends only on the current
 986position and check the rest against the stack at run time.
 987
 988**How can constrained decoding make answers worse?**
 989It forces the model's choices into the allowed set even when the model
 990believed something else: an integer field makes it invent a number when the
 991honest answer was null, and token-by-token masking can commit early to a
 992path the model thought unlikely overall. Allowing null, or letting the model
 993reason before the structured part, helps.
 994
 995**When is validating and retrying good enough?**
 996When the model is valid almost every time (95% needs about 1.05 calls on
 997average), when you can't change the decoder, or when the rule can't be
 998written as a grammar. It gets expensive fast as the valid rate falls: 1/p
 999calls on average.
1000
1001**What does a strict schema guarantee about tool arguments, and what doesn't it?**
1002It guarantees the shape: field names, types, required fields, enum values.
1003It does not guarantee the values are true or allowed: a customer ID can be
1004well formed and not exist, and an amount can fit the type while breaking a
1005minimum. Your code still validates business rules.
1006
1007## The papers behind this lesson
1008
1009- **Willard and Louf, *Efficient Guided Generation for Large Language
1010  Models* (2023)**: https://arxiv.org/abs/2307.09702. Recast constrained
1011  generation as moving between the states of a finite-state machine, with the
1012  allowed tokens indexed per state in advance, the design behind Outlines.
1013- **Geng, Josifoski, Peyrard and West, *Grammar-Constrained Decoding for
1014  Structured NLP Tasks without Finetuning* (2023)**:
1015  https://arxiv.org/abs/2305.13971. Showed that constraining decoding with a
1016  formal grammar lets an off-the-shelf model produce complex structured
1017  outputs reliably, with no task-specific training.
1018- **Scholak, Schucher and Bahdanau, *PICARD: Parsing Incrementally for
1019  Constrained Auto-Regressive Decoding from Language Models* (2021)**:
1020  https://arxiv.org/abs/2109.05093. Rejected tokens that an incremental
1021  parser could not accept, making generated SQL valid as it was written.
1022- **Dong et al., *XGrammar: Flexible and Efficient Structured Generation
1023  Engine for Large Language Models* (2024)**: https://arxiv.org/abs/2411.15100.
1024  Made grammar-constrained decoding fast by prechecking context-independent
1025  tokens and checking only the context-dependent ones against the stack.
1026- **Park et al., *Grammar-Aligned Decoding* (2024)**:
1027  https://arxiv.org/abs/2405.21047. Showed that token-by-token masking
1028  distorts the model's distribution over valid outputs, and proposed a way
1029  to sample closer to the model's own conditional odds.
1030- **Tam et al., *Let Me Speak Freely? A Study on the Impact of Format
1031  Restrictions on Performance of Large Language Models* (2024)**:
1032  https://arxiv.org/abs/2408.02442. Measured how strict output formats can
1033  lower a model's reasoning performance.
1034
1035## Further reading
1036
1037- Russ Cox, *Regular Expression Matching Can Be Simple And Fast* (Thompson's construction, explained): https://swtch.com/~rsc/regexp/regexp1.html
1038- *Understanding JSON Schema*: https://json-schema.org/understanding-json-schema
1039- RFC 8259, *The JavaScript Object Notation (JSON) Data Interchange Format*: https://www.rfc-editor.org/rfc/rfc8259
1040- Claude structured outputs (JSON outputs and strict tool use): https://platform.claude.com/docs/en/build-with-claude/structured-outputs
1041- llama.cpp, *GBNF Guide* (grammars for local models): https://github.com/ggml-org/llama.cpp/blob/master/grammars/README.md
1042- Outlines, structured generation library: https://github.com/dottxt-ai/outlines
1043- Willard and Louf (2023): https://arxiv.org/abs/2307.09702
1044- Dong et al., *XGrammar* (2024): https://arxiv.org/abs/2411.15100
1045- Park et al., *Grammar-Aligned Decoding* (2024): https://arxiv.org/abs/2405.21047
1046"""
1047
1048from __future__ import annotations
1049
1050from collections import Counter
1051from dataclasses import dataclass, field, replace
1052
1053import numpy as np
1054
1055from primer._show import banner, say, table, takeaway
1056
1057# ---------------------------------------------------------------------------
1058# 1. The toy vocabulary and the toy model
1059# ---------------------------------------------------------------------------
1060
1061EOS = "<eos>"  # the "I'm finished" token; it adds no text
1062
1063# A toy tokenizer vocabulary: every single character the examples need, plus
1064# a few multi-character pieces, the way a real BPE vocabulary (see
1065# primer.ml.tokenization) holds both letters and common chunks.
1066_SINGLE = list('{}[]":, -.!\'\n') + list("0123456789") + list("abcdefghijklmnopqrstuvwxyz") + list("ABDEGPSU")
1067_MULTI = [
1068    "Sure", " Here", '"age"', "age", '": ', ": ", ", ", '"}', "42", "17", "20", "null", "true", "false",
1069    '"name"', "name", "Ada", '"amount"', "amount", '"currency"', "currency", '"USD"', "USD", "EUR", "GBP",
1070]
1071VOCAB: list[str] = _SINGLE + _MULTI + [EOS]
1072
1073AGE_PATTERN = r'\{"age": [0-9]+\}'
1074AGE_REFERENCE = '{"age": 42}'
1075
1076
1077def chance_all_valid(p: float, n: int) -> float:
1078    """P(every one of n independent steps is right) = p ** n."""
1079    return p**n
1080
1081
1082class ToyModel:
1083    """A pretend language model that has *mostly* learned to write one answer.
1084
1085    It copies `reference` (starting from its first "{"), preferring longer
1086    tokens the way a real model prefers its tokenizer's usual pieces. Random
1087    noise on every logit stands in for everything the model is unsure of, and
1088    `chatter` is its habit of opening with "Sure". Deterministic given the rng.
1089    """
1090
1091    def __init__(self, reference: str, vocab: list[str] = VOCAB, skill: float = 8.0, noise: float = 1.0, chatter: float = 6.0):
1092        self.reference, self.vocab = reference, vocab
1093        self.skill, self.noise, self.chatter = skill, noise, chatter
1094
1095    def logits(self, text: str, rng: np.random.Generator) -> np.ndarray:
1096        z = rng.normal(0.0, self.noise, len(self.vocab))
1097        brace = text.find("{")
1098        # Where the model thinks it is in its answer: characters written since the first brace.
1099        remaining = self.reference if brace < 0 else self.reference[len(text) - brace:]
1100        if text == "" and "Sure" in self.vocab:
1101            z[self.vocab.index("Sure")] += self.chatter
1102        for i, token in enumerate(self.vocab):
1103            if token != EOS and remaining.startswith(token):
1104                # Longer pieces get a small bonus: a trained model has seen "42" far more than "4" then "2".
1105                z[i] += self.skill + 0.5 * (len(token) - 1)
1106        if brace >= 0 and remaining == "":
1107            z[self.vocab.index(EOS)] += self.skill
1108        return z
1109
1110
1111# ---------------------------------------------------------------------------
1112# 2. Masked sampling
1113# ---------------------------------------------------------------------------
1114
1115
1116def masked_softmax(logits: np.ndarray, allowed: np.ndarray) -> np.ndarray:
1117    """softmax over the allowed tokens only; forbidden tokens get exactly 0."""
1118    z = np.where(allowed, logits, -np.inf)
1119    z = z - z[allowed].max()  # the largest allowed logit becomes 0, so exp never overflows
1120    e = np.exp(z)  # exp(-inf) = 0 for every forbidden token
1121    return e / e.sum()
1122
1123
1124@dataclass
1125class Decoded:
1126    """One generated answer: its text, the tokens that spelled it, and whether it ended by choice."""
1127
1128    text: str
1129    tokens: list[str] = field(default_factory=list)
1130    finished: bool = False
1131    removed_mass: list[float] = field(default_factory=list)
1132
1133
1134def sample_answer(model: ToyModel, rng: np.random.Generator, constraint=None, max_tokens: int = 40) -> Decoded:
1135    """Sample one token at a time; with a constraint, mask every token that can't continue a valid output."""
1136    out = Decoded("")
1137    everything = np.ones(len(model.vocab), dtype=bool)
1138    for _ in range(max_tokens):
1139        logits = model.logits(out.text, rng)
1140        allowed = everything if constraint is None else constraint.mask(out.text)
1141        # How much of the model's own belief the mask throws away at this step.
1142        out.removed_mass.append(float(1 - masked_softmax(logits, everything)[allowed].sum()))
1143        token = model.vocab[int(rng.choice(len(logits), p=masked_softmax(logits, allowed)))]
1144        if token == EOS:
1145            out.finished = True
1146            break
1147        out.text += token
1148        out.tokens.append(token)
1149    return out
1150
1151
1152# ---------------------------------------------------------------------------
1153# 3. From a pattern to a finite-state machine
1154# ---------------------------------------------------------------------------
1155
1156ENUM_PATTERN = "cat|car|dog"
1157# A vocabulary small enough to check by hand: letters, a few chunks, and two
1158# chunks ("at", "og") that are real pieces of the words but can never start one.
1159ENUM_VOCAB = ["c", "a", "t", "r", "d", "o", "g", "ca", "cat", "do", "dog", "at", "og", EOS]
1160
1161
1162class _PatternParser:
1163    """Thompson's construction: turn a pattern into an NFA, one small fragment per piece.
1164
1165    Supported: literal characters, `\\x` escapes, classes like `[0-9]`,
1166    grouping `( )`, alternation `|`, and the repeats `*`, `+` and `?`.
1167    Each fragment is a (start, end) pair of state numbers; `edges[s]` lists
1168    (label, target) with label None for a free "epsilon" jump.
1169    """
1170
1171    SPECIAL = set("()[]|*+?\\")
1172
1173    def __init__(self, pattern: str):
1174        self.pattern, self.i = pattern, 0
1175        self.edges: list[list[tuple[frozenset[str] | None, int]]] = []
1176
1177    def new_state(self) -> int:
1178        self.edges.append([])
1179        return len(self.edges) - 1
1180
1181    def peek(self) -> str | None:
1182        return self.pattern[self.i] if self.i < len(self.pattern) else None
1183
1184    def take(self) -> str:
1185        self.i += 1
1186        return self.pattern[self.i - 1]
1187
1188    def alternation(self) -> tuple[int, int]:
1189        branches = [self.sequence()]
1190        while self.peek() == "|":
1191            self.take()
1192            branches.append(self.sequence())
1193        if len(branches) == 1:
1194            return branches[0]
1195        start, end = self.new_state(), self.new_state()
1196        for b_start, b_end in branches:
1197            self.edges[start].append((None, b_start))  # jump freely into any branch
1198            self.edges[b_end].append((None, end))
1199        return start, end
1200
1201    def sequence(self) -> tuple[int, int]:
1202        parts = []
1203        while self.peek() is not None and self.peek() not in "|)":
1204            parts.append(self.repeat())
1205        if not parts:  # the empty pattern matches the empty string
1206            s = self.new_state()
1207            return s, s
1208        for (_, prev_end), (next_start, _) in zip(parts, parts[1:]):
1209            self.edges[prev_end].append((None, next_start))  # glue the pieces end to start
1210        return parts[0][0], parts[-1][1]
1211
1212    def repeat(self) -> tuple[int, int]:
1213        inner_start, inner_end = self.atom()
1214        while self.peek() is not None and self.peek() in "*+?":
1215            op = self.take()
1216            start, end = self.new_state(), self.new_state()
1217            self.edges[start].append((None, inner_start))
1218            self.edges[inner_end].append((None, end))
1219            if op in "*+":
1220                self.edges[inner_end].append((None, inner_start))  # loop back for another round
1221            if op in "*?":
1222                self.edges[start].append((None, end))  # or skip the piece entirely
1223            inner_start, inner_end = start, end
1224        return inner_start, inner_end
1225
1226    def atom(self) -> tuple[int, int]:
1227        ch = self.take()
1228        if ch == "(":
1229            fragment = self.alternation()
1230            if self.take() != ")":
1231                raise ValueError(f"unclosed group in {self.pattern!r}")
1232            return fragment
1233        if ch == "[":
1234            chars = self.char_class()
1235        elif ch == "\\":
1236            chars = frozenset(self.take())
1237        elif ch in self.SPECIAL:
1238            raise ValueError(f"unexpected {ch!r} at position {self.i - 1} of {self.pattern!r}")
1239        else:
1240            chars = frozenset(ch)
1241        start, end = self.new_state(), self.new_state()
1242        self.edges[start].append((chars, end))
1243        return start, end
1244
1245    def char_class(self) -> frozenset[str]:
1246        chars: set[str] = set()
1247        while self.peek() != "]":
1248            lo = self.take()
1249            if self.peek() == "-" and self.pattern[self.i + 1] != "]":
1250                self.take()
1251                hi = self.take()
1252                chars.update(chr(c) for c in range(ord(lo), ord(hi) + 1))  # a range like 0-9
1253            else:
1254                chars.add(lo)
1255        self.take()
1256        return frozenset(chars)
1257
1258
1259@dataclass
1260class DFA:
1261    """A deterministic finite-state machine: one current state, one move per character.
1262
1263    `transitions[s]` maps a character to the next state; a missing entry is
1264    the dead end (the text can no longer match). States are numbered from 0,
1265    the start.
1266    """
1267
1268    transitions: list[dict[str, int]]
1269    accepting: set[int]
1270    start: int = 0
1271
1272    @property
1273    def n_states(self) -> int:
1274        return len(self.transitions)
1275
1276    def walk(self, state: int | None, text: str) -> int | None:
1277        """Follow `text` one character at a time; None means the machine is stuck."""
1278        for ch in text:
1279            if state is None:
1280                return None
1281            state = self.transitions[state].get(ch)
1282        return state
1283
1284    def accepts(self, text: str) -> bool:
1285        return self.walk(self.start, text) in self.accepting
1286
1287
1288def compile_pattern(pattern: str) -> DFA:
1289    """Pattern -> NFA (Thompson) -> DFA (subset construction)."""
1290    parser = _PatternParser(pattern)
1291    nfa_start, nfa_accept = parser.alternation()
1292    if parser.i != len(pattern):
1293        raise ValueError(f"unexpected {pattern[parser.i]!r} in {pattern!r}")
1294    edges = parser.edges
1295
1296    def closure(states: set[int]) -> frozenset[int]:
1297        # Everywhere you can reach from these states by free jumps alone.
1298        stack, seen = list(states), set(states)
1299        while stack:
1300            for label, target in edges[stack.pop()]:
1301                if label is None and target not in seen:
1302                    seen.add(target)
1303                    stack.append(target)
1304        return frozenset(seen)
1305
1306    # Each DFA state is the *set* of NFA states you could be in at once.
1307    first = closure({nfa_start})
1308    ids, order, transitions = {first: 0}, [first], []
1309    for current in order:  # `order` grows as new sets are found: a breadth-first walk
1310        moves: dict[str, int] = {}
1311        chars = sorted({c for s in current for label, _ in edges[s] if label for c in label})
1312        for ch in chars:
1313            nxt = closure({t for s in current for label, t in edges[s] if label and ch in label})
1314            if nxt not in ids:
1315                ids[nxt] = len(order)
1316                order.append(nxt)
1317            moves[ch] = ids[nxt]
1318        transitions.append(moves)
1319    accepting = {ids[s] for s in order if nfa_accept in s}
1320    return DFA(transitions, accepting)
1321
1322
1323def allowed_tokens(dfa: DFA, state: int | None, vocab: list[str]) -> list[str]:
1324    """Every token the machine can consume *in full* from `state`; the end token only where the text is complete."""
1325    return [t for t in vocab if (state in dfa.accepting if t == EOS else dfa.walk(state, t) is not None)]
1326
1327
1328def mask_table(dfa: DFA, vocab: list[str]) -> np.ndarray:
1329    """Precompute the answer for every state once: row s is the allowed-token mask at state s."""
1330    table = np.zeros((dfa.n_states, len(vocab)), dtype=bool)
1331    for s in range(dfa.n_states):
1332        for j, t in enumerate(vocab):
1333            table[s, j] = s in dfa.accepting if t == EOS else dfa.walk(s, t) is not None
1334    return table
1335
1336
1337class RegexConstraint:
1338    """Constrain decoding to a pattern: compile once, build the table once, then each step is a lookup."""
1339
1340    def __init__(self, pattern: str, vocab: list[str]):
1341        self.dfa = compile_pattern(pattern)
1342        self.table = mask_table(self.dfa, vocab)
1343
1344    def mask(self, text: str) -> np.ndarray:
1345        state = self.dfa.walk(self.dfa.start, text)
1346        if state is None:
1347            raise ValueError(f"{text!r} already breaks the pattern")
1348        return self.table[state]
1349
1350
1351# ---------------------------------------------------------------------------
1352# 4. JSON Schema: nesting needs a stack
1353# ---------------------------------------------------------------------------
1354
1355PERSON_SCHEMA = {
1356    "type": "object",
1357    "properties": {
1358        "name": {"type": "string"},
1359        "age": {"type": "integer"},
1360        "pets": {"type": "array", "items": {"type": "string", "enum": ["cat", "dog"]}},
1361    },
1362    "required": ["name", "age"],
1363    # Strict: no fields beyond these. Constrained decoding needs to know every key it may spell.
1364    "additionalProperties": False,
1365}
1366PERSON_REFERENCE = '{"name": "Ada", "age": 36, "pets": ["cat"]}'
1367
1368# The leaves of JSON are regular, so the machines from section 3 read them.
1369# (JSON also allows an exponent, as in 1e5; this subset leaves it out.)
1370_NUMBER = compile_pattern(r"-?(0|[1-9][0-9]*)(\.[0-9]+)?")
1371_INTEGER = compile_pattern(r"-?(0|[1-9][0-9]*)")
1372_LITERALS = {"t": "true", "f": "false", "n": "null"}
1373
1374
1375@dataclass(frozen=True)
1376class Frame:
1377    """One entry on the stack: a value that has been opened but not yet finished.
1378
1379    kind is "root", "object", "array", "string", "number" or "literal".
1380    phase says what the frame expects next; text holds the characters of a
1381    key, string, number or literal read so far; seen holds an object's keys.
1382    """
1383
1384    kind: str
1385    schema: dict
1386    phase: str = ""
1387    text: str = ""
1388    seen: frozenset = frozenset()
1389
1390
1391Stack = tuple  # a tuple of Frames, innermost last; immutable, so every step returns a new one
1392
1393
1394def _types(schema: dict) -> set[str]:
1395    t = schema.get("type")
1396    if t is None:  # the empty schema {} allows any JSON value
1397        return {"string"} if "enum" in schema else {"object", "array", "string", "number", "integer", "boolean", "null"}
1398    return set(t) if isinstance(t, list) else {t}
1399
1400
1401class SchemaChecker:
1402    """A pushdown automaton for a JSON Schema subset: a state machine plus a stack.
1403
1404    Covers objects (properties, required, additionalProperties), arrays
1405    (items), strings (with optional enum), integers, numbers, booleans and
1406    null, with at most one space after ":" and ",". Strings have no escapes.
1407    Every step either refuses the character or leaves a stack from which the
1408    value can still be finished, so "not dead" always means "completable".
1409    """
1410
1411    def __init__(self, schema: dict):
1412        self.schema = schema
1413
1414    def start(self) -> Stack:
1415        return (Frame("root", self.schema, "value"),)
1416
1417    # --- helpers ----------------------------------------------------------
1418
1419    @staticmethod
1420    def _strict(schema: dict) -> bool:
1421        return schema.get("additionalProperties") is False
1422
1423    def _unseen_keys(self, frame: Frame) -> list[str]:
1424        return [k for k in frame.schema.get("properties", {}) if k not in frame.seen]
1425
1426    def _may_add_key(self, frame: Frame) -> bool:
1427        return not self._strict(frame.schema) or bool(self._unseen_keys(frame))
1428
1429    def _may_close(self, frame: Frame) -> bool:
1430        return set(frame.schema.get("required", [])) <= frame.seen
1431
1432    def _open_value(self, stack: Stack, schema: dict, ch: str) -> Stack | None:
1433        """The first character of a value decides which kind of frame to push."""
1434        types = _types(schema)
1435        if ch == "{" and "object" in types:
1436            return stack + (Frame("object", schema, "open"),)
1437        if ch == "[" and "array" in types:
1438            return stack + (Frame("array", schema, "open"),)
1439        if ch == '"' and "string" in types:
1440            return stack + (Frame("string", schema),)
1441        if (ch == "-" or ch.isdigit()) and types & {"number", "integer"}:
1442            return stack + (Frame("number", schema, text=ch),)
1443        if ch in _LITERALS and ("boolean" if ch in "tf" else "null") in types:
1444            return stack + (Frame("literal", schema, _LITERALS[ch], ch),)
1445        return None
1446
1447    def _close(self, stack: Stack) -> Stack:
1448        """The innermost value is finished: pop it and tell its parent."""
1449        stack = stack[:-1]
1450        parent = stack[-1]
1451        if parent.kind == "object":
1452            return stack[:-1] + (replace(parent, phase="after_value", seen=parent.seen | {parent.text}),)
1453        return stack[:-1] + (replace(parent, phase="done" if parent.kind == "root" else "after_value"),)
1454
1455    def _child_schema(self, frame: Frame) -> dict:
1456        if frame.kind == "array":
1457            return frame.schema.get("items", {})
1458        extra = frame.schema.get("additionalProperties", {})
1459        return frame.schema.get("properties", {}).get(frame.text, extra if isinstance(extra, dict) else {})
1460
1461    # --- one character ----------------------------------------------------
1462
1463    def step(self, stack: Stack | None, ch: str) -> Stack | None:
1464        """Read one character. Returns the new stack, or None if no valid value starts this way."""
1465        if stack is None:
1466            return None
1467        top = stack[-1]
1468        rest = stack[:-1]
1469        k, phase = top.kind, top.phase
1470
1471        if k == "root":
1472            return self._open_value(stack, top.schema, ch) if phase == "value" else None
1473
1474        if k == "string":
1475            if ch == '"':
1476                enum = top.schema.get("enum")
1477                return self._close(stack) if enum is None or top.text in enum else None
1478            if ch == "\\" or ord(ch) < 32:  # no escapes in this subset; raw control characters are never legal JSON
1479                return None
1480            text = top.text + ch
1481            enum = top.schema.get("enum")
1482            if enum is not None and not any(e.startswith(text) for e in enum):
1483                return None
1484            return rest + (replace(top, text=text),)
1485
1486        if k == "number":
1487            machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1488            if machine.walk(machine.start, top.text + ch) is not None:
1489                return rest + (replace(top, text=top.text + ch),)
1490            # The number can't grow: if it is finished, the character belongs to whatever comes after it.
1491            if machine.accepts(top.text):
1492                return self.step(self._close(stack), ch)
1493            return None
1494
1495        if k == "literal":
1496            text = top.text + ch
1497            if not top.phase.startswith(text):  # phase holds the target word here
1498                return None
1499            return self._close(stack) if text == top.phase else rest + (replace(top, text=text),)
1500
1501        if k == "object":
1502            if phase in ("open", "after_comma", "need_key") and ch == '"' and self._may_add_key(top):
1503                return rest + (replace(top, phase="key", text=""),)
1504            if phase == "after_comma" and ch == " ":
1505                return rest + (replace(top, phase="need_key"),)
1506            if phase in ("open", "after_value") and ch == "}" and self._may_close(top):
1507                return self._close(stack)
1508            if phase == "after_value" and ch == "," and self._may_add_key(top):
1509                return rest + (replace(top, phase="after_comma"),)
1510            if phase == "key":
1511                if ch == '"':
1512                    ok = not self._strict(top.schema) or top.text in self._unseen_keys(top)
1513                    return rest + (replace(top, phase="colon"),) if ok else None
1514                if ch == "\\" or ord(ch) < 32:
1515                    return None
1516                text = top.text + ch
1517                if self._strict(top.schema) and not any(key.startswith(text) for key in self._unseen_keys(top)):
1518                    return None
1519                return rest + (replace(top, text=text),)
1520            if phase == "colon":
1521                return rest + (replace(top, phase="space_or_value"),) if ch == ":" else None
1522            if phase == "space_or_value" and ch == " ":
1523                return rest + (replace(top, phase="value"),)
1524            if phase in ("space_or_value", "value"):
1525                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1526            return None
1527
1528        if k == "array":
1529            if phase in ("open", "after_value") and ch == "]":
1530                return self._close(stack)
1531            if phase == "after_value":
1532                return rest + (replace(top, phase="after_comma"),) if ch == "," else None
1533            if phase == "after_comma" and ch == " ":
1534                return rest + (replace(top, phase="need_value"),)
1535            if phase in ("open", "after_comma", "need_value"):
1536                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1537            return None
1538
1539        raise AssertionError(f"unknown frame kind {k!r}")
1540
1541    # --- whole prefixes ---------------------------------------------------
1542
1543    def feed(self, stack: Stack | None, text: str) -> Stack | None:
1544        for ch in text:
1545            stack = self.step(stack, ch)
1546            if stack is None:
1547                return None
1548        return stack
1549
1550    def can_end(self, stack: Stack | None) -> bool:
1551        """May the output stop here? Yes once the top-level value is finished."""
1552        if stack is None:
1553            return False
1554        if stack[-1].kind == "root":
1555            return stack[-1].phase == "done"
1556        # A top-level number has no closing character: "42" is finished when the output stops.
1557        top = stack[-1]
1558        machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1559        return len(stack) == 2 and top.kind == "number" and machine.accepts(top.text)
1560
1561    def status(self, prefix: str) -> str:
1562        """"dead" (no continuation can fix it), "complete" (may stop here) or "open" (valid so far)."""
1563        stack = self.feed(self.start(), prefix)
1564        return "dead" if stack is None else "complete" if self.can_end(stack) else "open"
1565
1566    def open_containers(self, prefix: str) -> list[str]:
1567        """The objects and arrays opened but not yet closed, outermost first: the visible part of the stack."""
1568        stack = self.feed(self.start(), prefix)
1569        if stack is None:
1570            raise ValueError(f"{prefix!r} can't be completed")
1571        return [f.kind for f in stack if f.kind in ("object", "array")]
1572
1573
1574class SchemaConstraint:
1575    """Constrain decoding to a JSON Schema: at each step, try every token against the stack."""
1576
1577    def __init__(self, schema: dict, vocab: list[str]):
1578        self.checker, self.vocab = SchemaChecker(schema), vocab
1579
1580    def mask(self, text: str) -> np.ndarray:
1581        stack = self.checker.feed(self.checker.start(), text)
1582        # Unlike a pattern's table, this can't be precomputed for every state: the stack can grow without limit.
1583        return np.array([self.checker.can_end(stack) if t == EOS else self.checker.feed(stack, t) is not None for t in self.vocab])
1584
1585
1586# ---------------------------------------------------------------------------
1587# 5. Costs and pitfalls
1588# ---------------------------------------------------------------------------
1589
1590NULLABLE_AGE_SCHEMA = {
1591    "type": "object",
1592    "properties": {"age": {"type": ["integer", "null"]}},  # "we don't know" is now a valid answer
1593    "required": ["age"],
1594    "additionalProperties": False,
1595}
1596
1597
1598def forced_choice(probs: dict[str, float], allowed: set[str]) -> tuple[dict[str, float], float]:
1599    """Renormalise a model's next-token belief over the allowed tokens; also return the share thrown away."""
1600    kept_mass = sum(p for t, p in probs.items() if t in allowed)
1601    return {t: p / kept_mass for t, p in probs.items() if t in allowed}, 1 - kept_mass
1602
1603
1604# A two-step model small enough to enumerate. Valid answers must end in "y".
1605TWO_STEP_FIRST = {"A": 0.9, "B": 0.1}
1606TWO_STEP_SECOND = {"A": {"x": 0.99, "y": 0.01}, "B": {"x": 0.1, "y": 0.9}}
1607
1608
1609def distortion_example() -> dict[str, dict[str, float]]:
1610    """Token-by-token masking versus the model's own odds among valid answers.
1611
1612    masked: at step 1 both A and B can still end in "y", so nothing is
1613    masked and A wins 90% of the time; at step 2 the mask forces "y".
1614    conditional: P(answer) / P(any valid answer), the distribution you would
1615    get by sampling whole answers and throwing away the invalid ones.
1616    """
1617    valid = {first + "y": TWO_STEP_FIRST[first] * TWO_STEP_SECOND[first]["y"] for first in TWO_STEP_FIRST}
1618    total = sum(valid.values())
1619    masked = {first + "y": TWO_STEP_FIRST[first] * 1.0 for first in TWO_STEP_FIRST}  # step 2 renormalises to y = 1
1620    return {"masked": masked, "conditional": {a: p / total for a, p in valid.items()}}
1621
1622
1623def tokenizations(text: str, vocab: list[str]) -> list[list[str]]:
1624    """Every way to spell `text` as a sequence of vocabulary tokens."""
1625    pieces = [t for t in vocab if t != EOS]
1626    ways: dict[int, list[list[str]]] = {len(text): [[]]}
1627    for i in range(len(text) - 1, -1, -1):  # from the end backwards, so each suffix is solved once
1628        ways[i] = [[t] + rest for t in pieces if text.startswith(t, i) for rest in ways[i + len(t)]]
1629    return ways[0]
1630
1631
1632def mask_cost(n_steps: int, vocab_size: int, avg_len: float, n_states: int) -> dict[str, float]:
1633    """Character steps spent building masks: walk every token at every step, or once per state up front."""
1634    naive = n_steps * vocab_size * avg_len
1635    table = n_states * vocab_size * avg_len + n_steps * vocab_size  # build once, then read one row per step
1636    return {"naive": naive, "table": table}
1637
1638
1639def expected_attempts(p: float) -> float:
1640    """Average number of tries until the first valid answer, when each try is valid with chance p."""
1641    return 1 / p
1642
1643
1644def chance_still_failing(p: float, k: int) -> float:
1645    """Chance that k independent tries are all invalid."""
1646    return (1 - p) ** k
1647
1648
1649NO_AGE_REFERENCE = '{"name": "Ada"}'  # a model that "forgot" the required age
1650
1651
1652def _age_ok(text: str) -> bool:
1653    import re
1654
1655    return re.fullmatch(AGE_PATTERN, text) is not None  # Python's own regex engine is the referee
1656
1657
1658def _person_ok(text: str) -> bool:
1659    import json
1660
1661    from primer.agents.tools import validate  # the tools lesson's validator is the referee, not our checker
1662
1663    try:
1664        return validate(json.loads(text), PERSON_SCHEMA) == []
1665    except ValueError:
1666        return False
1667
1668
1669def validity_experiment(n: int = 200, seed: int = 0) -> list[dict]:
1670    """Sample n answers per task, with and without constraints, and sort the failures by kind."""
1671    tasks = [
1672        ("age pattern", AGE_REFERENCE, _age_ok, lambda: RegexConstraint(AGE_PATTERN, VOCAB)),
1673        ("person schema", PERSON_REFERENCE, _person_ok, lambda: SchemaConstraint(PERSON_SCHEMA, VOCAB)),
1674    ]
1675    rows = []
1676    for task, reference, ok, make_constraint in tasks:
1677        for constrained in (False, True):
1678            rng = np.random.default_rng(seed)
1679            model, constraint = ToyModel(reference), make_constraint() if constrained else None
1680            counts = {"valid": 0, "chatter": 0, "cut_off": 0, "broken": 0}
1681            for _ in range(n):
1682                answer = sample_answer(model, rng, constraint)
1683                if ok(answer.text) and answer.finished:
1684                    counts["valid"] += 1
1685                elif not answer.text.startswith("{"):
1686                    counts["chatter"] += 1  # words before the JSON, such as "Sure"
1687                elif not answer.finished:
1688                    counts["cut_off"] += 1  # ran out of token budget mid-value
1689                else:
1690                    counts["broken"] += 1  # finished, but a wrong character somewhere inside
1691            rows.append({"task": task, "constrained": constrained, **counts})
1692    return rows
1693
1694
1695# ---------------------------------------------------------------------------
1696# 6. Figures (rendered into the HTML docs by `make figures`)
1697# ---------------------------------------------------------------------------
1698
1699
1700def _state_labels(dfa: DFA, reference: str) -> dict[int, str]:
1701    """Name each state by the shortest prefix of `reference` that reaches it."""
1702    labels: dict[int, str] = {}
1703    for i in range(len(reference) + 1):
1704        labels.setdefault(dfa.walk(dfa.start, reference[:i]), reference[:i] or "(start)")
1705    return labels
1706
1707
1708def figures() -> dict:
1709    """Plot this lesson's data. matplotlib is imported here, and only here,
1710    so the lesson itself needs nothing beyond NumPy."""
1711    import matplotlib
1712
1713    matplotlib.use("Agg")
1714    import matplotlib.pyplot as plt
1715
1716    BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1717    figs = {}
1718
1719    # --- 1. p ** n: longer answers break more often -------------------------
1720    fig, ax = plt.subplots(figsize=(6, 3.4))
1721    n = np.arange(1, 201)
1722    for p, color in ((0.999, GREEN), (0.99, BLUE), (0.98, AMBER), (0.95, RED)):
1723        ax.plot(n, chance_all_valid(p, n), color=color, label=f"{p:.1%} per token")
1724    ax.plot([10], [chance_all_valid(0.98, 10)], "o", color="#111827")
1725    ax.annotate("10 tokens at 98%: 0.82", (10, 0.817), (60, 0.72), arrowprops=dict(arrowstyle="-", color=MUTED),
1726                zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1727    ax.set_xlabel("answer length n (tokens)")
1728    ax.set_ylabel("share of answers valid, pⁿ")
1729    ax.set_ylim(0, 1.05)
1730    ax.set_title("Every token is a chance to break the format")
1731    ax.legend(frameon=False, loc="lower left")
1732    figs["validity_vs_length"] = fig
1733
1734    # --- 2. One step of masking --------------------------------------------
1735    model, rng = ToyModel(AGE_REFERENCE), np.random.default_rng(0)
1736    logits = model.logits("", rng)
1737    before = masked_softmax(logits, np.ones(len(VOCAB), dtype=bool))
1738    after = masked_softmax(logits, RegexConstraint(AGE_PATTERN, VOCAB).mask(""))
1739    top = np.argsort(before)[::-1][:4]
1740    fig, ax = plt.subplots(figsize=(6, 3.2))
1741    x = np.arange(len(top))
1742    ax.bar(x - 0.2, before[top], 0.4, color=MUTED, label="what the model wanted")
1743    ax.bar(x + 0.2, after[top], 0.4, color=BLUE, label="after the mask")
1744    for xi, (b, a) in enumerate(zip(before[top], after[top])):
1745        ax.text(xi - 0.2, b + 0.02, f"{b:.2f}", ha="center", fontsize=8)
1746        ax.text(xi + 0.2, a + 0.02, f"{a:.2f}", ha="center", fontsize=8, color=BLUE)
1747    ax.set_xticks(x, [repr(VOCAB[i]) for i in top])
1748    ax.set_ylabel("probability")
1749    ax.set_ylim(0, 1.15)
1750    ax.set_title('First token of {"age": 42}: the mask removes the chatter')
1751    ax.legend(frameon=False)
1752    figs["masked_step"] = fig
1753
1754    # --- 3. Validity with and without constraints ----------------------------
1755    rows = validity_experiment(n=200)
1756    fig, ax = plt.subplots(figsize=(7, 3.6))
1757    names = [f"{r['task']}\n{'constrained' if r['constrained'] else 'unconstrained'}" for r in rows]
1758    bottom = np.zeros(len(rows))
1759    for key, color, label in (("valid", GREEN, "valid"), ("chatter", AMBER, "chatter before the JSON"),
1760                              ("broken", RED, "broken inside"), ("cut_off", MUTED, "cut off by the token budget")):
1761        share = np.array([r[key] / 200 for r in rows])
1762        ax.bar(names, share, bottom=bottom, color=color, label=label)
1763        bottom += share
1764    for i, r in enumerate(rows):
1765        ax.text(i, r["valid"] / 200 / 2, f"{r['valid'] / 200:.1%}", ha="center", color="white", fontweight="bold")
1766    ax.set_ylabel("share of 200 answers")
1767    ax.set_title("Constrained decoding: every finished answer is valid")
1768    ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.0, 1.0))
1769    figs["validity_bars"] = fig
1770
1771    # --- 4. The precomputed mask table for the age pattern -------------------
1772    dfa = compile_pattern(AGE_PATTERN)
1773    table = mask_table(dfa, VOCAB)
1774    shown = [j for j in range(len(VOCAB)) if table[:, j].any()] + [VOCAB.index("Sure"), VOCAB.index("x")]
1775    labels = _state_labels(dfa, AGE_REFERENCE)
1776    fig, ax = plt.subplots(figsize=(9, 4.2))
1777    ax.imshow(table[:, shown], cmap="Blues", vmin=0, vmax=1.3, aspect="auto")
1778    ax.set_xticks(range(len(shown)), [repr(VOCAB[j]) if VOCAB[j] != EOS else "end" for j in shown], rotation=90, fontsize=8)
1779    ax.set_yticks(range(dfa.n_states), [labels[s] for s in range(dfa.n_states)], fontsize=8)
1780    ax.set_xlabel("token (every token allowed somewhere, plus two that never are)")
1781    ax.set_ylabel("state, named by the text that reaches it")
1782    ax.set_title("The age pattern's mask table, computed once before generation")
1783    ax.grid(False)
1784    figs["mask_table"] = fig
1785
1786    # --- 5. Distortion: masked decoding vs the model's own odds --------------
1787    dist = distortion_example()
1788    fig, ax = plt.subplots(figsize=(5.5, 3.2))
1789    x = np.arange(2)
1790    answers = ["Ay", "By"]
1791    cond = [dist["conditional"][a] for a in answers]
1792    masked = [dist["masked"][a] for a in answers]
1793    ax.bar(x - 0.2, cond, 0.4, color=MUTED, label="model's own odds among valid answers")
1794    ax.bar(x + 0.2, masked, 0.4, color=BLUE, label="what token-by-token masking produces")
1795    for xi, (c, m) in enumerate(zip(cond, masked)):
1796        ax.text(xi - 0.2, c + 0.02, f"{c:.2f}", ha="center")
1797        ax.text(xi + 0.2, m + 0.02, f"{m:.2f}", ha="center")
1798    ax.set_xticks(x, answers)
1799    ax.set_ylim(0, 1.2)
1800    ax.set_ylabel("probability")
1801    ax.set_title("Valid, but not what the model preferred")
1802    ax.legend(frameon=False, fontsize=8)
1803    figs["distortion"] = fig
1804
1805    # --- 6. Mask cost: every step vs a table built once ----------------------
1806    answers_made = np.arange(0, 11)
1807    per_answer = mask_cost(n_steps=200, vocab_size=128_000, avg_len=4, n_states=50)
1808    build = 50 * 128_000 * 4
1809    naive = answers_made * per_answer["naive"]
1810    tabled = build + answers_made * (per_answer["table"] - build)
1811    fig, ax = plt.subplots(figsize=(6, 3.4))
1812    ax.plot(answers_made, naive / 1e6, "o-", color=RED, label="check every token at every step")
1813    ax.plot(answers_made, tabled / 1e6, "o-", color=BLUE, label="build a table once, read a row per step")
1814    ax.set_xlabel("answers generated with one schema")
1815    ax.set_ylabel("work (millions of steps)")
1816    ax.set_title("|V| = 128,000, 200 tokens per answer, 50 states")
1817    ax.legend(frameon=False)
1818    figs["mask_cost"] = fig
1819
1820    # --- 7. Retrying: expected attempts --------------------------------------
1821    p = np.linspace(0.05, 1.0, 200)
1822    fig, ax = plt.subplots(figsize=(6, 3.4))
1823    tries = expected_attempts(p)
1824    ax.plot(p, np.where(tries <= 12, tries, np.nan), color=BLUE)  # past 12 tries it leaves the chart, below the title
1825    for value, label, dx, dy in ((0.95, "95% valid", -0.12, 1.4), (0.665, "toy age answer", -0.3, 3.2), (0.18, "toy person answer", 0.04, 1.6)):
1826        ax.plot([value], [expected_attempts(value)], "o", color=RED)
1827        # An opaque box above the curve keeps each label whole where the curve passes behind it.
1828        ax.annotate(f"{label}: {expected_attempts(value):.2f}", (value, expected_attempts(value)),
1829                    (value + dx, expected_attempts(value) + dy), fontsize=8, arrowprops=dict(arrowstyle="-", color=MUTED),
1830                    zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1831    ax.set_xlabel("chance a single try is valid, p")
1832    ax.set_ylabel("average tries until valid, 1/p")
1833    ax.set_ylim(0, 12)
1834    ax.set_title("Retrying is cheap only when the model is usually right")
1835    figs["retry"] = fig
1836
1837    return figs
1838
1839
1840# ---------------------------------------------------------------------------
1841# 7. Narrated walkthrough
1842# ---------------------------------------------------------------------------
1843
1844
1845def demo() -> None:
1846    banner("1. Asking nicely: the toy model, unconstrained")
1847    rng = np.random.default_rng(0)
1848    model = ToyModel(AGE_REFERENCE)
1849    answers = [sample_answer(model, rng).text for _ in range(8)]
1850    say(f"The toy model has mostly learned to write {AGE_REFERENCE}. Eight free answers:")
1851    for a in answers:
1852        print(f"  {'ok ' if _age_ok(a) else 'BAD'}  {a!r}")
1853    print()
1854    say(f"Per-token accuracy compounds: 0.98 ** 10 = {chance_all_valid(0.98, 10):.3f}, 0.98 ** 100 = {chance_all_valid(0.98, 100):.3f}.")
1855    takeaway("One bad token breaks the whole answer, so long outputs fail far more often than short ones.")
1856
1857    banner("2. Mask, then sample")
1858    z, m = np.array([2.0, 1.0, 0.5, 0.0]), np.array([False, False, True, True])
1859    table(["token", "logit", "allowed", "before", "after"],
1860          [(t, zi, bool(mi), b, a) for t, zi, mi, b, a in
1861           zip(["Sure", " Here", "{", "["], z, m, masked_softmax(z, np.ones(4, dtype=bool)), masked_softmax(z, m))],
1862          floatfmt=".3f")
1863    rows = validity_experiment(n=200)
1864    table(["task", "constrained", "valid", "chatter", "broken", "cut off"],
1865          [(r["task"], "yes" if r["constrained"] else "no", r["valid"], r["chatter"], r["broken"], r["cut_off"]) for r in rows])
1866    takeaway("Forbidden tokens get probability zero, so every answer that finishes is valid by construction.")
1867
1868    banner("3. A pattern becomes a state machine")
1869    dfa = compile_pattern(ENUM_PATTERN)
1870    say(f"'{ENUM_PATTERN}' compiles to {dfa.n_states} states. Allowed tokens, walking every character:")
1871    for prefix in ("", "ca", "dog"):
1872        print(f"  after {prefix!r:6}: {allowed_tokens(dfa, dfa.walk(dfa.start, prefix), ENUM_VOCAB)}")
1873    print()
1874    age = compile_pattern(AGE_PATTERN)
1875    say(f"The age pattern {AGE_PATTERN} has {age.n_states} states; its whole mask table is "
1876        f"{age.n_states} x {len(VOCAB)} entries, built once before the first token.")
1877    takeaway("A token is allowed if the machine can swallow all of its characters; precompute that per state.")
1878
1879    banner("4. JSON Schema needs a stack")
1880    anything = SchemaChecker({})
1881    say(f"Stack after '{{\"a\": [1, {{\"b\": ': {anything.open_containers('{\"a\": [1, {\"b\": ')}")
1882    person = SchemaChecker(PERSON_SCHEMA)
1883    table(["prefix", "status"], [(p, person.status(p)) for p in
1884          ('{"na', '{"nx', '{"name": "Ada"}', '{"age": 3.', '{"age": 01', '{"name": "Ada", "age": 36}')])
1885    rng = np.random.default_rng(0)
1886    constraint = SchemaConstraint(PERSON_SCHEMA, VOCAB)
1887    for _ in range(3):
1888        print("  constrained:", sample_answer(ToyModel(PERSON_REFERENCE), rng, constraint).text)
1889    print()
1890    takeaway("Nesting needs memory without limit; the schema decides which keys, types and closers are legal.")
1891
1892    banner("5. Costs and pitfalls")
1893    kept, removed = forced_choice({"null": 0.80, "3": 0.12, "5": 0.08}, {"3", "5"})
1894    say(f"Forcing an integer when the model wanted null throws away {removed:.0%} of its belief; "
1895        f"it now answers 3 with {kept['3']:.0%} and 5 with {kept['5']:.0%}: an invented age.")
1896    d = distortion_example()
1897    say(f"Two-step model: its own odds for Ay among valid answers are {d['conditional']['Ay']:.3f}, "
1898        f"but masking one token at a time produces Ay {d['masked']['Ay']:.0%} of the time.")
1899    say(f"'42' can be spelled {tokenizations('42', VOCAB)}; the whole age answer has "
1900        f"{len(tokenizations(AGE_REFERENCE, VOCAB))} spellings, and the mask allows them all.")
1901    c = mask_cost(n_steps=200, vocab_size=128_000, avg_len=4, n_states=50)
1902    say(f"Mask work for one 200-token answer: {c['naive']:,.0f} steps checking every token, "
1903        f"{c['table']:,.0f} with a table (including building it).")
1904    say(f"Retrying instead: {expected_attempts(0.665):.2f} tries on average at 66.5% valid, "
1905        f"{expected_attempts(0.18):.2f} at 18%.")
1906    takeaway("The mask guarantees shape, not good answers: design schemas that allow the honest answer.")
1907
1908    banner("6. JSON mode versus a strict schema")
1909    import json
1910
1911    from primer.agents.tools import PAYMENT_SCHEMA, validate
1912
1913    forgetful = ToyModel(NO_AGE_REFERENCE)
1914    say(f"A model whose answer leaves out the required age, {NO_AGE_REFERENCE}, asked 50 times:")
1915    for name, schema in (("JSON mode, schema {}", {}), ("strict person schema", PERSON_SCHEMA)):
1916        rng = np.random.default_rng(0)
1917        constraint = SchemaConstraint(schema, VOCAB)
1918        answers = [sample_answer(forgetful, rng, constraint) for _ in range(50)]
1919        finished = Counter(a.text for a in answers if a.finished)
1920        print(f"  {name}: {sum(finished.values())} finished, {50 - sum(finished.values())} cut off. Most common:")
1921        for text, count in finished.most_common(3):
1922            print(f"    {count:2d} x {text}")
1923    print()
1924    say("JSON mode parses but misses the age; the schema forces an age, which the model has to invent.")
1925    arguments = '{"amount": 0, "currency": "USD"}'
1926    say(f"Tool arguments {arguments}: the schema checker says '{SchemaChecker(PAYMENT_SCHEMA).status(arguments)}', "
1927        f"yet validation says {validate(json.loads(arguments), PAYMENT_SCHEMA)}.")
1928    takeaway("JSON mode guarantees parseable; a strict schema guarantees the shape; only your code checks the truth.")
1929
1930
1931if __name__ == "__main__":
1932    demo()
Level 3: the code, function by function.
EOS = '<eos>'
VOCAB: list[str] = ['{', '}', '[', ']', '"', ':', ',', ' ', '-', '.', '!', "'", '\n', '0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z', 'A', 'B', 'D', 'E', 'G', 'P', 'S', 'U', 'Sure', ' Here', '"age"', 'age', '": ', ': ', ', ', '"}', '42', '17', '20', 'null', 'true', 'false', '"name"', 'name', 'Ada', '"amount"', 'amount', '"currency"', 'currency', '"USD"', 'USD', 'EUR', 'GBP', '<eos>']
AGE_PATTERN = '\\{"age": [0-9]+\\}'
AGE_REFERENCE = '{"age": 42}'
def chance_all_valid(p: float, n: int) -> float: on GitHub
1078def chance_all_valid(p: float, n: int) -> float:
1079    """P(every one of n independent steps is right) = p ** n."""
1080    return p**n

P(every one of n independent steps is right) = p ** n.

class ToyModel: on GitHub
1083class ToyModel:
1084    """A pretend language model that has *mostly* learned to write one answer.
1085
1086    It copies `reference` (starting from its first "{"), preferring longer
1087    tokens the way a real model prefers its tokenizer's usual pieces. Random
1088    noise on every logit stands in for everything the model is unsure of, and
1089    `chatter` is its habit of opening with "Sure". Deterministic given the rng.
1090    """
1091
1092    def __init__(self, reference: str, vocab: list[str] = VOCAB, skill: float = 8.0, noise: float = 1.0, chatter: float = 6.0):
1093        self.reference, self.vocab = reference, vocab
1094        self.skill, self.noise, self.chatter = skill, noise, chatter
1095
1096    def logits(self, text: str, rng: np.random.Generator) -> np.ndarray:
1097        z = rng.normal(0.0, self.noise, len(self.vocab))
1098        brace = text.find("{")
1099        # Where the model thinks it is in its answer: characters written since the first brace.
1100        remaining = self.reference if brace < 0 else self.reference[len(text) - brace:]
1101        if text == "" and "Sure" in self.vocab:
1102            z[self.vocab.index("Sure")] += self.chatter
1103        for i, token in enumerate(self.vocab):
1104            if token != EOS and remaining.startswith(token):
1105                # Longer pieces get a small bonus: a trained model has seen "42" far more than "4" then "2".
1106                z[i] += self.skill + 0.5 * (len(token) - 1)
1107        if brace >= 0 and remaining == "":
1108            z[self.vocab.index(EOS)] += self.skill
1109        return z

A pretend language model that has mostly learned to write one answer.

It copies reference (starting from its first "{"), preferring longer tokens the way a real model prefers its tokenizer's usual pieces. Random noise on every logit stands in for everything the model is unsure of, and chatter is its habit of opening with "Sure". Deterministic given the rng.

ToyModel( reference: str, vocab: list[str] = ['{', '}', '[', ']', '"', ':', ',', ' ', '-', '.', '!', "'", '\n', '0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z', 'A', 'B', 'D', 'E', 'G', 'P', 'S', 'U', 'Sure', ' Here', '"age"', 'age', '": ', ': ', ', ', '"}', '42', '17', '20', 'null', 'true', 'false', '"name"', 'name', 'Ada', '"amount"', 'amount', '"currency"', 'currency', '"USD"', 'USD', 'EUR', 'GBP', '<eos>'], skill: float = 8.0, noise: float = 1.0, chatter: float = 6.0) on GitHub
1092    def __init__(self, reference: str, vocab: list[str] = VOCAB, skill: float = 8.0, noise: float = 1.0, chatter: float = 6.0):
1093        self.reference, self.vocab = reference, vocab
1094        self.skill, self.noise, self.chatter = skill, noise, chatter
def logits(self, text: str, rng: numpy.random._generator.Generator) -> numpy.ndarray: on GitHub
1096    def logits(self, text: str, rng: np.random.Generator) -> np.ndarray:
1097        z = rng.normal(0.0, self.noise, len(self.vocab))
1098        brace = text.find("{")
1099        # Where the model thinks it is in its answer: characters written since the first brace.
1100        remaining = self.reference if brace < 0 else self.reference[len(text) - brace:]
1101        if text == "" and "Sure" in self.vocab:
1102            z[self.vocab.index("Sure")] += self.chatter
1103        for i, token in enumerate(self.vocab):
1104            if token != EOS and remaining.startswith(token):
1105                # Longer pieces get a small bonus: a trained model has seen "42" far more than "4" then "2".
1106                z[i] += self.skill + 0.5 * (len(token) - 1)
1107        if brace >= 0 and remaining == "":
1108            z[self.vocab.index(EOS)] += self.skill
1109        return z
def masked_softmax(logits: numpy.ndarray, allowed: numpy.ndarray) -> numpy.ndarray: on GitHub
1117def masked_softmax(logits: np.ndarray, allowed: np.ndarray) -> np.ndarray:
1118    """softmax over the allowed tokens only; forbidden tokens get exactly 0."""
1119    z = np.where(allowed, logits, -np.inf)
1120    z = z - z[allowed].max()  # the largest allowed logit becomes 0, so exp never overflows
1121    e = np.exp(z)  # exp(-inf) = 0 for every forbidden token
1122    return e / e.sum()

softmax over the allowed tokens only; forbidden tokens get exactly 0.

@dataclass
class Decoded: on GitHub
1125@dataclass
1126class Decoded:
1127    """One generated answer: its text, the tokens that spelled it, and whether it ended by choice."""
1128
1129    text: str
1130    tokens: list[str] = field(default_factory=list)
1131    finished: bool = False
1132    removed_mass: list[float] = field(default_factory=list)

One generated answer: its text, the tokens that spelled it, and whether it ended by choice.

Decoded( text: str, tokens: list[str] = <factory>, finished: bool = False, removed_mass: list[float] = <factory>)
text: str
tokens: list[str]
finished: bool = False
removed_mass: list[float]
def sample_answer( model: ToyModel, rng: numpy.random._generator.Generator, constraint=None, max_tokens: int = 40) -> Decoded: on GitHub
1135def sample_answer(model: ToyModel, rng: np.random.Generator, constraint=None, max_tokens: int = 40) -> Decoded:
1136    """Sample one token at a time; with a constraint, mask every token that can't continue a valid output."""
1137    out = Decoded("")
1138    everything = np.ones(len(model.vocab), dtype=bool)
1139    for _ in range(max_tokens):
1140        logits = model.logits(out.text, rng)
1141        allowed = everything if constraint is None else constraint.mask(out.text)
1142        # How much of the model's own belief the mask throws away at this step.
1143        out.removed_mass.append(float(1 - masked_softmax(logits, everything)[allowed].sum()))
1144        token = model.vocab[int(rng.choice(len(logits), p=masked_softmax(logits, allowed)))]
1145        if token == EOS:
1146            out.finished = True
1147            break
1148        out.text += token
1149        out.tokens.append(token)
1150    return out

Sample one token at a time; with a constraint, mask every token that can't continue a valid output.

ENUM_PATTERN = 'cat|car|dog'
ENUM_VOCAB = ['c', 'a', 't', 'r', 'd', 'o', 'g', 'ca', 'cat', 'do', 'dog', 'at', 'og', '<eos>']
@dataclass
class DFA: on GitHub
1260@dataclass
1261class DFA:
1262    """A deterministic finite-state machine: one current state, one move per character.
1263
1264    `transitions[s]` maps a character to the next state; a missing entry is
1265    the dead end (the text can no longer match). States are numbered from 0,
1266    the start.
1267    """
1268
1269    transitions: list[dict[str, int]]
1270    accepting: set[int]
1271    start: int = 0
1272
1273    @property
1274    def n_states(self) -> int:
1275        return len(self.transitions)
1276
1277    def walk(self, state: int | None, text: str) -> int | None:
1278        """Follow `text` one character at a time; None means the machine is stuck."""
1279        for ch in text:
1280            if state is None:
1281                return None
1282            state = self.transitions[state].get(ch)
1283        return state
1284
1285    def accepts(self, text: str) -> bool:
1286        return self.walk(self.start, text) in self.accepting

A deterministic finite-state machine: one current state, one move per character.

transitions[s] maps a character to the next state; a missing entry is the dead end (the text can no longer match). States are numbered from 0, the start.

DFA( transitions: list[dict[str, int]], accepting: set[int], start: int = 0)
transitions: list[dict[str, int]]
accepting: set[int]
start: int = 0
n_states: int on GitHub
1273    @property
1274    def n_states(self) -> int:
1275        return len(self.transitions)
def walk(self, state: int | None, text: str) -> int | None: on GitHub
1277    def walk(self, state: int | None, text: str) -> int | None:
1278        """Follow `text` one character at a time; None means the machine is stuck."""
1279        for ch in text:
1280            if state is None:
1281                return None
1282            state = self.transitions[state].get(ch)
1283        return state

Follow text one character at a time; None means the machine is stuck.

def accepts(self, text: str) -> bool: on GitHub
1285    def accepts(self, text: str) -> bool:
1286        return self.walk(self.start, text) in self.accepting
def compile_pattern(pattern: str) -> DFA: on GitHub
1289def compile_pattern(pattern: str) -> DFA:
1290    """Pattern -> NFA (Thompson) -> DFA (subset construction)."""
1291    parser = _PatternParser(pattern)
1292    nfa_start, nfa_accept = parser.alternation()
1293    if parser.i != len(pattern):
1294        raise ValueError(f"unexpected {pattern[parser.i]!r} in {pattern!r}")
1295    edges = parser.edges
1296
1297    def closure(states: set[int]) -> frozenset[int]:
1298        # Everywhere you can reach from these states by free jumps alone.
1299        stack, seen = list(states), set(states)
1300        while stack:
1301            for label, target in edges[stack.pop()]:
1302                if label is None and target not in seen:
1303                    seen.add(target)
1304                    stack.append(target)
1305        return frozenset(seen)
1306
1307    # Each DFA state is the *set* of NFA states you could be in at once.
1308    first = closure({nfa_start})
1309    ids, order, transitions = {first: 0}, [first], []
1310    for current in order:  # `order` grows as new sets are found: a breadth-first walk
1311        moves: dict[str, int] = {}
1312        chars = sorted({c for s in current for label, _ in edges[s] if label for c in label})
1313        for ch in chars:
1314            nxt = closure({t for s in current for label, t in edges[s] if label and ch in label})
1315            if nxt not in ids:
1316                ids[nxt] = len(order)
1317                order.append(nxt)
1318            moves[ch] = ids[nxt]
1319        transitions.append(moves)
1320    accepting = {ids[s] for s in order if nfa_accept in s}
1321    return DFA(transitions, accepting)

Pattern -> NFA (Thompson) -> DFA (subset construction).

def allowed_tokens( dfa: DFA, state: int | None, vocab: list[str]) -> list[str]: on GitHub
1324def allowed_tokens(dfa: DFA, state: int | None, vocab: list[str]) -> list[str]:
1325    """Every token the machine can consume *in full* from `state`; the end token only where the text is complete."""
1326    return [t for t in vocab if (state in dfa.accepting if t == EOS else dfa.walk(state, t) is not None)]

Every token the machine can consume in full from state; the end token only where the text is complete.

def mask_table(dfa: DFA, vocab: list[str]) -> numpy.ndarray: on GitHub
1329def mask_table(dfa: DFA, vocab: list[str]) -> np.ndarray:
1330    """Precompute the answer for every state once: row s is the allowed-token mask at state s."""
1331    table = np.zeros((dfa.n_states, len(vocab)), dtype=bool)
1332    for s in range(dfa.n_states):
1333        for j, t in enumerate(vocab):
1334            table[s, j] = s in dfa.accepting if t == EOS else dfa.walk(s, t) is not None
1335    return table

Precompute the answer for every state once: row s is the allowed-token mask at state s.

class RegexConstraint: on GitHub
1338class RegexConstraint:
1339    """Constrain decoding to a pattern: compile once, build the table once, then each step is a lookup."""
1340
1341    def __init__(self, pattern: str, vocab: list[str]):
1342        self.dfa = compile_pattern(pattern)
1343        self.table = mask_table(self.dfa, vocab)
1344
1345    def mask(self, text: str) -> np.ndarray:
1346        state = self.dfa.walk(self.dfa.start, text)
1347        if state is None:
1348            raise ValueError(f"{text!r} already breaks the pattern")
1349        return self.table[state]

Constrain decoding to a pattern: compile once, build the table once, then each step is a lookup.

RegexConstraint(pattern: str, vocab: list[str]) on GitHub
1341    def __init__(self, pattern: str, vocab: list[str]):
1342        self.dfa = compile_pattern(pattern)
1343        self.table = mask_table(self.dfa, vocab)
dfa
table
def mask(self, text: str) -> numpy.ndarray: on GitHub
1345    def mask(self, text: str) -> np.ndarray:
1346        state = self.dfa.walk(self.dfa.start, text)
1347        if state is None:
1348            raise ValueError(f"{text!r} already breaks the pattern")
1349        return self.table[state]
PERSON_SCHEMA = {'type': 'object', 'properties': {'name': {'type': 'string'}, 'age': {'type': 'integer'}, 'pets': {'type': 'array', 'items': {'type': 'string', 'enum': ['cat', 'dog']}}}, 'required': ['name', 'age'], 'additionalProperties': False}
PERSON_REFERENCE = '{"name": "Ada", "age": 36, "pets": ["cat"]}'
@dataclass(frozen=True)
class Frame: on GitHub
1376@dataclass(frozen=True)
1377class Frame:
1378    """One entry on the stack: a value that has been opened but not yet finished.
1379
1380    kind is "root", "object", "array", "string", "number" or "literal".
1381    phase says what the frame expects next; text holds the characters of a
1382    key, string, number or literal read so far; seen holds an object's keys.
1383    """
1384
1385    kind: str
1386    schema: dict
1387    phase: str = ""
1388    text: str = ""
1389    seen: frozenset = frozenset()

One entry on the stack: a value that has been opened but not yet finished.

kind is "root", "object", "array", "string", "number" or "literal". phase says what the frame expects next; text holds the characters of a key, string, number or literal read so far; seen holds an object's keys.

Frame( kind: str, schema: dict, phase: str = '', text: str = '', seen: frozenset = frozenset())
kind: str
schema: dict
phase: str = ''
text: str = ''
seen: frozenset = frozenset()
Stack = <class 'tuple'>
class SchemaChecker: on GitHub
1402class SchemaChecker:
1403    """A pushdown automaton for a JSON Schema subset: a state machine plus a stack.
1404
1405    Covers objects (properties, required, additionalProperties), arrays
1406    (items), strings (with optional enum), integers, numbers, booleans and
1407    null, with at most one space after ":" and ",". Strings have no escapes.
1408    Every step either refuses the character or leaves a stack from which the
1409    value can still be finished, so "not dead" always means "completable".
1410    """
1411
1412    def __init__(self, schema: dict):
1413        self.schema = schema
1414
1415    def start(self) -> Stack:
1416        return (Frame("root", self.schema, "value"),)
1417
1418    # --- helpers ----------------------------------------------------------
1419
1420    @staticmethod
1421    def _strict(schema: dict) -> bool:
1422        return schema.get("additionalProperties") is False
1423
1424    def _unseen_keys(self, frame: Frame) -> list[str]:
1425        return [k for k in frame.schema.get("properties", {}) if k not in frame.seen]
1426
1427    def _may_add_key(self, frame: Frame) -> bool:
1428        return not self._strict(frame.schema) or bool(self._unseen_keys(frame))
1429
1430    def _may_close(self, frame: Frame) -> bool:
1431        return set(frame.schema.get("required", [])) <= frame.seen
1432
1433    def _open_value(self, stack: Stack, schema: dict, ch: str) -> Stack | None:
1434        """The first character of a value decides which kind of frame to push."""
1435        types = _types(schema)
1436        if ch == "{" and "object" in types:
1437            return stack + (Frame("object", schema, "open"),)
1438        if ch == "[" and "array" in types:
1439            return stack + (Frame("array", schema, "open"),)
1440        if ch == '"' and "string" in types:
1441            return stack + (Frame("string", schema),)
1442        if (ch == "-" or ch.isdigit()) and types & {"number", "integer"}:
1443            return stack + (Frame("number", schema, text=ch),)
1444        if ch in _LITERALS and ("boolean" if ch in "tf" else "null") in types:
1445            return stack + (Frame("literal", schema, _LITERALS[ch], ch),)
1446        return None
1447
1448    def _close(self, stack: Stack) -> Stack:
1449        """The innermost value is finished: pop it and tell its parent."""
1450        stack = stack[:-1]
1451        parent = stack[-1]
1452        if parent.kind == "object":
1453            return stack[:-1] + (replace(parent, phase="after_value", seen=parent.seen | {parent.text}),)
1454        return stack[:-1] + (replace(parent, phase="done" if parent.kind == "root" else "after_value"),)
1455
1456    def _child_schema(self, frame: Frame) -> dict:
1457        if frame.kind == "array":
1458            return frame.schema.get("items", {})
1459        extra = frame.schema.get("additionalProperties", {})
1460        return frame.schema.get("properties", {}).get(frame.text, extra if isinstance(extra, dict) else {})
1461
1462    # --- one character ----------------------------------------------------
1463
1464    def step(self, stack: Stack | None, ch: str) -> Stack | None:
1465        """Read one character. Returns the new stack, or None if no valid value starts this way."""
1466        if stack is None:
1467            return None
1468        top = stack[-1]
1469        rest = stack[:-1]
1470        k, phase = top.kind, top.phase
1471
1472        if k == "root":
1473            return self._open_value(stack, top.schema, ch) if phase == "value" else None
1474
1475        if k == "string":
1476            if ch == '"':
1477                enum = top.schema.get("enum")
1478                return self._close(stack) if enum is None or top.text in enum else None
1479            if ch == "\\" or ord(ch) < 32:  # no escapes in this subset; raw control characters are never legal JSON
1480                return None
1481            text = top.text + ch
1482            enum = top.schema.get("enum")
1483            if enum is not None and not any(e.startswith(text) for e in enum):
1484                return None
1485            return rest + (replace(top, text=text),)
1486
1487        if k == "number":
1488            machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1489            if machine.walk(machine.start, top.text + ch) is not None:
1490                return rest + (replace(top, text=top.text + ch),)
1491            # The number can't grow: if it is finished, the character belongs to whatever comes after it.
1492            if machine.accepts(top.text):
1493                return self.step(self._close(stack), ch)
1494            return None
1495
1496        if k == "literal":
1497            text = top.text + ch
1498            if not top.phase.startswith(text):  # phase holds the target word here
1499                return None
1500            return self._close(stack) if text == top.phase else rest + (replace(top, text=text),)
1501
1502        if k == "object":
1503            if phase in ("open", "after_comma", "need_key") and ch == '"' and self._may_add_key(top):
1504                return rest + (replace(top, phase="key", text=""),)
1505            if phase == "after_comma" and ch == " ":
1506                return rest + (replace(top, phase="need_key"),)
1507            if phase in ("open", "after_value") and ch == "}" and self._may_close(top):
1508                return self._close(stack)
1509            if phase == "after_value" and ch == "," and self._may_add_key(top):
1510                return rest + (replace(top, phase="after_comma"),)
1511            if phase == "key":
1512                if ch == '"':
1513                    ok = not self._strict(top.schema) or top.text in self._unseen_keys(top)
1514                    return rest + (replace(top, phase="colon"),) if ok else None
1515                if ch == "\\" or ord(ch) < 32:
1516                    return None
1517                text = top.text + ch
1518                if self._strict(top.schema) and not any(key.startswith(text) for key in self._unseen_keys(top)):
1519                    return None
1520                return rest + (replace(top, text=text),)
1521            if phase == "colon":
1522                return rest + (replace(top, phase="space_or_value"),) if ch == ":" else None
1523            if phase == "space_or_value" and ch == " ":
1524                return rest + (replace(top, phase="value"),)
1525            if phase in ("space_or_value", "value"):
1526                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1527            return None
1528
1529        if k == "array":
1530            if phase in ("open", "after_value") and ch == "]":
1531                return self._close(stack)
1532            if phase == "after_value":
1533                return rest + (replace(top, phase="after_comma"),) if ch == "," else None
1534            if phase == "after_comma" and ch == " ":
1535                return rest + (replace(top, phase="need_value"),)
1536            if phase in ("open", "after_comma", "need_value"):
1537                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1538            return None
1539
1540        raise AssertionError(f"unknown frame kind {k!r}")
1541
1542    # --- whole prefixes ---------------------------------------------------
1543
1544    def feed(self, stack: Stack | None, text: str) -> Stack | None:
1545        for ch in text:
1546            stack = self.step(stack, ch)
1547            if stack is None:
1548                return None
1549        return stack
1550
1551    def can_end(self, stack: Stack | None) -> bool:
1552        """May the output stop here? Yes once the top-level value is finished."""
1553        if stack is None:
1554            return False
1555        if stack[-1].kind == "root":
1556            return stack[-1].phase == "done"
1557        # A top-level number has no closing character: "42" is finished when the output stops.
1558        top = stack[-1]
1559        machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1560        return len(stack) == 2 and top.kind == "number" and machine.accepts(top.text)
1561
1562    def status(self, prefix: str) -> str:
1563        """"dead" (no continuation can fix it), "complete" (may stop here) or "open" (valid so far)."""
1564        stack = self.feed(self.start(), prefix)
1565        return "dead" if stack is None else "complete" if self.can_end(stack) else "open"
1566
1567    def open_containers(self, prefix: str) -> list[str]:
1568        """The objects and arrays opened but not yet closed, outermost first: the visible part of the stack."""
1569        stack = self.feed(self.start(), prefix)
1570        if stack is None:
1571            raise ValueError(f"{prefix!r} can't be completed")
1572        return [f.kind for f in stack if f.kind in ("object", "array")]

A pushdown automaton for a JSON Schema subset: a state machine plus a stack.

Covers objects (properties, required, additionalProperties), arrays (items), strings (with optional enum), integers, numbers, booleans and null, with at most one space after ":" and ",". Strings have no escapes. Every step either refuses the character or leaves a stack from which the value can still be finished, so "not dead" always means "completable".

SchemaChecker(schema: dict) on GitHub
1412    def __init__(self, schema: dict):
1413        self.schema = schema
schema
def start(self) -> tuple: on GitHub
1415    def start(self) -> Stack:
1416        return (Frame("root", self.schema, "value"),)
def step(self, stack: tuple | None, ch: str) -> tuple | None: on GitHub
1464    def step(self, stack: Stack | None, ch: str) -> Stack | None:
1465        """Read one character. Returns the new stack, or None if no valid value starts this way."""
1466        if stack is None:
1467            return None
1468        top = stack[-1]
1469        rest = stack[:-1]
1470        k, phase = top.kind, top.phase
1471
1472        if k == "root":
1473            return self._open_value(stack, top.schema, ch) if phase == "value" else None
1474
1475        if k == "string":
1476            if ch == '"':
1477                enum = top.schema.get("enum")
1478                return self._close(stack) if enum is None or top.text in enum else None
1479            if ch == "\\" or ord(ch) < 32:  # no escapes in this subset; raw control characters are never legal JSON
1480                return None
1481            text = top.text + ch
1482            enum = top.schema.get("enum")
1483            if enum is not None and not any(e.startswith(text) for e in enum):
1484                return None
1485            return rest + (replace(top, text=text),)
1486
1487        if k == "number":
1488            machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1489            if machine.walk(machine.start, top.text + ch) is not None:
1490                return rest + (replace(top, text=top.text + ch),)
1491            # The number can't grow: if it is finished, the character belongs to whatever comes after it.
1492            if machine.accepts(top.text):
1493                return self.step(self._close(stack), ch)
1494            return None
1495
1496        if k == "literal":
1497            text = top.text + ch
1498            if not top.phase.startswith(text):  # phase holds the target word here
1499                return None
1500            return self._close(stack) if text == top.phase else rest + (replace(top, text=text),)
1501
1502        if k == "object":
1503            if phase in ("open", "after_comma", "need_key") and ch == '"' and self._may_add_key(top):
1504                return rest + (replace(top, phase="key", text=""),)
1505            if phase == "after_comma" and ch == " ":
1506                return rest + (replace(top, phase="need_key"),)
1507            if phase in ("open", "after_value") and ch == "}" and self._may_close(top):
1508                return self._close(stack)
1509            if phase == "after_value" and ch == "," and self._may_add_key(top):
1510                return rest + (replace(top, phase="after_comma"),)
1511            if phase == "key":
1512                if ch == '"':
1513                    ok = not self._strict(top.schema) or top.text in self._unseen_keys(top)
1514                    return rest + (replace(top, phase="colon"),) if ok else None
1515                if ch == "\\" or ord(ch) < 32:
1516                    return None
1517                text = top.text + ch
1518                if self._strict(top.schema) and not any(key.startswith(text) for key in self._unseen_keys(top)):
1519                    return None
1520                return rest + (replace(top, text=text),)
1521            if phase == "colon":
1522                return rest + (replace(top, phase="space_or_value"),) if ch == ":" else None
1523            if phase == "space_or_value" and ch == " ":
1524                return rest + (replace(top, phase="value"),)
1525            if phase in ("space_or_value", "value"):
1526                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1527            return None
1528
1529        if k == "array":
1530            if phase in ("open", "after_value") and ch == "]":
1531                return self._close(stack)
1532            if phase == "after_value":
1533                return rest + (replace(top, phase="after_comma"),) if ch == "," else None
1534            if phase == "after_comma" and ch == " ":
1535                return rest + (replace(top, phase="need_value"),)
1536            if phase in ("open", "after_comma", "need_value"):
1537                return self._open_value(rest + (replace(top, phase="value"),), self._child_schema(top), ch)
1538            return None
1539
1540        raise AssertionError(f"unknown frame kind {k!r}")

Read one character. Returns the new stack, or None if no valid value starts this way.

def feed(self, stack: tuple | None, text: str) -> tuple | None: on GitHub
1544    def feed(self, stack: Stack | None, text: str) -> Stack | None:
1545        for ch in text:
1546            stack = self.step(stack, ch)
1547            if stack is None:
1548                return None
1549        return stack
def can_end(self, stack: tuple | None) -> bool: on GitHub
1551    def can_end(self, stack: Stack | None) -> bool:
1552        """May the output stop here? Yes once the top-level value is finished."""
1553        if stack is None:
1554            return False
1555        if stack[-1].kind == "root":
1556            return stack[-1].phase == "done"
1557        # A top-level number has no closing character: "42" is finished when the output stops.
1558        top = stack[-1]
1559        machine = _INTEGER if _types(top.schema) == {"integer"} else _NUMBER
1560        return len(stack) == 2 and top.kind == "number" and machine.accepts(top.text)

May the output stop here? Yes once the top-level value is finished.

def status(self, prefix: str) -> str: on GitHub
1562    def status(self, prefix: str) -> str:
1563        """"dead" (no continuation can fix it), "complete" (may stop here) or "open" (valid so far)."""
1564        stack = self.feed(self.start(), prefix)
1565        return "dead" if stack is None else "complete" if self.can_end(stack) else "open"

"dead" (no continuation can fix it), "complete" (may stop here) or "open" (valid so far).

def open_containers(self, prefix: str) -> list[str]: on GitHub
1567    def open_containers(self, prefix: str) -> list[str]:
1568        """The objects and arrays opened but not yet closed, outermost first: the visible part of the stack."""
1569        stack = self.feed(self.start(), prefix)
1570        if stack is None:
1571            raise ValueError(f"{prefix!r} can't be completed")
1572        return [f.kind for f in stack if f.kind in ("object", "array")]

The objects and arrays opened but not yet closed, outermost first: the visible part of the stack.

class SchemaConstraint: on GitHub
1575class SchemaConstraint:
1576    """Constrain decoding to a JSON Schema: at each step, try every token against the stack."""
1577
1578    def __init__(self, schema: dict, vocab: list[str]):
1579        self.checker, self.vocab = SchemaChecker(schema), vocab
1580
1581    def mask(self, text: str) -> np.ndarray:
1582        stack = self.checker.feed(self.checker.start(), text)
1583        # Unlike a pattern's table, this can't be precomputed for every state: the stack can grow without limit.
1584        return np.array([self.checker.can_end(stack) if t == EOS else self.checker.feed(stack, t) is not None for t in self.vocab])

Constrain decoding to a JSON Schema: at each step, try every token against the stack.

SchemaConstraint(schema: dict, vocab: list[str]) on GitHub
1578    def __init__(self, schema: dict, vocab: list[str]):
1579        self.checker, self.vocab = SchemaChecker(schema), vocab
def mask(self, text: str) -> numpy.ndarray: on GitHub
1581    def mask(self, text: str) -> np.ndarray:
1582        stack = self.checker.feed(self.checker.start(), text)
1583        # Unlike a pattern's table, this can't be precomputed for every state: the stack can grow without limit.
1584        return np.array([self.checker.can_end(stack) if t == EOS else self.checker.feed(stack, t) is not None for t in self.vocab])
NULLABLE_AGE_SCHEMA = {'type': 'object', 'properties': {'age': {'type': ['integer', 'null']}}, 'required': ['age'], 'additionalProperties': False}
def forced_choice( probs: dict[str, float], allowed: set[str]) -> tuple[dict[str, float], float]: on GitHub
1599def forced_choice(probs: dict[str, float], allowed: set[str]) -> tuple[dict[str, float], float]:
1600    """Renormalise a model's next-token belief over the allowed tokens; also return the share thrown away."""
1601    kept_mass = sum(p for t, p in probs.items() if t in allowed)
1602    return {t: p / kept_mass for t, p in probs.items() if t in allowed}, 1 - kept_mass

Renormalise a model's next-token belief over the allowed tokens; also return the share thrown away.

TWO_STEP_FIRST = {'A': 0.9, 'B': 0.1}
TWO_STEP_SECOND = {'A': {'x': 0.99, 'y': 0.01}, 'B': {'x': 0.1, 'y': 0.9}}
def distortion_example() -> dict[str, dict[str, float]]: on GitHub
1610def distortion_example() -> dict[str, dict[str, float]]:
1611    """Token-by-token masking versus the model's own odds among valid answers.
1612
1613    masked: at step 1 both A and B can still end in "y", so nothing is
1614    masked and A wins 90% of the time; at step 2 the mask forces "y".
1615    conditional: P(answer) / P(any valid answer), the distribution you would
1616    get by sampling whole answers and throwing away the invalid ones.
1617    """
1618    valid = {first + "y": TWO_STEP_FIRST[first] * TWO_STEP_SECOND[first]["y"] for first in TWO_STEP_FIRST}
1619    total = sum(valid.values())
1620    masked = {first + "y": TWO_STEP_FIRST[first] * 1.0 for first in TWO_STEP_FIRST}  # step 2 renormalises to y = 1
1621    return {"masked": masked, "conditional": {a: p / total for a, p in valid.items()}}

Token-by-token masking versus the model's own odds among valid answers.

masked: at step 1 both A and B can still end in "y", so nothing is masked and A wins 90% of the time; at step 2 the mask forces "y". conditional: P(answer) / P(any valid answer), the distribution you would get by sampling whole answers and throwing away the invalid ones.

def tokenizations(text: str, vocab: list[str]) -> list[list[str]]: on GitHub
1624def tokenizations(text: str, vocab: list[str]) -> list[list[str]]:
1625    """Every way to spell `text` as a sequence of vocabulary tokens."""
1626    pieces = [t for t in vocab if t != EOS]
1627    ways: dict[int, list[list[str]]] = {len(text): [[]]}
1628    for i in range(len(text) - 1, -1, -1):  # from the end backwards, so each suffix is solved once
1629        ways[i] = [[t] + rest for t in pieces if text.startswith(t, i) for rest in ways[i + len(t)]]
1630    return ways[0]

Every way to spell text as a sequence of vocabulary tokens.

def mask_cost( n_steps: int, vocab_size: int, avg_len: float, n_states: int) -> dict[str, float]: on GitHub
1633def mask_cost(n_steps: int, vocab_size: int, avg_len: float, n_states: int) -> dict[str, float]:
1634    """Character steps spent building masks: walk every token at every step, or once per state up front."""
1635    naive = n_steps * vocab_size * avg_len
1636    table = n_states * vocab_size * avg_len + n_steps * vocab_size  # build once, then read one row per step
1637    return {"naive": naive, "table": table}

Character steps spent building masks: walk every token at every step, or once per state up front.

def expected_attempts(p: float) -> float: on GitHub
1640def expected_attempts(p: float) -> float:
1641    """Average number of tries until the first valid answer, when each try is valid with chance p."""
1642    return 1 / p

Average number of tries until the first valid answer, when each try is valid with chance p.

def chance_still_failing(p: float, k: int) -> float: on GitHub
1645def chance_still_failing(p: float, k: int) -> float:
1646    """Chance that k independent tries are all invalid."""
1647    return (1 - p) ** k

Chance that k independent tries are all invalid.

NO_AGE_REFERENCE = '{"name": "Ada"}'
def validity_experiment(n: int = 200, seed: int = 0) -> list[dict]: on GitHub
1670def validity_experiment(n: int = 200, seed: int = 0) -> list[dict]:
1671    """Sample n answers per task, with and without constraints, and sort the failures by kind."""
1672    tasks = [
1673        ("age pattern", AGE_REFERENCE, _age_ok, lambda: RegexConstraint(AGE_PATTERN, VOCAB)),
1674        ("person schema", PERSON_REFERENCE, _person_ok, lambda: SchemaConstraint(PERSON_SCHEMA, VOCAB)),
1675    ]
1676    rows = []
1677    for task, reference, ok, make_constraint in tasks:
1678        for constrained in (False, True):
1679            rng = np.random.default_rng(seed)
1680            model, constraint = ToyModel(reference), make_constraint() if constrained else None
1681            counts = {"valid": 0, "chatter": 0, "cut_off": 0, "broken": 0}
1682            for _ in range(n):
1683                answer = sample_answer(model, rng, constraint)
1684                if ok(answer.text) and answer.finished:
1685                    counts["valid"] += 1
1686                elif not answer.text.startswith("{"):
1687                    counts["chatter"] += 1  # words before the JSON, such as "Sure"
1688                elif not answer.finished:
1689                    counts["cut_off"] += 1  # ran out of token budget mid-value
1690                else:
1691                    counts["broken"] += 1  # finished, but a wrong character somewhere inside
1692            rows.append({"task": task, "constrained": constrained, **counts})
1693    return rows

Sample n answers per task, with and without constraints, and sort the failures by kind.

def figures() -> dict: on GitHub
1709def figures() -> dict:
1710    """Plot this lesson's data. matplotlib is imported here, and only here,
1711    so the lesson itself needs nothing beyond NumPy."""
1712    import matplotlib
1713
1714    matplotlib.use("Agg")
1715    import matplotlib.pyplot as plt
1716
1717    BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1718    figs = {}
1719
1720    # --- 1. p ** n: longer answers break more often -------------------------
1721    fig, ax = plt.subplots(figsize=(6, 3.4))
1722    n = np.arange(1, 201)
1723    for p, color in ((0.999, GREEN), (0.99, BLUE), (0.98, AMBER), (0.95, RED)):
1724        ax.plot(n, chance_all_valid(p, n), color=color, label=f"{p:.1%} per token")
1725    ax.plot([10], [chance_all_valid(0.98, 10)], "o", color="#111827")
1726    ax.annotate("10 tokens at 98%: 0.82", (10, 0.817), (60, 0.72), arrowprops=dict(arrowstyle="-", color=MUTED),
1727                zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1728    ax.set_xlabel("answer length n (tokens)")
1729    ax.set_ylabel("share of answers valid, pⁿ")
1730    ax.set_ylim(0, 1.05)
1731    ax.set_title("Every token is a chance to break the format")
1732    ax.legend(frameon=False, loc="lower left")
1733    figs["validity_vs_length"] = fig
1734
1735    # --- 2. One step of masking --------------------------------------------
1736    model, rng = ToyModel(AGE_REFERENCE), np.random.default_rng(0)
1737    logits = model.logits("", rng)
1738    before = masked_softmax(logits, np.ones(len(VOCAB), dtype=bool))
1739    after = masked_softmax(logits, RegexConstraint(AGE_PATTERN, VOCAB).mask(""))
1740    top = np.argsort(before)[::-1][:4]
1741    fig, ax = plt.subplots(figsize=(6, 3.2))
1742    x = np.arange(len(top))
1743    ax.bar(x - 0.2, before[top], 0.4, color=MUTED, label="what the model wanted")
1744    ax.bar(x + 0.2, after[top], 0.4, color=BLUE, label="after the mask")
1745    for xi, (b, a) in enumerate(zip(before[top], after[top])):
1746        ax.text(xi - 0.2, b + 0.02, f"{b:.2f}", ha="center", fontsize=8)
1747        ax.text(xi + 0.2, a + 0.02, f"{a:.2f}", ha="center", fontsize=8, color=BLUE)
1748    ax.set_xticks(x, [repr(VOCAB[i]) for i in top])
1749    ax.set_ylabel("probability")
1750    ax.set_ylim(0, 1.15)
1751    ax.set_title('First token of {"age": 42}: the mask removes the chatter')
1752    ax.legend(frameon=False)
1753    figs["masked_step"] = fig
1754
1755    # --- 3. Validity with and without constraints ----------------------------
1756    rows = validity_experiment(n=200)
1757    fig, ax = plt.subplots(figsize=(7, 3.6))
1758    names = [f"{r['task']}\n{'constrained' if r['constrained'] else 'unconstrained'}" for r in rows]
1759    bottom = np.zeros(len(rows))
1760    for key, color, label in (("valid", GREEN, "valid"), ("chatter", AMBER, "chatter before the JSON"),
1761                              ("broken", RED, "broken inside"), ("cut_off", MUTED, "cut off by the token budget")):
1762        share = np.array([r[key] / 200 for r in rows])
1763        ax.bar(names, share, bottom=bottom, color=color, label=label)
1764        bottom += share
1765    for i, r in enumerate(rows):
1766        ax.text(i, r["valid"] / 200 / 2, f"{r['valid'] / 200:.1%}", ha="center", color="white", fontweight="bold")
1767    ax.set_ylabel("share of 200 answers")
1768    ax.set_title("Constrained decoding: every finished answer is valid")
1769    ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.0, 1.0))
1770    figs["validity_bars"] = fig
1771
1772    # --- 4. The precomputed mask table for the age pattern -------------------
1773    dfa = compile_pattern(AGE_PATTERN)
1774    table = mask_table(dfa, VOCAB)
1775    shown = [j for j in range(len(VOCAB)) if table[:, j].any()] + [VOCAB.index("Sure"), VOCAB.index("x")]
1776    labels = _state_labels(dfa, AGE_REFERENCE)
1777    fig, ax = plt.subplots(figsize=(9, 4.2))
1778    ax.imshow(table[:, shown], cmap="Blues", vmin=0, vmax=1.3, aspect="auto")
1779    ax.set_xticks(range(len(shown)), [repr(VOCAB[j]) if VOCAB[j] != EOS else "end" for j in shown], rotation=90, fontsize=8)
1780    ax.set_yticks(range(dfa.n_states), [labels[s] for s in range(dfa.n_states)], fontsize=8)
1781    ax.set_xlabel("token (every token allowed somewhere, plus two that never are)")
1782    ax.set_ylabel("state, named by the text that reaches it")
1783    ax.set_title("The age pattern's mask table, computed once before generation")
1784    ax.grid(False)
1785    figs["mask_table"] = fig
1786
1787    # --- 5. Distortion: masked decoding vs the model's own odds --------------
1788    dist = distortion_example()
1789    fig, ax = plt.subplots(figsize=(5.5, 3.2))
1790    x = np.arange(2)
1791    answers = ["Ay", "By"]
1792    cond = [dist["conditional"][a] for a in answers]
1793    masked = [dist["masked"][a] for a in answers]
1794    ax.bar(x - 0.2, cond, 0.4, color=MUTED, label="model's own odds among valid answers")
1795    ax.bar(x + 0.2, masked, 0.4, color=BLUE, label="what token-by-token masking produces")
1796    for xi, (c, m) in enumerate(zip(cond, masked)):
1797        ax.text(xi - 0.2, c + 0.02, f"{c:.2f}", ha="center")
1798        ax.text(xi + 0.2, m + 0.02, f"{m:.2f}", ha="center")
1799    ax.set_xticks(x, answers)
1800    ax.set_ylim(0, 1.2)
1801    ax.set_ylabel("probability")
1802    ax.set_title("Valid, but not what the model preferred")
1803    ax.legend(frameon=False, fontsize=8)
1804    figs["distortion"] = fig
1805
1806    # --- 6. Mask cost: every step vs a table built once ----------------------
1807    answers_made = np.arange(0, 11)
1808    per_answer = mask_cost(n_steps=200, vocab_size=128_000, avg_len=4, n_states=50)
1809    build = 50 * 128_000 * 4
1810    naive = answers_made * per_answer["naive"]
1811    tabled = build + answers_made * (per_answer["table"] - build)
1812    fig, ax = plt.subplots(figsize=(6, 3.4))
1813    ax.plot(answers_made, naive / 1e6, "o-", color=RED, label="check every token at every step")
1814    ax.plot(answers_made, tabled / 1e6, "o-", color=BLUE, label="build a table once, read a row per step")
1815    ax.set_xlabel("answers generated with one schema")
1816    ax.set_ylabel("work (millions of steps)")
1817    ax.set_title("|V| = 128,000, 200 tokens per answer, 50 states")
1818    ax.legend(frameon=False)
1819    figs["mask_cost"] = fig
1820
1821    # --- 7. Retrying: expected attempts --------------------------------------
1822    p = np.linspace(0.05, 1.0, 200)
1823    fig, ax = plt.subplots(figsize=(6, 3.4))
1824    tries = expected_attempts(p)
1825    ax.plot(p, np.where(tries <= 12, tries, np.nan), color=BLUE)  # past 12 tries it leaves the chart, below the title
1826    for value, label, dx, dy in ((0.95, "95% valid", -0.12, 1.4), (0.665, "toy age answer", -0.3, 3.2), (0.18, "toy person answer", 0.04, 1.6)):
1827        ax.plot([value], [expected_attempts(value)], "o", color=RED)
1828        # An opaque box above the curve keeps each label whole where the curve passes behind it.
1829        ax.annotate(f"{label}: {expected_attempts(value):.2f}", (value, expected_attempts(value)),
1830                    (value + dx, expected_attempts(value) + dy), fontsize=8, arrowprops=dict(arrowstyle="-", color=MUTED),
1831                    zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1832    ax.set_xlabel("chance a single try is valid, p")
1833    ax.set_ylabel("average tries until valid, 1/p")
1834    ax.set_ylim(0, 12)
1835    ax.set_title("Retrying is cheap only when the model is usually right")
1836    figs["retry"] = fig
1837
1838    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1846def demo() -> None:
1847    banner("1. Asking nicely: the toy model, unconstrained")
1848    rng = np.random.default_rng(0)
1849    model = ToyModel(AGE_REFERENCE)
1850    answers = [sample_answer(model, rng).text for _ in range(8)]
1851    say(f"The toy model has mostly learned to write {AGE_REFERENCE}. Eight free answers:")
1852    for a in answers:
1853        print(f"  {'ok ' if _age_ok(a) else 'BAD'}  {a!r}")
1854    print()
1855    say(f"Per-token accuracy compounds: 0.98 ** 10 = {chance_all_valid(0.98, 10):.3f}, 0.98 ** 100 = {chance_all_valid(0.98, 100):.3f}.")
1856    takeaway("One bad token breaks the whole answer, so long outputs fail far more often than short ones.")
1857
1858    banner("2. Mask, then sample")
1859    z, m = np.array([2.0, 1.0, 0.5, 0.0]), np.array([False, False, True, True])
1860    table(["token", "logit", "allowed", "before", "after"],
1861          [(t, zi, bool(mi), b, a) for t, zi, mi, b, a in
1862           zip(["Sure", " Here", "{", "["], z, m, masked_softmax(z, np.ones(4, dtype=bool)), masked_softmax(z, m))],
1863          floatfmt=".3f")
1864    rows = validity_experiment(n=200)
1865    table(["task", "constrained", "valid", "chatter", "broken", "cut off"],
1866          [(r["task"], "yes" if r["constrained"] else "no", r["valid"], r["chatter"], r["broken"], r["cut_off"]) for r in rows])
1867    takeaway("Forbidden tokens get probability zero, so every answer that finishes is valid by construction.")
1868
1869    banner("3. A pattern becomes a state machine")
1870    dfa = compile_pattern(ENUM_PATTERN)
1871    say(f"'{ENUM_PATTERN}' compiles to {dfa.n_states} states. Allowed tokens, walking every character:")
1872    for prefix in ("", "ca", "dog"):
1873        print(f"  after {prefix!r:6}: {allowed_tokens(dfa, dfa.walk(dfa.start, prefix), ENUM_VOCAB)}")
1874    print()
1875    age = compile_pattern(AGE_PATTERN)
1876    say(f"The age pattern {AGE_PATTERN} has {age.n_states} states; its whole mask table is "
1877        f"{age.n_states} x {len(VOCAB)} entries, built once before the first token.")
1878    takeaway("A token is allowed if the machine can swallow all of its characters; precompute that per state.")
1879
1880    banner("4. JSON Schema needs a stack")
1881    anything = SchemaChecker({})
1882    say(f"Stack after '{{\"a\": [1, {{\"b\": ': {anything.open_containers('{\"a\": [1, {\"b\": ')}")
1883    person = SchemaChecker(PERSON_SCHEMA)
1884    table(["prefix", "status"], [(p, person.status(p)) for p in
1885          ('{"na', '{"nx', '{"name": "Ada"}', '{"age": 3.', '{"age": 01', '{"name": "Ada", "age": 36}')])
1886    rng = np.random.default_rng(0)
1887    constraint = SchemaConstraint(PERSON_SCHEMA, VOCAB)
1888    for _ in range(3):
1889        print("  constrained:", sample_answer(ToyModel(PERSON_REFERENCE), rng, constraint).text)
1890    print()
1891    takeaway("Nesting needs memory without limit; the schema decides which keys, types and closers are legal.")
1892
1893    banner("5. Costs and pitfalls")
1894    kept, removed = forced_choice({"null": 0.80, "3": 0.12, "5": 0.08}, {"3", "5"})
1895    say(f"Forcing an integer when the model wanted null throws away {removed:.0%} of its belief; "
1896        f"it now answers 3 with {kept['3']:.0%} and 5 with {kept['5']:.0%}: an invented age.")
1897    d = distortion_example()
1898    say(f"Two-step model: its own odds for Ay among valid answers are {d['conditional']['Ay']:.3f}, "
1899        f"but masking one token at a time produces Ay {d['masked']['Ay']:.0%} of the time.")
1900    say(f"'42' can be spelled {tokenizations('42', VOCAB)}; the whole age answer has "
1901        f"{len(tokenizations(AGE_REFERENCE, VOCAB))} spellings, and the mask allows them all.")
1902    c = mask_cost(n_steps=200, vocab_size=128_000, avg_len=4, n_states=50)
1903    say(f"Mask work for one 200-token answer: {c['naive']:,.0f} steps checking every token, "
1904        f"{c['table']:,.0f} with a table (including building it).")
1905    say(f"Retrying instead: {expected_attempts(0.665):.2f} tries on average at 66.5% valid, "
1906        f"{expected_attempts(0.18):.2f} at 18%.")
1907    takeaway("The mask guarantees shape, not good answers: design schemas that allow the honest answer.")
1908
1909    banner("6. JSON mode versus a strict schema")
1910    import json
1911
1912    from primer.agents.tools import PAYMENT_SCHEMA, validate
1913
1914    forgetful = ToyModel(NO_AGE_REFERENCE)
1915    say(f"A model whose answer leaves out the required age, {NO_AGE_REFERENCE}, asked 50 times:")
1916    for name, schema in (("JSON mode, schema {}", {}), ("strict person schema", PERSON_SCHEMA)):
1917        rng = np.random.default_rng(0)
1918        constraint = SchemaConstraint(schema, VOCAB)
1919        answers = [sample_answer(forgetful, rng, constraint) for _ in range(50)]
1920        finished = Counter(a.text for a in answers if a.finished)
1921        print(f"  {name}: {sum(finished.values())} finished, {50 - sum(finished.values())} cut off. Most common:")
1922        for text, count in finished.most_common(3):
1923            print(f"    {count:2d} x {text}")
1924    print()
1925    say("JSON mode parses but misses the age; the schema forces an age, which the model has to invent.")
1926    arguments = '{"amount": 0, "currency": "USD"}'
1927    say(f"Tool arguments {arguments}: the schema checker says '{SchemaChecker(PAYMENT_SCHEMA).status(arguments)}', "
1928        f"yet validation says {validate(json.loads(arguments), PAYMENT_SCHEMA)}.")
1929    takeaway("JSON mode guarantees parseable; a strict schema guarantees the shape; only your code checks the truth.")