Autoregressive Generation: How GPT Produces Text

Michael BrenndoerferUpdated July 30, 202577 min read

Part of Language AI Handbook

Covers mechanics of autoregressive generation in transformers, including the generation loop, KV caching for efficiency, stopping criteria.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Autoregressive Generation

When you ask a language model to complete the sentence "The cat sat on the", it doesn't produce the entire response at once. Instead, it generates text one token at a time, each new token conditioned on all the tokens that came before. This process, called autoregressive generation, is the basic mechanism by which decoder-only transformers like GPT produce coherent, contextually appropriate text. The word "autoregressive" comes from time-series modeling: a process is autoregressive when each new value in a sequence depends on the values that came before it, feeding back into itself. In language modeling, every new token is conditioned on the entire history of previously generated tokens, making the model literally "regress" on its own prior outputs.

Understanding autoregressive generation is needed for working with modern language models. The process seems simple on the surface: predict the next token, append it to the sequence, and repeat. But underneath this simple loop lies a computational challenge that has driven significant research into efficiency optimizations. The key insight is that naive generation would recompute the same attention calculations repeatedly, wasting enormous amounts of compute. The solution, key-value caching, transforms generation from quadratically expensive to linearly efficient. This single optimization is what makes practical deployment of large language models possible: without it, generating a single paragraph would take minutes even on powerful hardware.

Think of autoregressive generation like a skilled improviser building a story out loud. Each word they speak constrains and informs the next. They can't go back and revise what they've already said; each new word must follow from the entire context established so far. The model faces exactly the same constraint: each token decision is permanent and becomes part of the context for all future decisions. This one-way, committed nature of generation is both a strength, because it keeps outputs coherent and contextually grounded, and a limitation, because errors early in a sequence can compound and distort everything that follows.

The mathematical foundation of autoregressive generation is the chain rule of probability. Any joint probability over a sequence of tokens can be exactly factored into a product of conditional probabilities: the probability of the first token, times the probability of the second token given the first, times the probability of the third given the first two, and so on. This factorization is not an approximation; it is mathematically exact. The practical question is how well a neural network can approximate each conditional probability. GPT-style transformers turn out to be extraordinarily effective at this task, capturing subtle long-range dependencies through their attention mechanism.

Beyond the mathematical elegance, there is a practical engineering story here. Early language models computed autoregressive generation straightforwardly, paying a quadratic cost in sequence length at every step. As contexts grew from dozens of tokens to thousands, this cost became prohibitive. The invention of key-value caching reduced that cost to linear, making the long-context capabilities of modern chat assistants, code generation tools, and document summarizers practical. Understanding this optimization is not just academic: it shapes deployment decisions, memory budgets, and the architectural choices made in models like GPT-4 and Llama.

This chapter explores the mechanics of autoregressive generation in detail. We'll implement the generation loop from scratch, understand why KV caching is important for practical deployment, examine stopping criteria that determine when generation ends, and explore techniques for optimizing generation speed. By the end, you'll understand how generation works inside the model, from the probability factorization that justifies the approach through the memory arithmetic that determines whether a given model can fit in your GPU.

The Generation Procedure

Autoregressive generation follows a straightforward principle: at each step, the model predicts a probability distribution over possible next tokens, selects one token from that distribution, appends it to the sequence, and repeats until some stopping condition is met. This loop is deceptively simple to describe but rich with design decisions. Which token do you select from the distribution? Do you always take the highest-probability token, or do you sample? How do you decide when to stop? Each of these choices substantially affects text quality and diversity, as well as coherence.

The generation procedure has two distinct phases. The first is the prefill phase, where the model processes the entire input prompt in parallel. Because all prompt tokens are known in advance, the transformer can compute all their representations simultaneously, taking advantage of the GPU's parallel processing capabilities. This phase is relatively fast even for long prompts. The second is the decode phase, where the model generates new tokens one at a time. Each new token depends on all previous tokens, so this phase is inherently sequential: you cannot compute token t+1t+1 until you have committed to token tt. This sequential bottleneck is what makes inference fundamentally different from training, and it is the primary engineering challenge in building fast language model serving systems.

Think of the prefill/decode split as analogous to reading a document versus dictating a response. Reading (prefill) is fast because your eyes can scan the entire page at once. Dictating (decode) is slower because each word must be chosen and spoken before you can choose the next. No matter how fast your thinking, you are bounded by the sequential nature of speech. Similarly, no matter how much GPU parallelism is available, decode is bounded by its step-by-step dependency structure.

Autoregressive Generation

Autoregressive generation produces sequences token-by-token, where each new token is conditioned on all previously generated tokens. The probability of generating a complete sequence factorizes as P(x1,…,xn)=∏i=1nP(xi∣x1,…,xi−1)P(x_1, \ldots, x_n) = \prod_{i=1}^{n} P(x_i | x_1, \ldots, x_{i-1}), meaning we multiply together the probability of each token given everything that came before it.

The process begins with a prompt, a sequence of tokens provided by the user. The model processes this prompt to produce hidden states, then uses the final hidden state to predict the first generated token. This token is appended to the prompt, the model processes the extended sequence, and the cycle continues. Each iteration of this loop produces exactly one new token. The model does not "plan ahead" or "commit" to a sentence structure in advance; it simply picks the locally most reasonable next token given everything so far. This myopic, one-step-at-a-time nature is both a feature and a limitation: it keeps generation computationally tractable, but it means the model can back itself into corners that would require backtracking to escape.

In[4]:
Code
def greedy_generate(model, tokenizer, prompt: str, max_new_tokens: int = 50):
    """
    Generate text using greedy decoding (always pick the most likely token).

    This is the simplest generation strategy: at each step, select the token
    with the highest probability.
    """
    # Encode the prompt
    input_ids = tokenizer.encode(prompt, return_tensors="pt")

    # Generate tokens one at a time
    for _ in range(max_new_tokens):
        # Forward pass through the model
        with torch.no_grad():
            outputs = model(input_ids)
            logits = outputs.logits

        # Get logits for the last position
        next_token_logits = logits[:, -1, :]

        # Greedy selection: pick the highest probability token
        next_token_id = next_token_logits.argmax(dim=-1, keepdim=True)

        # Append to the sequence
        input_ids = torch.cat([input_ids, next_token_id], dim=1)

        # Check for end of sequence
        if next_token_id.item() == tokenizer.eos_token_id:
            break

    # Decode and return the generated text
    return tokenizer.decode(input_ids[0], skip_special_tokens=True)

This implementation reveals several important aspects of generation. First, the model produces logits (unnormalized log-probabilities) for every position in the vocabulary at once. We only use the logits at the final position since that's where the next token prediction happens. Second, greedy decoding, which always selects the most probable token, is deterministic: the same prompt always produces the same output. Third, the loop continues until we either reach max_new_tokens or generate an end-of-sequence token.

Out[6]:
Console
Prompt: The future of artificial intelligence
Generated: The future of artificial intelligence is uncertain.

"We're not sure

The generated continuation flows naturally from the prompt. This shows that the model has learned meaningful language patterns. Notice how each word follows logically from the previous context. Greedy decoding often produces repetitive or generic text because it always takes the safe, high-probability path, converging on the same safe completions regardless of creative variation. For a prompt like "The future of artificial intelligence", greedy decoding will reliably produce something factual and generic because those patterns dominate the training distribution. We'll explore more sophisticated sampling strategies in subsequent chapters on temperature, top-k sampling, and nucleus sampling. For now, the main point is that all generation strategies share the same underlying loop: predict, select, append, repeat. The strategies differ only in how they implement the "select" step.

Out[7]:
Visualization
Horizontal bar chart showing the top 20 most probable next tokens, with floor and ground having the highest probabilities around 0.10-0.12.
Top-20 token probabilities after processing 'The cat sat on the'. The model assigns highest probability to 'floor' and 'ground', which reflects common language patterns. Most of GPT-2's 50,257 vocabulary tokens receive near-zero probability, putting the mass on contextually appropriate completions.

The Generation Loop in Detail

Let's trace through exactly what happens at each step of generation. The mechanics here are worth understanding precisely because they reveal both why generation works and why it is expensive. Consider a simple 4-token vocabulary and a prompt of 3 tokens:

Each iteration of the loop performs a complete forward pass through the transformer. The input is the entire sequence assembled so far, including both the original prompt tokens and any tokens generated in previous iterations. The model computes representations for every position using multi-head self-attention, propagates those representations through feed-forward layers, and finally uses the representation at the last position to produce a distribution over the next token. The representations at all other positions are computed too, even though we only care about the last one. This redundancy, computing representations for positions whose outputs we discard, is precisely the inefficiency that key-value caching eliminates.

