Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

34. Sampling at scale and structured output

In this chapter

  • Everything an inference API lets a request ask for, and how one batched sampler serves a batch where every request asked for something different.
  • Penalties, logit bias, min_tokens, logprobs and prompt logprobs, and seeds that reproduce whatever else is in the batch.
  • Parallel samples (n > 1) and beam search, built on the prefix cache instead of special cache plumbing.
  • Constrained decoding: compiling regular expressions to a byte-level DFA, turning the DFA into per-token masks for a 150,000-token vocabulary, and JSON Schema to regex.

You will build

apply_processors, apply_penalties, top_k_top_p_min_p and Sampler.sample_device (engine/serve/sampler.py); beam_search (engine/serve/beam.py); utf8_sequences, TokenIndex.next_states, Guide.advance and json_schema_to_regex (engine/serve/structured.py).

Time: 7-9 hours. GPU: not needed.

What a request can ask for

Chapter 8’s sample chose one token from one row of logits with four settings. An API request can set many more, and in a serving engine every row of the batch belongs to a different request with different settings. SamplingParams collects them, with the names the OpenAI and vLLM APIs use:

@dataclass
class SamplingParams:
    max_tokens: int = 16
    temperature: float = 0.0          # 0 = greedy
    top_p: float = 1.0
    top_k: int = 0                    # 0 = off
    min_p: float = 0.0
    seed: int | None = None           # None: draw from the engine's generator
    n: int = 1                        # parallel samples of one prompt (Chapter 34)
    stop_token_ids: tuple = ()
    stop: tuple = ()                  # stop strings, checked on detokenized text (Chapter 35)
    include_stop_str_in_output: bool = False
    ignore_eos: bool = False
    min_tokens: int = 0               # EOS and stop tokens are masked until this many tokens exist
    presence_penalty: float = 0.0     # OpenAI: subtract once per distinct generated token
    frequency_penalty: float = 0.0    # OpenAI: subtract once per occurrence
    repetition_penalty: float = 1.0   # CTRL/HF: divide positive / multiply negative logits of seen tokens
    logit_bias: dict = field(default_factory=dict)   # {token_id: bias}
    allowed_token_ids: tuple | None = None
    logprobs: int | None = None       # return the chosen token's logprob and the top-n alternatives
    prompt_logprobs: int | None = None
    guided_regex: str | None = None   # constrained decoding (Chapter 34)
    guided_json: dict | str | None = None
    skip_special_tokens: bool = True

    def __post_init__(self):
        if self.max_tokens < 1:
            raise ValueError("max_tokens must be at least 1")
        if self.temperature < 0 or not 0 < self.top_p <= 1 or self.top_k < 0 or not 0 <= self.min_p <= 1:
            raise ValueError("Need temperature >= 0, 0 < top_p <= 1, top_k >= 0 and 0 <= min_p <= 1")
        if self.n < 1 or self.repetition_penalty <= 0 or self.min_tokens < 0:
            raise ValueError("Need n >= 1, repetition_penalty > 0 and min_tokens >= 0")
        if self.guided_regex is not None and self.guided_json is not None:
            raise ValueError("Choose one of guided_regex and guided_json")
        self.stop_token_ids = tuple(self.stop_token_ids)
        self.stop = (self.stop,) if isinstance(self.stop, str) else tuple(self.stop)

    def child(self, i):
        """Settings for sample i of an n > 1 request: one sample each, distinct seeds."""
        from dataclasses import replace
        return replace(self, n=1, seed=None if self.seed is None else self.seed + i)

    @property
    def greedy(self):
        return self.temperature == 0

    @property
    def needs_penalties(self):
        return bool(self.presence_penalty or self.frequency_penalty or self.repetition_penalty != 1.0)
settingwhat it doeswhere it acts
temperature, top_k, top_p, min_pshape and truncate the distribution (Chapter 8)filters
seedthis request’s draws are reproducible, whatever else is in the batchthe draw
presence_penalty, frequency_penaltysubtract a constant per distinct / per repeated output token (OpenAI)penalties
repetition_penaltyshrink the logits of every token seen in prompt or output (CTRL, Hugging Face)penalties
logit_bias, allowed_token_idsadd to chosen logits; forbid everything elseprocessors
min_tokensEOS and stop tokens are impossible until this many tokens existprocessors
guided_regex, guided_jsonthe output must match a pattern or a JSON Schemaprocessors
logprobs, prompt_logprobsreport log-probabilities of chosen, alternative and prompt tokensreported, not applied
nseveral independent samples of one promptthe engine
stop, include_stop_str_in_output, skip_special_tokensstop on generated text; how to detokenizethe detokenizer (Chapter 35)

Chapter 31’s SimpleSampler handled the first two rows of this table, with a Python loop over sampled rows. This chapter’s Sampler handles all of it with batched tensor operations, so its cost barely depends on how many requests asked for what.

One sampler for a batch

The pipeline, in the order vLLM uses:

logits [L, V] --> processors: allowed ids, logit bias, min_tokens, grammar masks
              --> penalties:  repetition, presence, frequency
              --> greedy rows: argmax                        (done)
              --> sampled rows: / temperature, top-k, top-p, min-p, draw

Order matters in two places. Processors come first because they express hard constraints: a grammar mask must not be undone by a penalty, and a forbidden token must stay forbidden whatever the temperature. Logprobs are taken from the raw logits, before anything else touches them. That’s the OpenAI convention and vLLM’s default: a reported logprob says what the model thought, so it’s comparable across requests with different settings, and evaluations that score answers by logprob aren’t disturbed by temperature or penalties.

Processors

