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
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
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.
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.
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.
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}
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)
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
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
- Russ Cox, Regular Expression Matching Can Be Simple And Fast (Thompson's construction, explained): https://swtch.com/~rsc/regexp/regexp1.html
- Understanding JSON Schema: https://json-schema.org/understanding-json-schema
- RFC 8259, The JavaScript Object Notation (JSON) Data Interchange Format: https://www.rfc-editor.org/rfc/rfc8259
- Claude structured outputs (JSON outputs and strict tool use): https://platform.claude.com/docs/en/build-with-claude/structured-outputs
- llama.cpp, GBNF Guide (grammars for local models): https://github.com/ggml-org/llama.cpp/blob/master/grammars/README.md
- Outlines, structured generation library: https://github.com/dottxt-ai/outlines
- Willard and Louf (2023): https://arxiv.org/abs/2307.09702
- Dong et al., XGrammar (2024): https://arxiv.org/abs/2411.15100
- Park et al., Grammar-Aligned Decoding (2024): https://arxiv.org/abs/2405.21047
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 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 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 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 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 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 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 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()
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.
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.
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
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.
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.
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.
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.
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.
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).
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.
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.
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.
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.
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".
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.
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.
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).
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.
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.
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])
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.
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.
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.
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.
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.
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.
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.
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.
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.")