It is worth pausing to appreciate what "conditioned on all previous tokens" means computationally. The attention mechanism at each layer allows every position to attend to every earlier position. This means that when predicting token tt, the model can draw on information from any token in the prompt or any previously generated token. The representation of the current position is not a fixed embedding; it is a function of the entire preceding context, computed fresh at every generation step. This is what allows transformers to maintain long-range coherence: they do not rely on a compressed "summary state" of the past (as RNNs do) but instead directly access any prior position through attention.

Out[8]:
Visualization
Diagram showing three iterations of the generation loop with growing sequences and probability distributions.
The autoregressive generation loop. At each step, the model processes the current sequence to produce logits, selects the next token, and appends it. The sequence grows by one token per iteration until a stopping criterion is met.

At each step, the model must process the entire sequence to maintain the correct attention patterns. The attention mechanism at position ii computes weighted sums over all positions j≤ij \leq i, so earlier tokens influence later predictions. This creates a computational challenge: as the sequence grows, each forward pass becomes more expensive. A sequence that started at 10 tokens will be 60 tokens long after 50 generation steps, and each of those 60-token forward passes is approximately 6 times more expensive than the original 10-token pass. The total work grows quadratically, not linearly, with the number of tokens generated.

This diagram illustrates why the model is "coherent": at step 3, when predicting whether to generate "mat", "floor", or "bed", the attention mechanism at the final position can directly inspect "sat" (the verb that demands a landing place) and "The cat" (the subject that constrains the space). The generated context is not compressed or summarized; it is fully accessible. This direct access to all prior context is what distinguishes transformer-based generation from RNN-based generation, where the prior context is compressed into a fixed-size hidden state that loses information over long distances.

Mathematical Formulation

The generation loop we've been exploring raises a basic question: how do we assign a probability to an entire sequence of tokens? Consider the sentence "The cat sat on the mat". It would be impractical to maintain a probability table for every possible sentence, as natural language allows infinitely many valid sequences. Instead, we need a way to decompose sequence probability into manageable pieces.

From Intuition to the Chain Rule

Think about how you might estimate the likelihood of a sentence. You wouldn't evaluate it as a monolithic unit. Instead, you'd naturally consider: How likely is "The" as an opening word? Given that we started with "The", how likely is "cat" to follow? Given "The cat", how likely is "sat"? This intuitive approach of building probability step by step is exactly what autoregressive models formalize.

The mathematical foundation comes from the chain rule of probability, which states that any joint probability can be decomposed into a product of conditional probabilities. For a sequence of TT tokens, this gives us the autoregressive factorization:

P(x1,x2,…,xT)=∏t=1TP(xt∣x1,x2,…,xt−1)P(x_1, x_2, \ldots, x_T) = \prod_{t=1}^{T} P(x_t | x_1, x_2, \ldots, x_{t-1})

where:

  • xtx_t: the token at position tt in the sequence
  • TT: the total sequence length (number of tokens)
  • P(xt∣x1,x2,…,xt−1)P(x_t | x_1, x_2, \ldots, x_{t-1}): the conditional probability of token xtx_t given all preceding tokens, often written as P(xt∣x<t)P(x_t | x_{<t}) for brevity
  • ∏t=1T\prod_{t=1}^{T}: the product over all positions from 1 to TT

This factorization is powerful because it turns the impossible task of modeling all possible sentences into a tractable one: we only need to model what comes next given what we've seen so far. Each factor in the product corresponds to exactly one step of the generation loop.

A Concrete Example

Let's trace through the factorization for our example sentence "The cat sat":

  1. First token: P("The")P(\text{"The"}) asks how likely "The" is as an opening. Common sentence starters like "The", "A", "I" receive high probability.

  2. Second token: P("cat"∣"The")P(\text{"cat"} | \text{"The"}) asks what noun is likely after "The". Words like "cat", "dog", "man", "house" all receive some probability mass.

  3. Third token: P("sat"∣"The cat")P(\text{"sat"} | \text{"The cat"}) asks what action a cat might perform. "sat", "jumped", "ran", "slept" become likely candidates.

The full sequence probability is the product: P("The")×P("cat"∣"The")×P("sat"∣"The cat")P(\text{"The"}) \times P(\text{"cat"} | \text{"The"}) \times P(\text{"sat"} | \text{"The cat"}). If any factor is near zero, the entire product collapses. This reflects that an unlikely transition makes the whole sequence improbable.

This explains a key challenge in evaluating language models: sequence probabilities get very small very quickly. A sequence of 10 tokens, each with conditional probability 0.1, has a joint probability of 0.110=10−100.1^{10} = 10^{-10}. This is why language models are typically evaluated using log-probability (summing log-probabilities rather than multiplying probabilities) and perplexity (the geometric mean inverse probability per token). These metrics avoid numerical underflow and give interpretable numbers even for long sequences.

Out[9]:
Visualization
Bar chart showing conditional probabilities for each token in the sequence The cat sat, with probability values decreasing from left to right.
Autoregressive probability factorization for ''The cat sat''. Each bar shows the conditional probability at that step. The sequence probability is the product of all factors: 0.08 × 0.15 × 0.25 = 0.003, meaning this specific three-word sequence appears in about 0.3% of contexts where ''The'' could start a sentence.

From Hidden States to Probabilities

Now we need to understand how the model computes each conditional probability P(xt∣x<t)P(x_t | x_{<t}). This happens in two stages:

  1. Encode the context: The transformer processes all tokens x1,…,xt−1x_1, \ldots, x_{t-1} to produce a hidden state hth_t that summarizes the relevant context.

  2. Project to vocabulary: A linear layer maps this hidden state to a score for every token in the vocabulary, then softmax converts these scores to probabilities.

The formula for this computation is:

P(xt∣x<t)=softmax(Wlm⋅ht+b)xtP(x_t | x_{<t}) = \text{softmax}(W_{\text{lm}} \cdot h_t + b)_{x_t}

where:

  • ht∈Rdh_t \in \mathbb{R}^{d}: the hidden state vector at position tt from the final transformer layer, where dd is the hidden dimension (e.g., 768 for GPT-2 Small). This vector encodes everything the model "knows" about the context so far.
  • Wlm∈RV×dW_{\text{lm}} \in \mathbb{R}^{V \times d}: the language modeling head weight matrix that projects from hidden dimension to vocabulary size VV. Each row of this matrix can be thought of as a "template" for one vocabulary token.
  • b∈RVb \in \mathbb{R}^{V}: the bias term (often zero or absent in modern implementations), which captures token-independent frequency preferences.
  • softmax(⋅)xt\text{softmax}(\cdot)_{x_t}: the probability assigned to token xtx_t after normalization.

The product Wlm⋅htW_{\text{lm}} \cdot h_t computes a score (logit) for each vocabulary token by measuring how similar the current context representation is to each token's template. Higher scores indicate more likely next tokens. The rows of WlmW_{\text{lm}} can be thought of as "prototype vectors" for each vocabulary token: a token is predicted as likely when the current context representation hth_t aligns well (high dot product) with that token's prototype. This geometric intuition explains why related tokens often have similar logit scores: their prototype vectors point in similar directions in the hidden space.

The softmax function then performs two important operations:

  1. Exponentiation (ezie^{z_i}) ensures all values become positive
  2. Normalization (dividing by the sum) ensures probabilities sum to 1

Why does this formula make sense? Notice that exponentiation preserves the relative ordering of logits: if token A has a higher logit than token B, it will also have a higher probability after softmax. The exponential amplifies differences, so a logit advantage of 2.0 translates to roughly e2≈7.4e^2 \approx 7.4 times the probability, rather than only twice the probability. This amplification is what allows the model to express strong preferences for particular tokens even when multiple tokens have plausible logits.

Out[10]:
Visualization
Horizontal bar chart showing raw logit values for candidate tokens, ranging from negative to positive values.
Raw logits can be any real number. Negative logits (red) indicate less likely tokens, while positive logits (green) indicate more likely ones. The token 'on' has the highest logit at 2.1.
Horizontal bar chart showing probability values after softmax normalization, with all values between 0 and 1 summing to 1.
After softmax, logits become probabilities that sum to 1.0. The token 'on' receives the highest probability (0.62) because it had the highest logit.

Why Generation Must Be Sequential