def apply_processors(logits, requests):
    """Per-request logit edits that must happen before anything else.  (Your engine: Chapter 34)

    allowed_token_ids: everything else -> -inf. logit_bias: add bias[token]. min_tokens: EOS
    and stop tokens -> -inf until the request has that many output tokens. Grammar guides
    (structured.py): tokens that would break the pattern -> -inf.
    """
    vocab = logits.shape[-1]
    for i, request in enumerate(requests):
        p = request.params
        if p.allowed_token_ids is not None:
            keep = torch.zeros(vocab, dtype=torch.bool, device=logits.device)
            keep[list(p.allowed_token_ids)] = True
            logits[i].masked_fill_(~keep, float("-inf"))
        if p.logit_bias:
            ids = torch.tensor([int(t) for t in p.logit_bias], device=logits.device)
            logits[i, ids] += torch.tensor([float(b) for b in p.logit_bias.values()], device=logits.device)
        if request.num_output_tokens < p.min_tokens:
            banned = list(p.stop_token_ids) + ([request.eos_token_id] if request.eos_token_id is not None else [])
            if banned:
                logits[i, banned] = float("-inf")
        guide = request.extra.get("guide")
        if guide is not None:
            logits[i].masked_fill_(~guide.allowed(logits.device), float("-inf"))
    return logits

These are per-request edits with Python loops over the requests that set them. They’re cheap because most requests set none, and each edit is a single indexed write. min_tokens is worth a second look: setting EOS and every stop token to $-\infty$ makes stopping impossible rather than unlikely, so a model that wants to stop early must pick its next-best token. The engine’s check_stop also refuses to stop before min_tokens, which matters for stop strings (Chapter 35), which no logit mask can prevent.

Penalties

def apply_penalties(logits, requests):
    """Repetition, presence and frequency penalties, batched.  (Your engine: Chapter 34)

    counts[r, t] = occurrences of token t in request r's OUTPUT; seen[r, t] = t occurs in its
    prompt or output. Then, as in vLLM:
        repetition (HF/CTRL):  seen positive logits / p, seen negative logits * p
        frequency (OpenAI):    logits -= frequency_penalty * counts
        presence (OpenAI):     logits -= presence_penalty * (counts > 0)
    """
    rows = [i for i, r in enumerate(requests) if r.params.needs_penalties]
    if not rows:
        return logits
    device, vocab = logits.device, logits.shape[-1]
    sub = [requests[i] for i in rows]
    width = max(max(r.num_output_tokens for r in sub), 1)
    outputs = torch.full((len(sub), width), vocab, dtype=torch.long, device=device)      # vocab = padding
    for j, r in enumerate(sub):
        if r.num_output_tokens:
            outputs[j, :r.num_output_tokens] = torch.tensor(r.output_token_ids, device=device)
    counts = torch.zeros((len(sub), vocab + 1), device=device).scatter_add_(
        1, outputs, torch.ones_like(outputs, dtype=torch.float))[:, :vocab]
    seen = counts > 0
    for j, r in enumerate(sub):
        seen[j, torch.tensor(r.prompt_token_ids, device=device)] = True
    rep = torch.tensor([r.params.repetition_penalty for r in sub], device=device)[:, None]
    freq = torch.tensor([r.params.frequency_penalty for r in sub], device=device)[:, None]
    pres = torch.tensor([r.params.presence_penalty for r in sub], device=device)[:, None]
    x = logits[rows]
    x = torch.where(seen, torch.where(x > 0, x / rep, x * rep), x)
    x = x - freq * counts - pres * (counts > 0).float()
    logits[rows] = x
    return logits

The three penalties do different things and are easy to confuse:

  • Presence ($\alpha$) and frequency ($\beta$) come from the OpenAI API: $\ell_t \leftarrow \ell_t - \alpha \cdot [c_t > 0] - \beta \cdot c_t$, where $c_t$ counts token $t$ in the output so far. They’re additive and unbounded: a frequency penalty of 0.5 makes a token that already appeared 20 times 10 logits less likely.
  • Repetition ($\rho$), from CTRL (Keskar et al., 2019) and Hugging Face, is multiplicative and applies to every token seen in the prompt or output: positive logits are divided by $\rho$, negative ones multiplied. It’s bounded, but it penalizes copying from the prompt, which hurts tasks that should quote it, like extraction and summarization.

The counts come from one scatter_add_ over the batch’s output tokens padded into a rectangle (the padding value is vocab, a column that’s dropped afterwards). Rebuilding them from Python lists every step costs host time proportional to the outputs’ total length; vLLM keeps persistent count tensors and updates them with each step’s tokens (stretch exercise 1).

Filters and the draw

def top_k_top_p_min_p(logits, top_k, top_p, min_p):
    """Per-row filters with one descending sort.  (Your engine: Chapter 34)

    top_k [L] (vocab = off), top_p [L] (1 = off), min_p [L] (0 = off), applied in Chapter 8's
    order: keep the k best; renormalize; keep a token while the mass BEFORE it is < p (so the
    token that crosses p stays); drop tokens below min_p times the top probability (a ratio,
    so it doesn't care about renormalization). The top token always survives.
    """
    ordered, order = logits.sort(dim=-1, descending=True)
    rank = torch.arange(logits.shape[-1], device=logits.device)[None]
    ordered = ordered.masked_fill(rank >= top_k[:, None], float("-inf"))
    probs = ordered.softmax(-1)
    keep = (probs.cumsum(-1) - probs) < top_p[:, None]
    keep &= probs >= min_p[:, None] * probs[:, :1]
    keep[:, 0] = True
    filtered = ordered.masked_fill(~keep, float("-inf"))
    return torch.full_like(logits, float("-inf")).scatter(-1, order, filtered)