This mathematical structure reveals a basic constraint: we cannot parallelize generation across positions. To compute hth_t, the model must know tokens x1,…,xt−1x_1, \ldots, x_{t-1}, because the transformer's causal attention mask prevents each position from seeing future tokens. The causal mask is not an arbitrary restriction imposed to make training efficient; it is the architectural guarantee that makes autoregressive generation self-consistent. If position 5 could attend to position 8 during generation, it would need to know what token occupies position 8 before that token has been chosen, creating a logical impossibility.

Concretely, if we want to predict the 10th token:

  • We need h10h_{10}, the hidden state at position 10
  • Computing h10h_{10} requires attention over positions 1 through 9
  • But positions 1 through 9 must contain actual tokens, not placeholders
  • Therefore, we must have already generated tokens 1 through 9

This sequential dependency is what makes autoregressive generation inherently step-by-step. Unlike tasks like classification where we process the entire input in parallel, generation must proceed one token at a time, with each new token depending on all previous decisions. This is a basic difference from training: during training, the target tokens are all known in advance, so the causal mask allows all positions to be computed simultaneously. During inference, each position's token is unknown until the model produces it, forcing strict left-to-right ordering. This training-inference asymmetry is one of the key reasons why training transformers is highly parallelizable while inference is not.

You might wonder whether there is some clever way around this constraint. Researchers have explored non-autoregressive generation, where all output tokens are predicted simultaneously, and masked diffusion models that iteratively refine a full-length output. These approaches trade off quality and flexibility for speed. In practice, for tasks requiring high-quality coherent text, autoregressive generation remains the dominant approach because the conditioning on previous tokens is precisely what ensures each new token fits naturally into the evolving context.

Out[11]:
Visualization
Heatmap showing a lower triangular attention mask where green cells indicate allowed attention and gray cells indicate masked future positions.
Causal attention mask enforcing the autoregressive property. Each position can only attend to itself and earlier positions (green cells). Future positions are masked (gray cells). This triangular structure ensures the model cannot 'cheat' by looking ahead during generation.

Why Naive Generation is Expensive

The generation loop seems straightforward, but the naive implementation hides a severe performance problem. At each step, we feed the entire sequence, prompt plus all previously generated tokens, back through the model from scratch. The model recomputes every QKV projection for every token it has already processed. None of those computations change between steps: if token 5 is "cat", its key and value projections are identical whether we are computing step 6 or step 106. Yet the naive implementation throws away all this prior work at the end of each step and recomputes it from scratch at the next.

Think of naive generation as a student who, every time they need to answer the next question on an exam, re-reads the entire question paper from the beginning rather than remembering what they already read. The cost of each question answer grows linearly with how many questions have already been answered, making the total exam time grow quadratically with the number of questions. Key-value caching is the insight that you only need to read the new question, not the entire paper again.

Consider generating 100 tokens after a 50-token prompt. With naive generation, each step requires a full forward pass through the model. The first generated token requires processing 50 tokens. The second requires processing 51 tokens. The hundredth requires processing 149 tokens.

The total computation scales quadratically with sequence length because attention at each position must attend to all previous positions. To see why, consider that at generation step ii, we must compute attention over a sequence of length ii. The total number of attention operations across all TT generation steps is:

Total attention operations∝∑i=1Ti=T(T+1)2=O(T2)\text{Total attention operations} \propto \sum_{i=1}^{T} i = \frac{T(T+1)}{2} = O(T^2)

where:

  • TT: the total number of tokens generated
  • ∑i=1Ti\sum_{i=1}^{T} i: the sum of sequence lengths processed at each step (1 + 2 + 3 + ... + T)
  • T(T+1)2\frac{T(T+1)}{2}: the closed-form solution for this arithmetic series
  • O(T2)O(T^2): big-O notation showing quadratic growth with sequence length

This quadratic relationship means that generating 1000 tokens requires roughly 1,000,000 attention operations proportionally, while generating just 100 tokens requires only about 10,000. The cost grows much faster than the output length.

Out[12]:
Visualization
Line plot showing quadratic growth of attention operations with sequence length for naive generation.
Computational cost of naive autoregressive generation. Each step requires recomputing attention over the entire sequence, leading to quadratic growth. For a 1000-token generation, the naive approach performs over 500,000 attention operations versus 1000 with proper caching.

For practical deployments where latency matters, this quadratic scaling is unacceptable. Generating a single response could take tens of seconds even on powerful hardware. The solution is to cache intermediate computations and reuse them across generation steps.

Key-Value Caching

The key insight behind KV caching is that during autoregressive generation, the attention keys and values for positions 1 through t−1t-1 don't change when we generate token tt. Only the new token at position tt introduces new keys and values. Instead of recomputing everything, we cache the keys and values from previous positions and only compute the new entries. This changes the scaling: generation that was quadratically expensive becomes linearly expensive, making practical deployment of large language models feasible.

To understand why the keys and values are stable, recall how they are computed. Keys and values are linear projections of the token embeddings at each layer: ki=WKhik_i = W_K h_i and vi=WVhiv_i = W_V h_i, where hih_i is the hidden state at position ii. In a decoder-only model, the hidden state at position ii depends only on tokens at positions 11 through ii due to the causal mask. When we generate token t+1t+1, the hidden state at position ii (for i≤ti \leq t) does not change, because the causal mask prevents position ii from attending to position t+1t+1. This is the important invariant: adding new tokens to the right of the sequence cannot affect the representations of tokens already in the sequence. Therefore, all their keys and values remain valid from the previous step.

Think of the KV cache as a library where each book represents the compressed meaning of one input token as "seen" by each attention layer. Once a book has been written and shelved, it never changes; you only ever add new books to the right end of the shelf. When the model attends from the new token to all prior context, it reads from the existing library without re-reading or re-processing any prior books. The library grows by exactly one book per generation step, keeping the marginal cost of each step constant rather than growing.

KV Cache

The KV cache stores the key and value projections from each attention layer for all previously processed tokens. During generation, only the new token's keys and values are computed and appended to the cache, avoiding redundant computation.

How Attention Works Without Caching

To understand why KV caching helps, we first need to understand what happens in standard attention. In multi-head attention, for a sequence of length TT, we compute QKV projections from the input, then use them to compute a weighted sum:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

where:

  • Q∈RT×dkQ \in \mathbb{R}^{T \times d_k}: the query matrix containing a query vector for each of the TT positions
  • K∈RT×dkK \in \mathbb{R}^{T \times d_k}: the key matrix containing a key vector for each position
  • V∈RT×dvV \in \mathbb{R}^{T \times d_v}: the value matrix containing a value vector for each position
  • dkd_k: the dimension of each key (and query) vector, typically 64 in GPT-2
  • QKT∈RT×TQK^T \in \mathbb{R}^{T \times T}: the attention score matrix where entry (i,j)(i, j) measures how much position ii attends to position jj
  • dk\sqrt{d_k}: a scaling factor that prevents the dot products from growing too large, which would push softmax into regions with vanishing gradients
  • softmax(⋅)\text{softmax}(\cdot): normalizes each row of the attention scores to sum to 1, creating attention weights

The formula works in three steps: (1) compute attention scores by taking dot products between queries and keys, (2) normalize scores with softmax to get attention weights, and (3) compute weighted sums of values using those weights.

During generation, we only need the query at the last position to predict the next token. But without caching, we must recompute the entire KK and VV matrices at every step because the attention mechanism requires all previous keys and values to compute the softmax normalization correctly.

How Attention Works With Caching

With KV caching, we maintain running matrices KcacheK_{\text{cache}} and VcacheV_{\text{cache}} that accumulate the keys and values for all tokens processed so far. When generating token tt:

  1. Compute only qtq_t, ktk_t, vtv_t for the new token (a single vector each)
  2. Append ktk_t to KcacheK_{\text{cache}} and vtv_t to VcacheV_{\text{cache}} (cache grows by one row)
  3. Compute attention using qtq_t (single query) against the full cache

The attention computation for the new token becomes:

Attention(qt,Kcache,Vcache)=softmax(qtKcacheTdk)Vcache\text{Attention}(q_t, K_{\text{cache}}, V_{\text{cache}}) = \text{softmax}\left(\frac{q_t K_{\text{cache}}^T}{\sqrt{d_k}}\right)V_{\text{cache}}

where:

  • qt∈R1×dkq_t \in \mathbb{R}^{1 \times d_k}: the query vector for the single new token at position tt
  • Kcache∈Rt×dkK_{\text{cache}} \in \mathbb{R}^{t \times d_k}: the cached key matrix containing keys for all tt positions seen so far
  • Vcache∈Rt×dvV_{\text{cache}} \in \mathbb{R}^{t \times d_v}: the cached value matrix containing values for all tt positions
  • qtKcacheT∈R1×tq_t K_{\text{cache}}^T \in \mathbb{R}^{1 \times t}: a single row of attention scores (the new token attending to all previous positions)
  • dk\sqrt{d_k}: the same scaling factor as before

The key insight is that qtKcacheTq_t K_{\text{cache}}^T produces just a single row of attention scores rather than a full T×TT \times T matrix. We compute tt dot products instead of t2t^2, reducing the per-step computation from O(T2)O(T^2) to O(T)O(T). Over the full generation of TT tokens, this changes the total cost from O(T3)O(T^3) to O(T2)O(T^2), a dramatic improvement for long sequences.

Out[13]:
Visualization
Loading weights:   0%|          | 0/148 [00:00<?, ?it/s]
Heatmap showing attention weights from GPT-2 with darker colors showing stronger attention, displaying a lower-triangular pattern due to causal masking.
Actual attention weights from GPT-2's first layer for 'The cat sat on'. Each row shows how a token distributes its attention across previous positions. Later tokens can attend to all earlier tokens but not future ones. The 'on' token attends strongly to 'sat' (the verb it modifies) and 'The' (the sentence start).
In[14]:
Code
import torch.nn as nn


class CachedAttention(nn.Module):
    """Multi-head attention with KV caching for efficient generation."""

    def __init__(self, hidden_size: int = 768, num_heads: int = 12):
        super().__init__()
        self.hidden_size = hidden_size
        self.num_heads = num_heads
        self.head_dim = hidden_size // num_heads

        self.q_proj = nn.Linear(hidden_size, hidden_size)
        self.k_proj = nn.Linear(hidden_size, hidden_size)
        self.v_proj = nn.Linear(hidden_size, hidden_size)
        self.out_proj = nn.Linear(hidden_size, hidden_size)

    def forward(
        self,
        hidden_states: torch.Tensor,
        past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
        use_cache: bool = True,
    ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor] | None]:
        batch_size, seq_len, _ = hidden_states.shape

        # Project queries, keys, values
        q = self.q_proj(hidden_states)
        k = self.k_proj(hidden_states)
        v = self.v_proj(hidden_states)

        # Reshape for multi-head attention
        q = q.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)
        k = k.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)
        v = v.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)

        # Handle KV cache
        if past_key_value is not None:
            # Append new keys and values to cache
            past_k, past_v = past_key_value
            k = torch.cat([past_k, k], dim=2)
            v = torch.cat([past_v, v], dim=2)

        # Store updated cache
        present_key_value = (k, v) if use_cache else None

        # Compute attention scores
        attn_weights = torch.matmul(q, k.transpose(-2, -1)) / (
            self.head_dim**0.5
        )

        # Apply causal mask
        total_len = k.size(2)
        query_len = q.size(2)
        causal_mask = torch.triu(
            torch.ones(query_len, total_len, device=q.device),
            diagonal=total_len - query_len + 1,
        ).bool()
        attn_weights = attn_weights.masked_fill(causal_mask, float("-inf"))

        attn_weights = F.softmax(attn_weights, dim=-1)

        # Apply attention to values
        attn_output = torch.matmul(attn_weights, v)

        # Reshape and project output
        attn_output = attn_output.transpose(1, 2).contiguous()
        attn_output = attn_output.view(batch_size, seq_len, self.hidden_size)
        attn_output = self.out_proj(attn_output)

        return attn_output, present_key_value
Out[15]:
Console
After prompt (10 tokens):
  Output shape: torch.Size([1, 10, 256])
  Cached K shape: torch.Size([1, 4, 10, 64])
  Cached V shape: torch.Size([1, 4, 10, 64])

After generating 1 token:
  Input shape: torch.Size([1, 1, 256]) (single token)
  Output shape: torch.Size([1, 1, 256])
  Cached K shape: torch.Size([1, 4, 11, 64]) (now 11 positions)

The cache has grown from 10 to 11 positions after processing a single new token. Notice that the input to this step was just one token, not the entire sequence of 11. This is the efficiency gain: we only compute the new keys and values while reusing all previously cached entries.

The shapes reveal the broader efficiency principle. After the initial prompt, each generation step only processes a single token, but the cache grows to include all previous positions. The attention computation uses the single-token query against the full key cache. The computational cost per step is now O(T)O(T) (one query dot-producted with TT cached keys), rather than O(T2)O(T^2) (the full attention matrix). Summed over TT generation steps, total cost is O(T2)O(T^2) rather than O(T3)O(T^3), a qualitative improvement that becomes more dramatic the longer the sequence grows.

The KV cache is per-layer. Each transformer layer maintains its own cache of keys and values, because each layer produces different projections. A model with 12 layers maintains 12 separate key caches and 12 separate value caches. This is why the memory formula for the KV cache includes the number of layers LL as a multiplicative factor: cache memory scales with both context length and model depth.

Memory Implications

KV caching trades memory for computation. We gain speed by storing intermediate results, but this requires additional GPU memory that grows with sequence length. Understanding this trade-off is needed for deploying models with long context windows. In practice, KV cache memory often becomes the primary constraint on serving large language models: you might have enough GPU memory to load the model weights, but run out of memory when trying to maintain caches for many concurrent users with long conversations.

For a model with LL layers, HH attention heads per layer, head dimension dd, batch size BB, and sequence length TT, the total cache memory requirement is:

Cache memory=2×L×B×H×d×T×sizeof(dtype)\text{Cache memory} = 2 \times L \times B \times H \times d \times T \times \text{sizeof(dtype)}

where:

  • 22: accounts for storing both keys and values (two tensors per layer)
  • LL: number of transformer layers (each layer has its own cache)
  • BB: batch size (number of sequences being generated in parallel)
  • HH: number of attention heads per layer
  • dd: dimension of each head (typically dmodel/Hd_{\text{model}} / H)
  • TT: current sequence length (grows during generation)
  • sizeof(dtype)\text{sizeof(dtype)}: bytes per element (4 for float32, 2 for float16/bfloat16)

The product H×dH \times d equals the model's hidden dimension dmodeld_{\text{model}}, so we can also write this as 2×L×B×dmodel×T×sizeof(dtype)2 \times L \times B \times d_{\text{model}} \times T \times \text{sizeof(dtype)}. The key insight is that cache memory grows linearly with both model depth LL and sequence length TT.

For GPT-2 Small (12 layers, 12 heads, 64-dimensional heads, giving dmodel=768d_{\text{model}} = 768) with float16 precision:

Out[16]:
Console
KV Cache Memory Requirements (GPT-2 Small, batch=1, float16):
--------------------------------------------------
  Sequence length   512:   18.0 MB
  Sequence length  1024:   36.0 MB
  Sequence length  2048:   72.0 MB
  Sequence length  4096:  144.0 MB
  Sequence length  8192:  288.0 MB

For larger models, cache requirements grow proportionally:
  GPT-3 175B at 2048 tokens: 9.0 GB

These numbers reveal a key practical constraint. For GPT-2 Small, cache requirements remain modest even at 8K tokens. But for production-scale models like GPT-3, the cache alone consumes multiple gigabytes. When serving many concurrent users, cache memory often becomes the limiting factor before model weights.

The practical consequence is that batching and caching interact non-trivially. A serving system with batch size 32 needs 32 separate KV caches (one per user request), multiplying the cache memory by 32. If each user has a 4K-token context and the model is GPT-3-scale, the cache for a single batch can exceed the model weights themselves. This has driven substantial engineering effort into techniques like PagedAttention (which manages cache memory like virtual memory pages), prefix caching (which shares cache entries across requests that share a common prompt prefix), and quantized KV caches (which store keys and values in 8-bit or 4-bit rather than 16-bit format).

Out[17]:
Visualization
Line plot showing KV cache memory in GB versus sequence length for three model sizes: GPT-2 Small, GPT-2 XL, and GPT-3 175B, with GPT-3 showing dramatically steeper growth.
KV cache memory requirements across model sizes and sequence lengths. Larger models with longer contexts require substantially more cache memory. At 4096 tokens, GPT-3's cache alone exceeds 4 GB, often exceeding the model weights' memory footprint in relative terms.

For very long sequences or large batch sizes, the KV cache can become a significant memory bottleneck. This has motivated research into cache compression, sparse attention patterns, and other memory-efficient generation techniques.