One descending sort serves all three filters for every row. Each row’s top_k becomes a rank mask, its top_p a cumulative-mass mask on the renormalized survivors, and its min_p a ratio to the top probability. A row that asked for none of them has top_k = vocab, top_p = 1, min_p = 0, and keeps everything. The sort is skipped entirely when no row needs it.

class Sampler:
    def __init__(self, max_logprobs=20):
        self.max_logprobs = max_logprobs

    def sample_device(self, logits, requests, generators):
        """[L, V] logits -> [L] token ids on the device, no host sync.  (Your engine: Chapter 34)"""
        logits = apply_penalties(apply_processors(logits.float().clone(), requests), requests)
        device, vocab = logits.device, logits.shape[-1]
        greedy = torch.tensor([r.params.greedy for r in requests], device=device)
        temperature = torch.tensor([r.params.temperature or 1.0 for r in requests], device=device)
        x = logits / temperature[:, None]
        if any(r.params.top_k or r.params.top_p < 1 or r.params.min_p for r in requests):
            top_k = torch.tensor([r.params.top_k or vocab for r in requests], device=device)
            top_p = torch.tensor([r.params.top_p for r in requests], device=device)
            min_p = torch.tensor([r.params.min_p for r in requests], device=device)
            x = top_k_top_p_min_p(x, top_k, top_p, min_p)
        probs = x.softmax(-1)
        # Exponential race: argmax(p_i / E_i) with E_i ~ Exp(1) draws i with probability p_i.
        noise = torch.empty_like(probs).exponential_()
        for i, r in enumerate(requests):
            if r.request_id in generators:          # seeded: the same draws whatever the batch
                noise[i] = torch.empty(vocab, device=device).exponential_(generator=generators[r.request_id])
        sampled = (probs / noise).argmax(-1)
        self.last_logits = logits
        return torch.where(greedy, logits.argmax(-1), sampled)

    def __call__(self, logits, requests, generators):
        tokens = self.sample_device(logits, requests, generators)
        token_list = tokens.tolist()
        logprobs = None
        wanted = [r.params.logprobs for r in requests]
        if any(n is not None for n in wanted):
            logprobs = self.logprobs(logits, tokens, wanted)
        for request, token in zip(requests, token_list):
            guide = request.extra.get("guide")
            if guide is not None:
                guide.advance(token)
        return [[t] for t in token_list], logprobs

    def logprobs(self, raw_logits, tokens, wanted):
        """Chosen token's logprob and rank, and the top-n alternatives, from the raw logits."""
        lp = raw_logits.float().log_softmax(-1)
        chosen = lp.gather(-1, tokens[:, None])[:, 0]
        rank = (lp > chosen[:, None]).sum(-1) + 1
        n = min(max(w or 0 for w in wanted), self.max_logprobs)
        top_v, top_i = lp.topk(max(n, 1), dim=-1)
        out = []
        for i, w in enumerate(wanted):
            if w is None:
                out.append(None)
                continue
            out.append([{"token_id": int(tokens[i]), "logprob": float(chosen[i]), "rank": int(rank[i]),
                         "top": dict(zip(top_i[i, :w].tolist(), top_v[i, :w].tolist()))}])
        return out

The draw uses the exponential race: with $E_i \sim \text{Exp}(1)$ independent, $\arg\max_i p_i / E_i$ is token $i$ with probability exactly $p_i$. It’s the Gumbel-max trick of Chapter 19 in another form ($-\log E_i$ is a Gumbel variable), so it needs no sort and no cumulative sum, only elementwise operations and an argmax, and it never syncs with the host.

Seeded requests draw their noise from their own torch.Generator. That’s what makes seed meaningful in a server: the row’s random numbers depend only on its generator, not on its position in the batch or on how many other rows drew before it. The test checks exactly that: the same seeded request draws the same tokens alone and in a batch.

Logprobs, and why prompt logprobs disable the prefix cache

logprobs=n returns, for each generated token, its log-probability, its rank, and the $n$ most likely alternatives. prompt_logprobs=n returns the log-probability of each prompt token given the ones before it, which is how evaluation harnesses score multiple-choice answers and compute perplexity through an API.

Prompt logprobs need the logits of every prompt position, not just the last. build_batch therefore adds, for such requests, the rows whose next token is a prompt token after the sampling rows, and the engine turns them into log-probabilities as they’re computed, chunk by chunk:

def build_batch(scheduled, block_tables, block_size, device="cpu", pad_to=None):
    """Lay out [(request, n), ...] as one flattened batch.  (Your engine: Chapter 31)

    block_tables[request_id] lists the request's physical blocks in logical order. Request r's
    n new tokens are token_ids[c : c + n] (plus its draft tokens, if any), at positions
    c .. c + n - 1, where c is its num_computed_tokens. Position p lives in slot
    table[p // block_size] * block_size + p % block_size.

    pad_to (CUDA graphs, Chapter 33) appends dummy decode rows with slot -1 and length 0.
    """
    ids, positions, slots, starts, lengths, tables, logits_rows, counts = [], [], [], [0], [], [], [], []
    prompt_rows, prompt_spans = [], []
    for request, n in scheduled:
        c = request.num_computed_tokens
        tokens = (request.token_ids + request.spec_token_ids)[c:c + n]
        if len(tokens) != n:
            raise ValueError(f"{request.request_id}: scheduled {n} tokens but only {len(tokens)} exist")
        table = block_tables[request.request_id]
        ids.extend(tokens)
        positions.extend(range(c, c + n))
        slots.extend(table[p // block_size] * block_size + p % block_size for p in range(c, c + n))
        starts.append(starts[-1] + n)
        lengths.append(c + n)
        tables.append(table)
        # Rows from the last real token onwards produce samples: one for a finished prefill or a
        # decode, 1 + (draft tokens scheduled) when verifying. A mid-prompt chunk produces none.
        k = min(n, max(0, c + n - (request.num_tokens - 1)))
        logits_rows.extend(range(starts[-1] - k, starts[-1]))
        counts.append(k)
        if request.params.prompt_logprobs is not None:      # rows whose next token is a prompt token
            last = min(c + n, request.num_prompt_tokens - 1)
            if last > c:
                prompt_rows.extend(range(starts[-2], starts[-2] + last - c))
                prompt_spans.append((request, c, request.token_ids[c + 1:last + 1]))
    if pad_to is not None:
        for _ in range(pad_to - len(lengths)):
            ids.append(0), positions.append(0), slots.append(-1)
            starts.append(starts[-1] + 1)
            lengths.append(0)
            tables.append([])
    width = max(1, max(len(t) for t in tables))
    table_tensor = torch.zeros((len(tables), width), dtype=torch.int32)
    for r, table in enumerate(tables):
        table_tensor[r, :len(table)] = torch.tensor(table, dtype=torch.int32)
    meta = BatchMeta(
        query_start_loc=torch.tensor(starts, dtype=torch.int32, device=device),
        seq_lens=torch.tensor(lengths, dtype=torch.int32, device=device),
        block_table=table_tensor.to(device),
        slot_mapping=torch.tensor(slots, dtype=torch.int64, device=device),
        block_size=block_size, query_start_loc_cpu=starts, seq_lens_cpu=lengths,
        max_query_len=max(b - a for a, b in zip(starts, starts[1:])), max_seq_len=max(lengths))
    return Batch(torch.tensor(ids, dtype=torch.long, device=device),
                 torch.tensor(positions, dtype=torch.long, device=device), meta,
                 torch.tensor(logits_rows + prompt_rows, dtype=torch.long, device=device), counts,
                 [r.request_id for r, _ in scheduled], {}, prompt_spans)

A cached prefix is never computed, so its logits don’t exist. The block manager therefore gives no prefix-cache hits to requests that ask for prompt logprobs, the same rule vLLM uses. Such requests still publish their blocks, so they help later requests.

Parallel sampling (n=4: four independent answers to one prompt) looks like it needs forked caches, as in Chapter 25. With the prefix cache, it doesn’t: the engine turns it into four ordinary requests, rid:0 to rid:3, with seeds seed + i. They’re admitted in the same step, and since the scheduler publishes blocks as soon as it allocates them (Chapter 31), children 1-3 adopt child 0’s prompt blocks immediately. Only the last, partial block is computed four times. The frontend (Chapter 36) gathers the children back into one response.

Beam search keeps the $W$ most probable sequences, extending each by its most probable next tokens. It’s deterministic and favors high-likelihood outputs, which suits translation and some structured tasks, though it’s known to produce bland, repetitive text in open-ended generation (Holtzman et al., 2020). Inside an engine it’s awkward: beams fork and die every step, and the scheduler would need to know about them. vLLM V1 moved beam search out of the engine, and so does this book:

def beam_search(engine, prompt_ids, beam_width, max_tokens, eos_token_id=None, length_penalty=1.0):
    """Returns [(tokens, score)] for the best beam_width sequences, best first.  (Your engine: Chapter 34)

    score = sum of token logprobs / (generated length ** length_penalty). A beam that emits EOS
    is finished and stops growing; the search ends when W beams have finished or max_tokens
    is reached.
    """
    beams, finished = [([], 0.0)], []
    params = SamplingParams(max_tokens=1, logprobs=2 * beam_width, ignore_eos=True)
    for step in range(max_tokens):
        outputs = []
        prompts = {f"beam-{step}-{i}": list(prompt_ids) + tokens for i, (tokens, _) in enumerate(beams)}
        engine.generate(prompts, params, outputs=outputs)
        top = {o.request_id: o.logprobs[0]["top"] for o in outputs}
        candidates = []
        for i, (tokens, logprob) in enumerate(beams):
            for token, lp in top[f"beam-{step}-{i}"].items():
                candidates.append((tokens + [token], logprob + lp))
        candidates.sort(key=lambda c: -c[1])
        beams = []
        for tokens, logprob in candidates:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, logprob))
            else:
                beams.append((tokens, logprob))
            if len(beams) == beam_width:
                break
        if len(finished) >= beam_width or not beams:
            break
    pool = finished + beams
    scored = [(tokens, logprob / len(tokens) ** length_penalty) for tokens, logprob in pool]
    return sorted(scored, key=lambda c: -c[1])[:beam_width]

Each step submits one single-token request per beam, asking for the top $2W$ logprobs (enough that $W$ survive even if some beams end). A beam’s request shares everything but its newest token with its parent’s request of the previous step, so the prefix cache recomputes about one block per beam per step. The cost is one engine round-trip per generated token rather than a fused loop, which is the right trade-off for a feature that’s rarely used in serving.

Constrained decoding

A model asked for JSON usually produces JSON. A program that parses the output needs always: one missing quote in a million requests is a production incident. Constrained decoding guarantees it by masking, at every step, every token that would make the output impossible to complete validly. The model still chooses, among the tokens that keep the output valid.

The pieces: compile the pattern to an automaton, and at each step allow exactly the tokens that keep the automaton alive.

Why the automaton runs over bytes

Tokens are byte strings. In a byte-level BPE vocabulary (Chapter 4), é (UTF-8 C3 A9) may be one token, or two tokens C3 and A9, and many tokens end in the middle of a multi-byte character, especially for Chinese, Japanese and emoji. A character-level automaton can’t say whether the token C3 is allowed. A byte-level automaton can: after C3 it’s in a state that expects one continuation byte 80-BF.