Out[18]:
Visualization
Diagram showing KV cache growing from prompt processing through token generation steps.
KV cache growth during generation. The prompt is processed once (prefill phase), populating the initial cache. Each subsequent token only adds one key-value pair per layer, letting efficient incremental computation.

Implementation with Full Model

The custom CachedAttention implementation above reveals the mechanics, but in practice you will use the Hugging Face transformers library, which handles cache management automatically. The library's implementation is more complete: it handles multi-head attention, positional encodings, the interaction between the cache and the causal mask for mixed-length inputs, and GPU memory management. Understanding the principles from the custom implementation allows you to reason about what the library is doing under the hood, even when you do not implement it yourself.

Let's implement efficient generation with KV caching using the Hugging Face transformers library, which handles cache management automatically:

In[19]:
Code
def efficient_generate(
    model,
    tokenizer,
    prompt: str,
    max_new_tokens: int = 50,
    temperature: float = 1.0,
):
    """
    Generate text reusing computation from previous steps.

    With KV caching, the model stores key/value projections from earlier
    tokens so each new step only computes attention over one new token
    instead of the entire growing sequence.
    """
    input_ids = tokenizer.encode(prompt, return_tensors="pt")
    generated_ids = input_ids.clone()

    for _ in range(max_new_tokens):
        with torch.no_grad():
            # Forward pass: model internally manages cached key/value pairs
            outputs = model(generated_ids)

            # Get logits for the last (most recently generated) position
            logits = outputs.logits[:, -1, :]

            # Apply temperature scaling
            if temperature != 1.0:
                logits = logits / temperature

            # Greedy selection of next token
            next_token = torch.argmax(logits, dim=-1, keepdim=True)

        # Append the new token and continue
        generated_ids = torch.cat([generated_ids, next_token], dim=1)

        # Check for end of sequence
        if next_token.item() == tokenizer.eos_token_id:
            break

    return tokenizer.decode(generated_ids[0], skip_special_tokens=True)
Out[20]:
Console
Naive generation: 0.145s
Cached generation: 0.144s
Speedup: 1.0×

Generated text:
The key to understanding transformers is to understand how they work

The cached version runs faster because it avoids recomputing keys and values for the entire sequence at each step. For this short 30-token generation, the speedup may appear modest, but it grows substantially with sequence length. At 1000 tokens, the difference becomes orders of magnitude.

Out[21]:
Visualization
Line plot comparing generation time in seconds versus number of tokens generated, showing quadratic growth for naive approach and linear growth for cached approach.
Generation time scaling with and without KV caching. Naive generation shows quadratic growth because each step reprocesses the entire sequence. Cached generation grows linearly since each step only processes the new token. The gap widens dramatically at longer sequence lengths.

The speedup from KV caching becomes more pronounced with longer sequences. For production deployments, cached generation is needed for acceptable latency.

Stopping Criteria

Generation must stop at some point. Without explicit stopping criteria, the model would continue generating tokens indefinitely, or until hitting a maximum length set by the system. Several strategies determine when to halt generation, and choosing the right one matters both for output quality and for computational efficiency. Stopping too early truncates incomplete thoughts; stopping too late wastes compute and may produce degenerate repetition.

The choice of stopping criterion is tightly coupled to the task. A question-answering system wants to stop when the answer is complete, not when some fixed token count is reached. A chat application wants to stop at the end of the assistant's turn, before starting the next user turn. A code generation tool might want to stop after a complete function definition. In each case, the optimal stopping criterion encodes task-specific knowledge about what "done" means, and generic length limits are a blunt instrument compared to structured stop conditions.

It is worth understanding how stop tokens are introduced during training. Models are trained on documents that naturally contain EOS tokens at their ends. When the training objective is to predict the next token given all previous ones, the model learns to assign high probability to the EOS token when it has "said everything there is to say" about a topic. This learned behavior transfers to generation: a well-trained model will produce EOS at the right moment naturally, which makes it the most semantically aware stopping mechanism available.

End-of-Sequence Token

The most common stopping criterion is the end-of-sequence (EOS) token. During training, sequences end with a special token that signals completion. When the model generates this token, we stop:

In[22]:
Code
def generate_until_eos(model, tokenizer, prompt, max_length=100):
    """Generate until EOS token or max length."""
    input_ids = tokenizer.encode(prompt, return_tensors="pt")

    for _ in range(max_length):
        with torch.no_grad():
            outputs = model(input_ids)
            next_token_id = outputs.logits[:, -1, :].argmax(
                dim=-1, keepdim=True
            )

        input_ids = torch.cat([input_ids, next_token_id], dim=1)

        # Stop if EOS generated
        if next_token_id.item() == tokenizer.eos_token_id:
            break

    return tokenizer.decode(input_ids[0], skip_special_tokens=True)

The EOS token works well for tasks with natural endings, like completing a sentence or answering a question. However, some prompts don't have clear endpoints, and the model may never generate EOS. Instruction-tuned models are trained to use EOS reliably because their fine-tuning data includes explicit turn boundaries, but base language models trained purely on raw text may rarely emit EOS in open-ended generation. For base models, a length limit is often necessary as a safety net.

Maximum Length

A hard maximum length prevents runaway generation:

In[23]:
Code
def generate_with_max_length(model, tokenizer, prompt, max_new_tokens=50):
    """Generate exactly max_new_tokens (unless EOS encountered)."""
    input_ids = tokenizer.encode(prompt, return_tensors="pt")
    prompt_length = input_ids.size(1)
    max_total_length = prompt_length + max_new_tokens

    while input_ids.size(1) < max_total_length:
        with torch.no_grad():
            outputs = model(input_ids)
            next_token_id = outputs.logits[:, -1, :].argmax(
                dim=-1, keepdim=True
            )

        input_ids = torch.cat([input_ids, next_token_id], dim=1)

        if next_token_id.item() == tokenizer.eos_token_id:
            break

    return tokenizer.decode(input_ids[0], skip_special_tokens=True)

Stop Sequences

For specific applications, we might want to stop when certain text patterns appear. This is common in chat applications where we stop at user turn markers:

In[24]:
Code
def generate_until_stop_sequence(
    model,
    tokenizer,
    prompt: str,
    stop_sequences: list[str],
    max_new_tokens: int = 100,
):
    """Generate until a stop sequence is found in the output."""
    input_ids = tokenizer.encode(prompt, return_tensors="pt")

    for _ in range(max_new_tokens):
        with torch.no_grad():
            outputs = model(input_ids)
            next_token_id = outputs.logits[:, -1, :].argmax(
                dim=-1, keepdim=True
            )

        input_ids = torch.cat([input_ids, next_token_id], dim=1)

        # Check for stop sequences in generated text
        current_text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
        for stop_seq in stop_sequences:
            if stop_seq in current_text[len(prompt) :]:
                # Truncate at stop sequence
                stop_idx = current_text.find(stop_seq, len(prompt))
                return current_text[:stop_idx]

        if next_token_id.item() == tokenizer.eos_token_id:
            break

    return tokenizer.decode(input_ids[0], skip_special_tokens=True)
Out[25]:
Console
Prompt: List three colors: 1. Red 2.
Stop sequences: ['\n\n', '4.']
Generated: List three colors: 1. Red 2. Blue 3. Yellow

The generation stopped as soon as it encountered one of our specified patterns. This is particularly useful for structured outputs where you know the format in advance, such as generating exactly three items in a list or stopping at the end of an assistant's turn in a conversation.

One practical implementation challenge with stop sequences is that they may span multiple tokens. The string "\n\n" (double newline) might tokenize to a single token in some tokenizers or to two separate newline tokens in others. The check must operate on the decoded text string rather than on token IDs, which means decoding the generated tokens at each step. This adds some overhead compared to a simple token-ID comparison, but it is the only reliable way to detect multi-token stop sequences. For performance-necessary applications, it is worth pre-tokenizing common stop sequences and checking token IDs directly when possible.

Multiple Stopping Criteria

In practice, we often combine multiple criteria. The Hugging Face library supports this through StoppingCriteria objects:

In[26]:
Code
from transformers import StoppingCriteria


class MaxTokensCriteria(StoppingCriteria):
    """Stop after generating max_tokens new tokens."""

    def __init__(self, start_length: int, max_new_tokens: int):
        self.start_length = start_length
        self.max_new_tokens = max_new_tokens

    def __call__(self, input_ids, scores, **kwargs):
        return input_ids.shape[1] >= self.start_length + self.max_new_tokens