So the regex compiler turns every character class into byte sequences. A range of code points becomes a few sequences of byte ranges:

U+0000-U+007F   [00-7F]
U+0080-U+07FF   [C2-DF][80-BF]
U+0800-U+0FFF   [E0][A0-BF][80-BF]
U+1000-U+CFFF   [E1-EC][80-BF][80-BF]
...
def utf8_sequences(lo, hi):
    """Byte-range sequences whose concatenations are exactly the UTF-8 encodings of the code
    points lo..hi.  (Your engine: Chapter 34)

    Example: U+0080..U+07FF is [C2-DF][80-BF]. A range is first split where the encoded length
    changes (and around the surrogates, which have no encoding), then wherever a prefix
    byte doesn't cover a whole block of continuation bytes, until each piece is a product of
    byte ranges.
    """
    for a, b in ((0, 0x7F), (0x80, 0x7FF), (0x800, 0xD7FF), (0xE000, 0xFFFF), (0x10000, MAX_CODEPOINT)):
        if lo <= b and hi >= a:
            yield from _split(max(lo, a), min(hi, b))


def _split(lo, hi):
    n = len(chr(lo).encode())
    for i in range(1, n):
        mask = (1 << (6 * i)) - 1                         # the low i continuation bytes
        if lo & ~mask != hi & ~mask:
            if lo & mask:
                yield from _split(lo, lo | mask)
                yield from _split((lo | mask) + 1, hi)
                return
            if hi & mask != mask:
                yield from _split(lo, (hi & ~mask) - 1)
                yield from _split(hi & ~mask, hi)
                return
    yield list(zip(chr(lo).encode(), chr(hi).encode()))

The splitting rule is the one from RE2 and Rust’s regex-syntax: a range is cut where the encoded length changes and around the surrogates (U+D800-U+DFFF, which UTF-8 can’t encode), then wherever a prefix byte doesn’t cover a whole block of continuation bytes, until each piece is a product of byte ranges. The test checks it by brute force: the byte strings a set of sequences generates must be exactly the encodings of the code points in the range.

From regex to DFA

The parser is ordinary recursive descent: alternation of sequences of quantified atoms. Its output is a small tree:

class Parser:
    """Recursive descent over the regex text. AST nodes are tuples:
    ("chars", [(lo, hi), ...])  one code point from these ranges
    ("cat", [nodes]), ("alt", [nodes]), ("repeat", node, min, max or None)
    """

    CLASSES = {"d": [(48, 57)], "w": [(48, 57), (65, 90), (95, 95), (97, 122)],
               "s": [(9, 13), (32, 32)]}
    ESCAPES = {"n": 10, "t": 9, "r": 13, "f": 12, "v": 11, "0": 0}

    def __init__(self, text):
        self.text, self.i = text, 0

    def parse(self):
        node = self.alternation()
        if self.i != len(self.text):
            raise ValueError(f"Unexpected {self.text[self.i]!r} at {self.i} in {self.text!r}")
        return node

    def peek(self):
        return self.text[self.i] if self.i < len(self.text) else None

    def take(self):
        ch = self.text[self.i]
        self.i += 1
        return ch

    def alternation(self):
        branches = [self.sequence()]
        while self.peek() == "|":
            self.take()
            branches.append(self.sequence())
        return branches[0] if len(branches) == 1 else ("alt", branches)

    def sequence(self):
        items = []
        while self.peek() not in (None, "|", ")"):
            items.append(self.quantified())
        return ("cat", items)

    def quantified(self):
        node = self.atom()
        while self.peek() in ("*", "+", "?", "{"):
            ch = self.peek()
            if ch == "{":
                match = re.match(r"\{(\d+)(,(\d*))?\}", self.text[self.i:])
                if not match:
                    break                                  # a literal brace
                self.i += match.end()
                lo = int(match.group(1))
                hi = lo if match.group(2) is None else (int(match.group(3)) if match.group(3) else None)
                node = ("repeat", node, lo, hi)
                continue
            self.take()
            node = ("repeat", node, {"*": 0, "+": 1, "?": 0}[ch], {"*": None, "+": None, "?": 1}[ch])
        return node

    def atom(self):
        ch = self.take()
        if ch == "(":
            if self.text.startswith("?:", self.i):
                self.i += 2
            node = self.alternation()
            if self.peek() != ")":
                raise ValueError(f"Unclosed group in {self.text!r}")
            self.take()
            return node
        if ch == "[":
            return ("chars", self.char_class())
        if ch == ".":
            return ("chars", [(0, 9), (11, MAX_CODEPOINT)])
        if ch == "\\":
            return ("chars", self.escape())
        if ch in "*+?":
            raise ValueError(f"Nothing to repeat at {self.i - 1} in {self.text!r}")
        return ("chars", [(ord(ch), ord(ch))])

    def escape(self):
        ch = self.take()
        if ch in self.CLASSES:
            return self.CLASSES[ch]
        if ch.lower() in self.CLASSES:
            return complement(self.CLASSES[ch.lower()])
        if ch == "u":
            code = int(self.text[self.i:self.i + 4], 16)
            self.i += 4
            return [(code, code)]
        if ch == "x":
            code = int(self.text[self.i:self.i + 2], 16)
            self.i += 2
            return [(code, code)]
        code = self.ESCAPES.get(ch, ord(ch))
        return [(code, code)]

    def char_class(self):
        negate = self.peek() == "^"
        if negate:
            self.take()
        ranges, first = [], True
        while first or self.peek() != "]":
            first = False
            if self.peek() is None:
                raise ValueError(f"Unclosed class in {self.text!r}")
            ch = self.take()
            lo = self.escape() if ch == "\\" else [(ord(ch), ord(ch))]
            if len(lo) == 1 and lo[0][0] == lo[0][1] and self.peek() == "-" and self.text[self.i + 1:self.i + 2] not in ("]", ""):
                self.take()
                ch = self.take()
                hi = self.escape() if ch == "\\" else [(ord(ch), ord(ch))]
                ranges.append((lo[0][0], hi[0][0]))
            else:
                ranges.extend(lo)
        self.take()
        return complement(ranges) if negate else normalize(ranges)