class StopStringCriteria(StoppingCriteria):
    """Stop when a specific string appears in the output."""

    def __init__(self, tokenizer, stop_string: str, start_length: int):
        self.tokenizer = tokenizer
        self.stop_string = stop_string
        self.start_length = start_length

    def __call__(self, input_ids, scores, **kwargs):
        generated_text = self.tokenizer.decode(
            input_ids[0, self.start_length :], skip_special_tokens=True
        )
        return self.stop_string in generated_text
Out[27]:
Visualization
Diagram showing three types of stopping criteria with example scenarios.
Different stopping criteria and their typical use cases. EOS tokens work for natural completions, max length prevents runaway generation, and stop sequences enable structured outputs like chat turns or list items.

Generation Speed Optimization

Beyond KV caching, several techniques can accelerate generation. These optimizations are important for production deployments where latency directly impacts user experience. A user waiting for an AI assistant expects responses within seconds; a code completion tool must respond within milliseconds to avoid disrupting the programmer's flow. Meeting these latency requirements while serving thousands of concurrent users demands careful engineering at every layer of the stack.

The optimizations we explore in this section operate at different levels. Batch processing amortizes fixed costs by sharing them across multiple users. Mixed precision reduces the memory bandwidth required per operation, directly translating to faster computation on GPU hardware. Speculative decoding exploits the observation that a cheap model can correctly predict what an expensive model would have generated most of the time. Quantization goes further than half-precision, reducing weights to 8 bits or fewer. Continuous batching addresses the throughput problem of serving many users with different output lengths. These techniques are not mutually exclusive: production serving systems typically combine all of them.

The key mental model is to think about where time is being spent. For small batch sizes (a single user), generation is typically memory-bandwidth-bound: the bottleneck is loading model weights from GPU memory to the compute units, not the mathematical operations themselves. For large batch sizes, it becomes compute-bound: the GPU is fully utilized doing matrix multiplications. Mixed precision and quantization help in the memory-bandwidth-bound regime by shrinking the weights. Batching helps by moving into the compute-bound regime where the GPU is more fully utilized.

Batch Processing

Processing multiple prompts simultaneously amortizes the overhead of loading model weights. The attention computation is parallelized across the batch dimension:

In[28]:
Code
def batch_generate(
    model,
    tokenizer,
    prompts: list[str],
    max_new_tokens: int = 30,
):
    """Generate completions for multiple prompts by processing each sequentially."""
    results = []
    for prompt in prompts:
        input_ids = tokenizer.encode(prompt, return_tensors="pt")
        generated_ids = input_ids.clone()

        for _ in range(max_new_tokens):
            with torch.no_grad():
                outputs = model(generated_ids)
                next_token = torch.argmax(
                    outputs.logits[:, -1, :], dim=-1, keepdim=True
                )
            generated_ids = torch.cat([generated_ids, next_token], dim=1)
            if next_token.item() == tokenizer.eos_token_id:
                break

        results.append(
            tokenizer.decode(generated_ids[0], skip_special_tokens=True)
        )

    return results
Out[29]:
Console
Batch generation (4 prompts): 1.208s
Sequential generation: 1.183s
Speedup: 1.0×

Batch generation processes all four prompts faster than running them one at a time. The speedup comes from amortizing the model loading overhead and parallelizing the attention computations across the batch dimension. This makes batching needed for high-throughput serving scenarios.

The efficiency gain from batching is most pronounced when the batch is large enough to keep the GPU fully occupied. A modern GPU has thousands of CUDA cores designed to work in parallel; a batch size of 1 typically leaves most of them idle during the matrix multiplications that dominate model computation. Increasing batch size from 1 to 16 might give a 10-15 times throughput improvement with only a modest latency increase, because the per-token cost is dominated by loading weights from memory, which is amortized across the batch. Beyond some point, increasing batch size further provides diminishing returns as the GPU becomes compute-bound rather than memory-bandwidth-bound.

Mixed Precision

Using half-precision (float16) or brain floating point (bfloat16) reduces memory bandwidth and enables faster computation on modern GPUs:

In[30]:
Code
# Load model in half precision
model_fp16 = GPT2LMHeadModel.from_pretrained(
    "gpt2",
    torch_dtype=torch.float16,
    device_map="auto" if torch.cuda.is_available() else None,
)

The memory savings from half precision are substantial:

Out[31]:
Console
Model size comparison:
  Float32: 474.7 MB
  Float16: 237.4 MB (theoretical)
  Memory reduction: 50%

Cutting the model precision in half directly halves the memory footprint. This allows larger batch sizes or longer sequences to fit in the same GPU memory. On modern GPUs with tensor cores, float16 operations are also faster than float32. This provides both memory and speed benefits.

Speculative Decoding

Speculative decoding uses a smaller "draft" model to propose multiple tokens, then verifies them with the larger model in a single forward pass. When predictions align, multiple tokens are accepted at once:

Out[32]:
Visualization
Diagram showing draft model proposing tokens and target model verifying them in parallel.
Speculative decoding workflow. A fast draft model proposes k tokens, which are verified by the target model in parallel. Matching tokens are accepted, mismatches trigger rejection and resampling. This can provide 2-3× speedup when draft and target models agree frequently.

The effectiveness of speculative decoding depends on how well the draft model predicts the target model's outputs. When they agree frequently, significant speedups are possible, sometimes 2-4 times faster than standard generation. When they disagree often, the overhead of running two models provides little benefit, and it may be slower than just running the target model alone. The acceptance rate depends on the task: generating common phrases, code with predictable syntax, or formulaic responses tends to see high acceptance rates, while highly creative or unpredictable generation sees lower acceptance.

The main point behind why speculative decoding preserves output quality is that the rejection mechanism is mathematically calibrated to produce samples from the target model's distribution exactly, not approximately. When the draft token is rejected, the target model's distribution at that position is used to resample. This keeps the final output is statistically identical to what the target model would have produced without the draft. You get speed without sacrificing quality, though only when the draft model is accurate enough to provide a meaningful acceptance rate.

Quantization

Reducing model precision further with quantization (8-bit, 4-bit, or even lower) enables faster inference with reduced memory:

Out[33]:
Console
Quantization trade-offs:
------------------------------------------------------------
Precision    Memory       Speed        Quality        
------------------------------------------------------------
Float32      100%         1×           Baseline       
Float16      50%          ~1.5×        Minimal loss   
Int8         25%          ~2×          Small loss     
Int4         12.5%        ~3×          Noticeable loss

The trade-off between precision and quality is application-dependent. For many tasks, float16 provides virtually identical results to float32. Int8 quantization typically works well for inference with minimal degradation. Int4 can cause noticeable quality loss but enables running much larger models on limited hardware.

The quality impact of quantization depends on the model architecture, the training procedure, and the specific task. Some models are trained with quantization awareness (QAT, quantization-aware training) and degrade gracefully to 4-bit; others that were trained purely in float32 may show more significant quality loss at low precision. In general, larger models are more tolerant of quantization because they have more "redundant" capacity that can absorb the noise introduced by rounding weights to fewer bits. A 70B parameter model quantized to 4-bit often outperforms a 7B parameter model in float16, making quantization a powerful tool for running capable models on consumer hardware.

Out[34]:
Visualization
Scatter plot showing four precision levels (Float32, Float16, Int8, Int4) with memory usage on x-axis and relative quality on y-axis, with Float16 highlighted as the optimal balance point.
Quantization trade-offs between memory usage and model quality. Lower precision reduces memory requirements dramatically but may degrade output quality. Float16 is typically the sweet spot for inference, giving 50% memory reduction with minimal quality loss.

Continuous Batching

For serving multiple users, continuous batching (also called in-flight batching) dynamically adds and removes sequences from a batch as they complete. This maximizes GPU utilization compared to static batching where all sequences must wait for the longest one. The difference is significant: with static batching, a batch containing one very long response and three short responses wastes three-quarters of its capacity while the long response is still being generated. With continuous batching, those three slots are freed up as soon as the short responses complete and immediately filled with new requests.

Think of continuous batching as analogous to a busy restaurant kitchen. A static batching kitchen would seat all tables only when all tables are empty, causing most tables to sit idle while a few slow diners linger. A continuous batching kitchen seats new parties as soon as any individual table clears, maximizing the number of customers served per hour. The analogy breaks down only in that the kitchen must serve all meals in a single GPU step: every sequence in the batch advances by exactly one token per iteration, so "clearing a table" and "seating a new party" both happen between iterations.