Thompson’s construction turns the tree into an NFA with byte-range edges and epsilon edges, one small fragment per node; {m,n} becomes $m$ mandatory copies and $n - m$ optional ones. Subset construction then turns the NFA into a DFA whose states are sets of NFA states, with a full 256-entry transition row per state:

class DFA:
    """table[s, byte] = next state; state 0 is dead (absorbing), state 1 is the start."""

    def __init__(self, pattern, max_states=20000):
        nfa = NFA()
        start, end = nfa.state(), nfa.state()
        nfa.build(Parser(pattern).parse(), start, end)
        first = nfa.closure([start])
        ids = {first: 1}
        rows, accepting, todo = [[0] * 256, [0] * 256], [False, end in first], [first]
        while todo:                                       # subset construction
            current = todo.pop()
            row = rows[ids[current]]
            moves = {}
            for s in current:
                for a, b, nxt in nfa.edges[s]:
                    for byte in range(a, b + 1):
                        moves.setdefault(byte, set()).add(nxt)
            for byte, targets in moves.items():
                target = nfa.closure(targets)
                if target not in ids:
                    if len(ids) + 1 > max_states:
                        raise ValueError("Pattern needs too many DFA states; simplify it")
                    ids[target] = len(rows)
                    rows.append([0] * 256)
                    accepting.append(end in target)
                    todo.append(target)
                row[byte] = ids[target]
        self.table = torch.tensor(rows, dtype=torch.int32)
        self.accepting = torch.tensor(accepting)
        # A state is "final" when it accepts and no byte leads anywhere: only EOS may follow.
        self.final = self.accepting & (self.table == 0).all(-1)

    def match(self, data):
        state = 1
        for byte in data:
            state = int(self.table[state, byte])
        return bool(self.accepting[state])

State 0 is dead: every byte leads back to it, and any token that reaches it is forbidden. A state is final when it accepts and no byte leads anywhere; the output is complete, and only EOS may follow. Russ Cox’s Regular Expression Matching Can Be Simple And Fast (2007) explains why this construction is linear in the input, unlike the backtracking matchers in most languages’ standard libraries.

From DFA to token masks

For a DFA state $s$ and a vocabulary of $V$ tokens, the mask says which tokens’ bytes, fed one by one from $s$, never reach the dead state. Walking 150,000 tokens one at a time in Python would take a second per state. Walking them all at once takes milliseconds: put every token’s bytes in a padded [V, max_len] tensor, start every token at $s$, and advance all of them one byte position at a time with a single table lookup:

class TokenIndex:
    """Which tokens each DFA state allows, computed by walking ALL tokens' bytes at once."""

    def __init__(self, dfa, vocab_bytes, eos_token_id, vocab_size=None):
        self.dfa, self.eos = dfa, eos_token_id
        self.vocab_size = vocab_size or len(vocab_bytes)
        width = max(1, max(len(b) for b in vocab_bytes))
        self.lengths = torch.tensor([len(b) for b in vocab_bytes])
        self.bytes = torch.zeros((len(vocab_bytes), width), dtype=torch.long)
        for i, b in enumerate(vocab_bytes):
            if b:
                self.bytes[i, :len(b)] = torch.tensor(list(b))
        self.cache = {}

    def next_states(self, state):
        """[vocab] next DFA state after each token from `state` (0 = forbidden).  (Your engine: Chapter 34)"""
        if state not in self.cache:
            current = torch.full((self.bytes.shape[0],), state, dtype=torch.int32)
            for position in range(self.bytes.shape[1]):
                step = self.dfa.table[current.long(), self.bytes[:, position]]
                current = torch.where(position < self.lengths, step, current)
            current[self.lengths == 0] = 0                        # empty tokens (specials) never advance a pattern
            self.cache[state] = current
        return self.cache[state]

    def allowed(self, state, device="cpu"):
        """[vocab_size] bool: tokens that keep the pattern alive, plus EOS in accepting states."""
        mask = torch.zeros(self.vocab_size, dtype=torch.bool)
        nxt = self.next_states(state)
        mask[:nxt.shape[0]] = nxt != 0
        if self.eos is not None:
            mask[self.eos] = bool(self.dfa.accepting[state])
        return mask.to(device)

The result also gives each token’s next state, so advancing a request after sampling is one lookup. Results are cached per state, and states are computed lazily, only when some request reaches them. A pattern with 440 DFA states rarely visits more than a few dozen. EOS is allowed exactly in accepting states.

class Guide:
    """One request's position in its pattern."""

    def __init__(self, index):
        self.index, self.state = index, 1

    def allowed(self, device="cpu"):
        return self.index.allowed(self.state, device)

    def advance(self, token):
        """(Your engine: Chapter 34)"""
        if token == self.index.eos:
            self.state = -1                                       # done
            return
        nxt = int(self.index.next_states(self.state)[token])
        if nxt == 0:
            raise ValueError(f"Token {token} violates the pattern")
        self.state = nxt

    @property
    def finished(self):
        """No byte can follow: the output is complete even without an EOS token."""
        return self.state == -1 or bool(self.index.dfa.final[self.state])

Each request holds a Guide with its current state. The sampler masks with guide.allowed() before anything else and calls guide.advance(token) after the draw; the engine stops the request when the guide is final, even without an EOS token.

JSON Schema to regex

Most structured output in practice is “JSON matching this schema”, from OpenAI’s response_format and from tool calling (Chapter 35). A practical subset of JSON Schema is regular, so it can be translated into a regex, the approach of Outlines (Willard and Louf, 2023):

def json_schema_to_regex(schema, defs=None, depth=0):
    """A regex whose matches are exactly the (compact, optionally single-spaced) JSON documents
    valid under a practical subset of JSON Schema.  (Your engine: Chapter 34)

    Supported: type (one or a list), properties + required (in declared order), items with
    minItems/maxItems, enum, const, anyOf/oneOf, $ref into $defs, string minLength/maxLength/
    pattern, and the primitives. Objects are closed (no extra keys).
    """
    if depth > 16:
        raise ValueError("Schema nests too deeply (recursive $ref?)")
    defs = defs if defs is not None else schema.get("$defs", schema.get("definitions", {}))
    recurse = lambda s: json_schema_to_regex(s, defs, depth + 1)                     # noqa: E731
    if "$ref" in schema:
        return recurse(defs[schema["$ref"].split("/")[-1]])
    if "const" in schema:
        return regex_escape(json.dumps(schema["const"]))
    if "enum" in schema:
        return "(?:" + "|".join(regex_escape(json.dumps(v)) for v in schema["enum"]) + ")"
    for key in ("anyOf", "oneOf"):
        if key in schema:
            return "(?:" + "|".join(recurse(s) for s in schema[key]) + ")"
    kind = schema.get("type")
    if isinstance(kind, list):
        return "(?:" + "|".join(recurse({**schema, "type": k}) for k in kind) + ")"
    if kind == "object" or (kind is None and "properties" in schema):
        props = schema.get("properties", {})
        required = set(schema.get("required", []))
        items = [(f'"{regex_escape(name)}"{WS}:{WS}{recurse(sub)}', name in required) for name, sub in props.items()]
        if not items:
            return r"\{" + WS + r"\}"
        branches = []
        for first, (pattern, _) in enumerate(items):          # which property comes first decides the commas
            if any(req for _, req in items[:first]):
                break                                         # a required property can't be skipped
            body = pattern
            for later, req in items[first + 1:]:
                part = f"{WS},{WS}{later}"
                body += part if req else f"(?:{part})?"
            branches.append(body)
        inner = "(?:" + "|".join(branches) + ")"
        if not required:
            inner += "?"
        return r"\{" + WS + inner + WS + r"\}"
    if kind == "array":
        item = recurse(schema.get("items", {"type": "string"}))
        lo, hi = schema.get("minItems", 0), schema.get("maxItems")
        more = f"(?:{WS},{WS}{item})"
        if hi is not None and hi == 0:
            body = ""
        elif lo == 0:
            body = f"(?:{item}{more}{{0,{hi - 1}}})?" if hi is not None else f"(?:{item}{more}*)?"
        else:
            body = item + (more + (f"{{{lo - 1},{hi - 1}}}" if hi is not None else f"{{{lo - 1},}}"))
        return r"\[" + WS + body + WS + r"\]"
    if kind == "string":
        if "pattern" in schema:
            return f'"{schema["pattern"].lstrip("^").rstrip("$")}"'
        lo, hi = schema.get("minLength"), schema.get("maxLength")
        if lo is not None or hi is not None:
            return f'"{STRING_CHAR}{{{lo or 0},{"" if hi is None else hi}}}"'
        return PRIMITIVES["string"]
    if kind in PRIMITIVES:
        return PRIMITIVES[kind]
    if kind is None:                                          # {} = any primitive (bounded: no nesting)
        return "(?:" + "|".join(PRIMITIVES.values()) + ")"
    raise ValueError(f"Unsupported schema type {kind!r}")

The only subtle part is objects with optional properties: a comma separates two properties only if both are present. The translation chooses which property comes first (the alternation), after which every later property is , "key": value, mandatory if required and optional otherwise. Properties stay in declared order, which keeps the regex, and the DFA, small. Objects are closed: no keys the schema doesn’t list.

What doesn’t fit: recursive schemas (a tree whose nodes contain trees) and arbitrary JSON ({"type": "json_object"} with unbounded nesting) aren’t regular, because a finite automaton can’t count brackets. They need a pushdown automaton, or a context-free grammar engine such as XGrammar (Dong et al., 2024) or llguidance, which keep a stack, precompute masks for the tokens whose validity doesn’t depend on it, and check the rest at runtime. The interface stays the same: a mask per step and an advance per token (stretch exercise 3).

What constrained decoding costs

Three costs, all addressed by production engines:

  1. Compile time: building the DFA and its first masks. Cache compiled patterns (compile_pattern and the factory’s indexes do); tool schemas repeat across requests.
  2. Mask time per step: a few milliseconds for a new state, under a millisecond for a cached one, on the CPU. Engines compute the next step’s masks while the GPU runs the forward pass, and apply them as bitmasks (32 tokens per int32).
  3. Synchronous scheduling: the next mask depends on the token just sampled, so a guided request can’t use Chapter 33’s async scheduling, whose next step is planned before the token is known. This engine refuses the combination; production engines advance the grammar on the device or delay the mask by a step.