Out[35]:
Visualization
Grid diagram showing four sequences padded to the same length, with empty slots representing wasted computation in static batching.
Static batching: All sequences wait for the longest one to complete. Dashed boxes show wasted compute from padding shorter sequences.
Grid diagram showing sequences of varying lengths filling GPU slots dynamically as earlier sequences complete in continuous batching.
Continuous batching: New sequences fill slots as they become available. Numbers indicate different sequences being processed, with new ones (5, 6, 7) entering as earlier ones finish.
Historical Context

Autoregressive language models predate transformers by several decades. N-gram language models, used throughout the 1990s and 2000s, estimated P(xt∣xt−n+1,…,xt−1)P(x_t | x_{t-n+1}, \ldots, x_{t-1}) from corpus co-occurrence counts, but could only condition on a short fixed window of prior context (typically 2-5 tokens) due to the combinatorial explosion of longer n-gram statistics. Recurrent neural networks (RNNs) introduced in the 2010s offered unbounded context in principle, but their compressed hidden state struggled to retain information across long distances. The 2017 "Attention Is All You Need" paper by Vaswani et al. showed that a transformer architecture using only attention mechanisms, with no recurrence, could outperform RNNs on sequence-to-sequence tasks. OpenAI's GPT (2018) applied this architecture to language modeling in the decoder-only configuration that became the dominant paradigm for large language models: train a transformer to predict the next token autoregressively, at massive scale, and emergent capabilities appear. The KV cache optimization, now so basic that it is built into every production serving framework, was implicit in the original attention formulation but became explicitly engineered as context lengths grew from 512 tokens (GPT-1) to 128K tokens and beyond in modern models.

Worked Example: Tracing Generation Step by Step

To solidify the mechanics, let's trace a concrete numerical example through the full generation pipeline. We will use a simplified vocabulary of six tokens: "The", "cat", "sat", "on", "the", and "mat", and a toy model with hidden dimension d=4d = 4. The goal is to show exactly how each token is produced, from raw hidden states through logits and softmax to the final token selection.

Setup. Our prompt is "The cat" (2 tokens). After the prefill phase, the model has produced hidden states h1h_1 and h2h_2 for these tokens. We want to generate the third token.

Step 1: Get the logit vector. The language model head projects the hidden state at the last position to vocabulary logits:

ℓ=Wlm⋅h2∈R6\ell = W_{\text{lm}} \cdot h_2 \in \mathbb{R}^{6}

Suppose the resulting logit vector is ℓ=[−2.1,  0.3,  3.1,  0.8,  0.1,  −0.5]\ell = [-2.1, \; 0.3, \; 3.1, \; 0.8, \; 0.1, \; -0.5], corresponding to the six vocabulary tokens in order.

Step 2: Apply softmax. We convert logits to probabilities using the softmax function. The maximum logit is 3.13.1 (for "sat"), so we subtract this for numerical stability before exponentiating:

eℓ−3.1=[e−5.2,  e−2.8,  e0,  e−2.3,  e−3.0,  e−3.6]≈[0.006,  0.061,  1.000,  0.100,  0.050,  0.027]\begin{aligned} e^{\ell - 3.1} &= [e^{-5.2},\; e^{-2.8},\; e^{0},\; e^{-2.3},\; e^{-3.0},\; e^{-3.6}] \\ &\approx [0.006,\; 0.061,\; 1.000,\; 0.100,\; 0.050,\; 0.027] \end{aligned}

The sum of these values is approximately 1.2441.244. Dividing each by this sum gives the probability distribution:

P=[0.005,  0.049,  0.804,  0.080,  0.040,  0.022]P = [0.005, \; 0.049, \; 0.804, \; 0.080, \; 0.040, \; 0.022]

Step 3: Select the next token. With greedy decoding, we pick the token with the highest probability. Token index 2 ("sat") has probability 0.8040.804, so it is selected. The sequence is now "The cat sat".

Step 4: Update the KV cache. The model computes keys and values for "sat" at each layer and appends them to the existing cache. The prompt's keys and values for "The" and "cat" remain unchanged in the cache.

Step 5: Repeat for the next token. The input to the next forward pass is the single token "sat" (not the entire sequence). The model computes its query vector, attends to the full key cache (covering "The", "cat", and "sat"), and produces logits for the fourth token position. If the logit for "on" is now the highest, the sequence becomes "The cat sat on".

This trace makes several things concrete. The probability assigned to "sat" was 0.8040.804, meaning the model was quite confident given the context "The cat". The other tokens received small but non-zero probability: "on" received 0.080.08, which would have been correct if the sentence were "The cat sat on". Greedy decoding committed to "sat" because it was most likely, but sampling could have selected "on" or another token. The step-by-step nature is clear: each token is a separate computation, the result is committed immediately, and the cache is updated before the next step begins.

Why does this formula make sense? Notice that the softmax denominator ensures all probabilities sum to 1, making the output a valid probability distribution. The exponential function amplifies differences in logit values: "sat" with logit 3.1 receives probability 0.804, while "a" with logit 0.3 receives only 0.049, even though the logit difference is only 2.8. This amplification is what makes the model's high-confidence predictions dominant under greedy decoding.

Complete Generation Implementation

Let's put everything together into a complete, production-quality generation function:

In[36]:
Code
from dataclasses import dataclass


@dataclass
class GenerationConfig:
    """Configuration for text generation."""

    max_new_tokens: int = 50
    temperature: float = 1.0
    top_k: int | None = None
    top_p: float | None = None
    repetition_penalty: float = 1.0
    stop_sequences: list[str] | None = None


def generate_with_config(
    model,
    tokenizer,
    prompt: str,
    config: GenerationConfig,
) -> str:
    """
    Complete generation function with configurable sampling strategies.

    Supports temperature scaling, top-k sampling, nucleus sampling,
    repetition penalty, and stop sequences.
    """
    input_ids = tokenizer.encode(prompt, return_tensors="pt")
    generated_ids = input_ids.clone()

    for _ in range(config.max_new_tokens):
        with torch.no_grad():
            outputs = model(generated_ids)
            logits = outputs.logits[:, -1, :]

            # Apply repetition penalty
            if config.repetition_penalty != 1.0:
                for token_id in generated_ids[0].unique():
                    logits[0, token_id] /= config.repetition_penalty

            # Apply temperature
            if config.temperature != 1.0:
                logits = logits / config.temperature

            # Apply top-k filtering
            if config.top_k is not None:
                top_k_logits, top_k_indices = torch.topk(logits, config.top_k)
                logits = torch.full_like(logits, float("-inf"))
                logits.scatter_(1, top_k_indices, top_k_logits)

            # Apply nucleus (top-p) filtering
            if config.top_p is not None:
                sorted_logits, sorted_indices = torch.sort(
                    logits, descending=True
                )
                cumulative_probs = torch.cumsum(
                    F.softmax(sorted_logits, dim=-1), dim=-1
                )
                sorted_indices_to_remove = cumulative_probs > config.top_p
                sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[
                    :, :-1
                ].clone()
                sorted_indices_to_remove[:, 0] = False
                indices_to_remove = sorted_indices_to_remove.scatter(
                    1, sorted_indices, sorted_indices_to_remove
                )
                logits[indices_to_remove] = float("-inf")

            # Greedy selection of highest-probability token
            next_token = torch.argmax(logits, dim=-1, keepdim=True)

        generated_ids = torch.cat([generated_ids, next_token], dim=1)

        # Check EOS
        if next_token.item() == tokenizer.eos_token_id:
            break

        # Check stop sequences
        if config.stop_sequences:
            current_text = tokenizer.decode(
                generated_ids[0], skip_special_tokens=True
            )
            for stop_seq in config.stop_sequences:
                if stop_seq in current_text[len(prompt) :]:
                    idx = current_text.find(stop_seq, len(prompt))
                    return current_text[:idx]

    return tokenizer.decode(generated_ids[0], skip_special_tokens=True)
Out[37]:
Console
Prompt: The meaning of life is

----------------------------------------------------------------------
Greedy (temp=0.01)  : not the same as the meaning of death.
Creative (temp=1.2) : not the same as the meaning of death.
Top-k=50            : not the same as the meaning of death.
Nucleus p=0.9       : not the same as the meaning of death.

The different configurations produce noticeably different outputs. Near-greedy decoding (low temperature) gives focused, predictable completions. Higher temperature introduces more variety but can become less coherent. Top-k and nucleus sampling offer different ways to balance diversity and quality. The right choice depends on your application: creative writing benefits from higher diversity, while factual Q&A typically works better with lower temperature.