SGLang adds a speedup worth knowing: when the DFA has only one path forward (inside a fixed key like "name": "), the engine can append all its tokens at once without asking the model, called jump-forward decoding (stretch exercise 4).

Run it

python run.py guided
{"regex_chars": 295, "dfa_states": 440, "compile_ms": 22.4}
{"state": "start", "allowed_tokens": 195, "first_mask_ms": 17.3, "cached_mask_ms": 0.67}
{"state": "after '{\"name\": \"'", "allowed_tokens": 141111, "first_mask_ms": 15.9, "cached_mask_ms": 0.45}
{"seed": 0, "output": "{\"name\" : \"age\",\"age\": -243,\"admin\":true }", "valid_json": true}
{"seed": 1, "output": "{\"name\":\"name\" ,\"age\":-6638951103811364110 }", "valid_json": true}
{"seed": 2, "output": "{\"name\" : \"c++\", \"age\" :3 , \"admin\":true}", "valid_json": true}
{"greedy": [244, 158, 145, 235, 484, 225], "beam_best": [244, 158, 145, 235, 484, 225], "beam_best_mean_logprob": -4.52}

A schema with four properties, an enum array and optional fields compiles to a 440-state DFA in 22 ms. With a synthetic vocabulary the size of Qwen3’s (151,936 byte strings), the first mask for a state takes about 17 ms and a cached one under a millisecond. At the start, only the 195 tokens that begin with { and continue validly are allowed; inside a string value, 141,111 are. Then the random 2-layer test model, which knows nothing about JSON, produces valid, schema-conforming JSON with three different seeds. (Its choices are silly; the guarantee isn’t.) Finally, beam search with four beams finds the greedy sequence here, as it often does when one continuation dominates.

Build it

Engine milestone 34: the sampler and structured output. Implement apply_processors, apply_penalties, top_k_top_p_min_p and Sampler.sample_device in engine/serve/sampler.py; beam_search in engine/serve/beam.py; and utf8_sequences, TokenIndex.next_states, Guide.advance and json_schema_to_regex in engine/serve/structured.py (the regex parser, NFA, DFA, logprob formatting and the engine wiring are provided). Use it with EngineConfig(sampler="full") and EngineCore(..., vocab_bytes=...).

pytest tests/test_ch34_sampling.py
python run.py guided --impl engine

The tests check UTF-8 sequences by brute force, the DFA against Python’s re on ten patterns and twenty-five strings, token masks against a byte-by-byte walk (including half of a two-byte character), EOS only in accepting states, the JSON Schema translation on valid and invalid documents with required, optional and $ref properties, each processor and penalty against its formula, the filters against Chapter 8’s order, the sampling distribution and seed independence from the batch, logprobs and prompt logprobs against a full forward pass, n=4 sharing the prompt’s blocks, beam search against a brute-force implementation, and guided generation producing pattern- and schema-valid output.

Stretch exercises

  1. ★ Keep penalty counts in a persistent [max_num_seqs, vocab] tensor indexed by request slot, updated with each step’s sampled tokens. Measure host time per step against apply_penalties at 256 requests with 1,000-token outputs. Where: penalty state in Sampler in engine/serve/sampler.py, with request-slot lifecycle in engine/serve/engine.py.
  2. ★★ Pack masks as int32 bitmasks (32 tokens per word) and apply them with a small Triton kernel that expands bits to $-\infty$. Compare memory and time against boolean masks for 256 guided requests. Where: mask storage in engine/serve/structured.py; new engine/kernels/triton_masks.py, called from engine/serve/sampler.py.
  3. ★★★ Support arbitrary JSON ({"type": "json_object"}) with a pushdown guide: a byte-level automaton for tokens plus a stack of open { and [. Precompute, per automaton state, the tokens that don’t touch the stack, and check the rest at runtime (the core idea of XGrammar). Where: guide state and GuideFactory in engine/serve/structured.py.
  4. ★★ Implement jump-forward decoding: when the guide’s state has exactly one path for the next several bytes, tokenize that text and append its tokens directly, skipping the model for those positions. Measure steps saved on the schema of run.py guided. Where: deterministic-path detection in engine/serve/structured.py, with token/cache progression in engine/serve/engine.py.

Check your understanding

  1. Why are logprobs reported from the raw logits rather than after temperature and filters?
  2. What’s the difference between repetition_penalty and frequency_penalty, and which one would you avoid for summarization?
  3. Why does the exponential race draw token $i$ with probability $p_i$? Why is it graph-friendly?
  4. Why does a seeded request use its own generator, and what would go wrong with one shared generator?
  5. Why can’t a request with prompt_logprobs use the prefix cache?
  6. How does n=4 share the prompt’s KV cache without forking, and which block is computed four times?
  7. Why must the constraint automaton run over bytes rather than characters?
  8. Why can a regex describe “an object with these properties” but not “any JSON value”?

Going deeper

  • Willard and Louf, Efficient Guided Generation for Large Language Models (2023), the Outlines paper; Dong et al., XGrammar: Flexible and Efficient Structured Generation Engine for Large Language Models (2024); Microsoft’s llguidance; LMSYS, Fast JSON Decoding for Local LLMs with Compressed Finite State Machine (2024), for jump-forward decoding.
  • Russ Cox, Regular Expression Matching Can Be Simple And Fast (2007), and the utf8 module of Rust’s regex-syntax crate, for Thompson NFAs and UTF-8 range compilation.
  • Holtzman et al., The Curious Case of Neural Text Degeneration (ICLR 2020), for nucleus sampling and beam search’s failure modes; Nguyen et al., Turning Up the Heat: Min-p Sampling for Creative and Coherent LLM Outputs (2024); Keskar et al., CTRL (2019), for the repetition penalty.
  • vLLM’s vllm/v1/sample/ (sampler, penalties, logprobs) and vllm/v1/structured_output/, which follow the structure of this chapter.