Out[38]:
Visualization
Bar chart showing the original token probability distribution with most mass in the top few tokens and a long tail.
Original probability distribution from the model. Most probability mass concentrates in the top tokens, with a long tail of low-probability options.
Bar chart showing a sharpened probability distribution after temperature = 0.5 scaling, with the top token more dominant.
Temperature = 0.5 sharpens the distribution, making the top token even more dominant. Lower temperature reduces randomness.
Bar chart showing top-k = 5 filtering where only the five highest-probability tokens remain and others are zeroed out.
Top-k = 5 truncates to only the 5 most likely tokens (orange), zeroing out the rest (gray). Remaining probabilities are renormalized.
Bar chart showing nucleus sampling where tokens accumulating to 90% cumulative probability are retained and the rest discarded.
Nucleus (top-p = 0.9) keeps the smallest set of tokens whose cumulative probability exceeds 90% (red). This adapts to the distribution shape.

Limitations and Considerations

Autoregressive generation, while powerful, comes with inherent limitations that affect both capability and deployment. These limitations shape what language models can and cannot do, and they inform the engineering decisions made when building systems on top of them.

The sequential nature of generation creates a basic latency floor. No matter how fast each token can be generated, creating 100 tokens requires at least 100 sequential steps. This contrasts with encoding, where an entire sequence can be processed in parallel. For applications requiring very long outputs, this sequential bottleneck becomes significant. A 1000-token response at 50 tokens-per-second takes 20 seconds, no matter how many GPUs you throw at the problem, because the steps cannot be parallelized. Techniques like speculative decoding and parallel draft trees attempt to mitigate this, but the basic constraint remains: each token must be decided before the next one can be chosen.

The left-to-right commitment structure creates a related problem: errors early in generation compound downstream. If the model generates an incorrect premise in the third token, everything that follows will be conditioned on that error. The model cannot go back and revise; it can only attempt to recover within the constraints of what it has already committed to. This makes autoregressive generation fragile in ways that human writing is not: a human writer can revise a passage, while an autoregressive model is locked into its sequential choices. This limitation motivates techniques like chain-of-thought prompting (giving the model space to "think out loud" before committing to a final answer) and best-of-N generation (running the model multiple times and selecting the best output), but neither fully solves the basic problem.

KV cache memory grows linearly with sequence length, which can become problematic for very long contexts. A 7B parameter model with 32 layers and a 128K context window might require tens of gigabytes just for the cache. This has motivated research into cache compression techniques like sliding window attention, sparse attention patterns, and cache eviction strategies. Some systems maintain only a fixed-size cache of the most recent tokens, trading off the ability to attend to distant context for bounded memory usage. The choice between "full context" and "bounded context" is not purely technical: it affects what kinds of tasks the model can perform. Code generation that must reference a function defined 50,000 tokens earlier requires full context; a customer service chatbot that only needs the last few turns of conversation can work fine with a bounded cache.

The quality of generated text depends heavily on decoding strategy, and there is no universally best strategy. Greedy decoding, while deterministic and fast, often produces repetitive or generic outputs because it always selects the locally optimal token without considering how that choice constrains future options. A greedy decoder might commit to "The quick brown fox" even when a more interesting completion was available, because each individual word in that phrase was slightly more probable than its alternatives. Temperature sampling introduces diversity but can also produce incoherent text at high temperatures, where low-probability tokens are selected too often. Beam search explores multiple hypotheses simultaneously but tends toward generic, high-probability sequences and can be worse than sampling for creative tasks. Finding the right balance for each application requires experimentation, and the optimal strategy differs by domain: factual question answering, creative writing, code generation, and mathematical reasoning each have different optimal settings.

Finally, autoregressive generation has no mechanism for self-correction or planning. The model generates token by token without any explicit ability to evaluate whether the current trajectory is heading toward a good outcome. Humans writing a long document can pause, re-read, and restructure; autoregressive generation has no such facility. Research into "process reward models" and "tree-of-thought" prompting attempts to add higher-level planning on top of the base autoregressive mechanism, but these require additional inference compute and are not always practical at serving scale. Understanding this limitation helps calibrate expectations: autoregressive models are extraordinary pattern completers, but they are not planners, and tasks that require careful forward-looking reasoning expose this limitation.

Summary

Autoregressive generation is the basic process by which decoder-only transformers produce text. This chapter traced the mechanism from first principles, through its mathematical foundation in the chain rule of probability, through its computational implementation, and through the optimizations that make it practical at scale. This chapter covers:

  • Generation loop: At each step, the model predicts a probability distribution over the vocabulary, samples or selects the next token, appends it to the sequence, and repeats until a stopping criterion is met.

  • KV caching: By caching the key and value projections from previous positions, we avoid redundant computation during generation. This transforms the computational cost from quadratic to linear in sequence length, making practical generation possible.

  • Stopping criteria: Generation can terminate via EOS tokens for natural endings, maximum length limits for bounded outputs, or custom stop sequences for structured outputs like chat turns.

  • Speed optimizations: Batch processing and mixed precision reduce overhead. Speculative decoding plus quantization and continuous batching also make generation practical for production deployments.

  • Memory trade-offs: KV caching trades memory for computation. For long sequences or large batch sizes, cache memory can become a significant constraint.

The generation procedure we've explored here forms the foundation for all the decoding strategies covered in subsequent chapters. Temperature scaling, top-k sampling, nucleus sampling, and repetition penalties all operate within this basic framework, modifying how we select the next token from the model's predicted distribution. The KV cache, the stopping criteria, and the optimization techniques are all orthogonal to these sampling choices: you can combine any decoding strategy with KV caching, any stopping criterion, and any precision level. This modularity is one of the strengths of the autoregressive framework: the core loop is simple, and each component can be optimized independently.

Two mental models are worth carrying forward. First, think of generation as a probability chain: every generated sequence is the product of a long chain of conditional probabilities, and the quality of the output depends on how well each link in the chain was estimated. Second, think of the KV cache as earned credit: every token the model has processed represents work that need never be repeated, and the cache is the mechanism by which that earned credit is stored and reused. Together, these two ideas capture both why autoregressive generation works, grounded probability theory applied at each step, and why it is efficient in practice, reuse of prior computation via caching.

Key Parameters

When implementing autoregressive generation, these parameters have the most significant impact on behavior and performance:

  • max_new_tokens: Maximum number of tokens to generate. Set based on your application's typical output length. Shorter limits reduce latency and memory usage. Longer limits allow more complete responses but increase the risk of degenerate outputs.

  • temperature: Controls the randomness of token selection (covered in detail in the next chapter). Values below 1.0 make the distribution sharper, favoring high-probability tokens. Values above 1.0 flatten the distribution, increasing diversity. Use 0.0-0.3 for factual tasks, 0.7-1.0 for creative tasks.

  • use_cache: Whether to use KV caching for efficient generation. Should almost always be True for inference. Only disable for debugging or when memory is extremely constrained.

  • pad_token_id: Token ID used for padding in batched generation. Must be set to a valid token ID (often the EOS token ID) to avoid errors with variable-length sequences.

  • do_sample: Whether to sample from the probability distribution (True) or use greedy decoding (False). Greedy decoding is deterministic but often repetitive. Sampling introduces variety but requires tuning temperature and other sampling parameters.

  • stop_sequences: List of strings that trigger generation to stop. Useful for structured outputs, chat applications, and preventing runaway generation. Processed after each token, so slight performance overhead.

  • batch_size: Number of sequences to generate in parallel. Larger batches improve throughput but increase memory usage. Find the sweet spot based on your GPU memory and latency requirements.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about autoregressive generation and KV caching.

Autoregressive Generation Quiz

Question 1 of 100 of 10 completed
What is the key insight behind autoregressive generation?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2025autoregressivegeneration, author = {Michael Brenndoerfer}, title = {Autoregressive Generation: How GPT Produces Text}, year = {2025}, url = {https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-27} }
APAAcademic
Michael Brenndoerfer (2025). Autoregressive Generation: How GPT Produces Text. Retrieved from https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation
MLAAcademic
Michael Brenndoerfer. "Autoregressive Generation: How GPT Produces Text." 2026. Web. September 27, 2026. <https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation>.
CHICAGOAcademic
Michael Brenndoerfer. "Autoregressive Generation: How GPT Produces Text." Accessed September 27, 2026. https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation.
HARVARDAcademic
Michael Brenndoerfer (2025) 'Autoregressive Generation: How GPT Produces Text'. Available at: https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation (Accessed: September 27, 2026).
SimpleBasic
Michael Brenndoerfer (2025). Autoregressive Generation: How GPT Produces Text. https://mbrenndoerfer.com/writing/autoregressive-generation-gpt-text-generation

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.