In the first part of this series, we sped up the inference of our GPTlite model with caching, batching, faster attention, distillation and pruning. This post continues with techniques that generate several tokens per forward pass, use fewer bits per value, and remove the overheads around the math. As before, the code of each technique is in its own file, which imports GPTlite, and the post shows only the parts that matter, in collapsible sections. All the code is in the repository of this post.

As a reminder, inference has two phases: prefill processes the whole prompt in a single, compute-bound forward pass, and decode generates the answer one token per forward pass. Decode is memory-bandwidth-bound: each step reads all the weights of the model, and the KV cache, to do very little arithmetic with them. Most techniques in this post either read fewer bytes per generated token, or generate more tokens per byte read.

This post covers:

  1. Speculative decoding and multi-token prediction: generating several tokens per forward pass.
  2. Quantization: fewer bits per weight, activation and cached value.
  3. Compilers and kernels: kernel fusion, torch.compile and CUDA graphs.
  4. Architectures built for fast inference: mixture of experts, hybrid and diffusion models.
  5. A summary of the techniques of both posts.

Speculative decoding: more than one token per forward pass

The idea

Decode is memory-bound: a forward pass over one token costs about the same as a forward pass over a handful of tokens, because the time goes into reading the weights. Speculative decoding exploits this. A cheap method drafts \(K\) tokens, and the large target model checks all of them in a single forward pass. Every draft token the target agrees with is a token generated almost for free. When the target disagrees, we keep the tokens before the disagreement, take the target’s own token at the disagreement, and draft again. With the right acceptance rule, the output is exactly the same as with the target model alone: speculative decoding is lossless.

If each draft token is accepted with probability \(\alpha\), the expected number of tokens produced per forward pass of the target model is:

\[\frac{1 - \alpha^{K+1}}{1 - \alpha}\]

The speedup therefore depends on how often the drafts are right and on how cheap they are. Speculative decoding helps most at small batch sizes, where decode is most memory-bound; at large batch sizes, the GPU is already busy, and the extra verification work competes with useful work. The methods below differ in where the drafts come from.

Draft-model speculative decoding

The original method (Leviathan et al., Chen et al.) drafts with a small language model that uses the same tokenizer as the target, ideally a distilled version of the target, such as the student we train in the distillation section of part 1. Each iteration:

  1. the draft model generates \(K\) tokens autoregressively;
  2. the target model runs one forward pass over the context and the \(K\) draft tokens, which gives its distribution at each of the \(K\) draft positions, plus one more after the last draft;
  3. the draft tokens are checked from left to right.

The acceptance rule depends on how we decode:

  • With greedy decoding, draft tokens are accepted while they are equal to the target’s most likely token. At the first mismatch, the draft token is replaced by the target’s most likely token, and the remaining drafts are discarded. If all \(K\) drafts are accepted, the target’s prediction after the last draft is appended as a bonus token. The output is identical to greedy decoding with the target alone.
  • With sampling (speculative sampling), the draft model samples each token \(x\) from its distribution \(p\), and the target accepts it with probability \(\min\left(1, q(x) / p(x)\right)\), where \(q\) is the target’s distribution. On rejection, a replacement token is sampled from the normalized residual distribution \(\max(0, q - p)\), and the remaining drafts are discarded. This modified rejection sampling produces tokens distributed exactly as \(q\). It requires the drafts to be sampled from \(p\): mixing greedy drafts with this probabilistic rule recovers neither greedy decoding nor sampling.

In a batch, each sequence accepts a different number of tokens, so the sequences progress at different speeds, and engines track their lengths separately, as in continuous batching (see part 1).

To keep it simple, our implementation generates one sequence at a time, and the whole sequence must fit in the context window. Both models use a KV cache: after each verification, the cache of the target is truncated to the accepted tokens, and the draft model reuses the part of its own cache that is still valid.

The following code implements greedy speculative decoding in GPTlite, with the distilled model as the draft. The full code is in main_speculative_decoding.py, which checks that the output is the same as greedy decoding with the target alone:

Show code
def generate_speculative(model, prompt, n_tokens, propose, n_drafts):
  """ Greedy speculative decoding of a single sequence (prompt of shape [1, S]). At every iteration,
      propose(tokens, k) returns up to k draft tokens, and the target model checks all of them in a single
      forward pass. The cache of the target holds all the accepted tokens except the last one, which is
      processed together with the drafts. The whole sequence must fit in the context window. """
  assert prompt.size(1) + n_tokens <= model.seqlen, "the sequence must fit in the context window"
  tokens, kv_cache = prompt, None
  if prompt.size(1) > 1:
    _, kv_cache = model(prompt[:, :-1])
  stats = {"target forward passes": 0, "proposed drafts": 0, "accepted drafts": 0}
  while tokens.size(1) < prompt.size(1) + n_tokens:
    drafts = propose(tokens, min(n_drafts, model.seqlen - tokens.size(1)))  # [1, k]
    # target predictions after the last accepted token and after each draft: [1, k+1]
    logits, kv_cache = model(torch.cat([tokens[:, -1:], drafts], dim=1), kv_cache=kv_cache)
    predictions = torch.argmax(logits, dim=-1)
    n_accepted = 0  # accept the drafts while they match the predictions of the target
    while n_accepted < drafts.size(1) and drafts[0, n_accepted] == predictions[0, n_accepted]:
      n_accepted += 1
    # keep the accepted drafts, plus the target's own token at the first mismatch (or the bonus token, after
    # the last draft if all were accepted), and discard the cache entries of the rejected drafts
    tokens = torch.cat([tokens, drafts[:, :n_accepted], predictions[:, n_accepted:n_accepted+1]], dim=1)
    kv_cache = truncate_cache(kv_cache, tokens.size(1) - 1)
    stats["target forward passes"] += 1
    stats["proposed drafts"] += drafts.size(1)
    stats["accepted drafts"] += n_accepted
  return tokens[:, :prompt.size(1) + n_tokens], stats

class DraftModel:
  """ Proposes draft tokens with a small model and its own KV cache. On every call, it reuses the cache
      of the longest prefix of the accepted tokens it has already processed """

  def __init__(self, model):
    self.model = model
    self.kv_cache, self.processed = None, None  # cache and the tokens it holds

  def __call__(self, tokens, n_drafts):
    if n_drafts == 0:
      return tokens[:, :0]
    n_cached = 0
    if self.processed is not None:  # number of leading tokens processed before and still valid
      n = min(self.processed.size(1), tokens.size(1) - 1)
      mismatches = (self.processed[0, :n] != tokens[0, :n]).nonzero()
      n_cached = mismatches[0].item() if len(mismatches) > 0 else n
    kv_cache = truncate_cache(self.kv_cache, n_cached) if n_cached > 0 else None
    new_tokens, drafts = tokens[:, n_cached:], []
    for _ in range(n_drafts):  # greedy generation with the draft model
      logits, kv_cache = self.model(new_tokens, kv_cache=kv_cache)
      new_tokens = torch.argmax(logits[:, -1:], dim=-1)  # [1, 1]
      drafts.append(new_tokens)
    drafts = torch.cat(drafts, dim=1)
    self.kv_cache, self.processed = kv_cache, torch.cat([tokens, drafts[:, :-1]], dim=1)  # the last draft was not processed
    return drafts

Prompt lookup decoding

Prompt lookup decoding needs no draft model at all. It takes the last few generated tokens (an n-gram), searches for the same n-gram earlier in the prompt or in the generated text, and proposes the tokens that followed it as drafts; the verification is the same as above. It costs almost nothing, and gives large speedups when the output copies from the input, as in summarization, code editing, question answering over documents, or rewriting a previous answer. Related methods draft from other sources: lookahead decoding collects n-grams from Jacobi iterations of the model itself, and REST retrieves continuations from a datastore of text.

The following code implements prompt lookup decoding, reusing the verification loop of speculative decoding with a different source of drafts. The full code is in main_prompt_lookup.py:

Show code
def prompt_lookup(tokens, n_drafts, max_ngram=3):
  """ Drafts the tokens that followed the most recent earlier occurrence of the last n-gram of the sequence,
      trying the longest n-grams first. Returns no drafts if the n-gram never occurred before """
  sequence = tokens[0]
  for n in range(max_ngram, 0, -1):
    if sequence.size(0) <= n:
      continue
    ngram = sequence[-n:]
    windows = sequence[:-1].unfold(0, n, 1)  # all the earlier n-grams of the sequence
    matches = (windows == ngram).all(dim=1).nonzero()
    if len(matches) > 0:
      start = matches[-1].item() + n  # first token after the most recent match
      return sequence[start:start + n_drafts].unsqueeze(0)
  return tokens[:, :0]

# usage: the same verification loop as speculative decoding, with a different source of drafts
tokens, stats = generate_speculative(model, prompt, n_tokens, prompt_lookup, n_drafts)

Medusa

Medusa adds a few extra decoding heads on top of the last hidden state of the target model: the original output head predicts the next token, and head \(i\) predicts the token \(i+1\) positions ahead. The top candidates of each head are combined into a tree of possible continuations, and all of them are verified in a single forward pass with tree attention, a mask that lets each candidate attend only to its own ancestors in the tree. The longest accepted path is kept. Training only the heads, with the model frozen (Medusa-1), gives speedups above 2.2 times; training the heads together with the model (Medusa-2) reaches 2.3 to 3.6 times. Medusa needs no separate draft model, but each head guesses its token independently, without seeing the tokens guessed by the other heads, so the accuracy drops quickly for later positions. For sampling, Medusa also proposes a faster typical acceptance rule, which is not exactly lossless.

EAGLE

EAGLE drafts at the level of features rather than tokens. Its draft model is a single lightweight transformer layer that takes the top-layer hidden states (features) of the target, together with the embeddings of the tokens sampled so far, predicts the next feature, and turns it into a token with the target’s own output layer. Feature sequences are more regular than token sequences, so these drafts are much more accurate than Medusa’s independent heads. EAGLE-2 shapes the draft tree dynamically, expanding the branches where the draft model is confident. EAGLE-3 (NeurIPS 2025) noticed that EAGLE barely improved with more training data, and traced this to its feature prediction objective. It drops feature prediction and predicts tokens directly, fuses low-, mid- and high-level features of the target instead of the top layer only, and uses training-time test: during training, the draft model is fed its own predictions, to simulate multi-step drafting. EAGLE-3 reaches speedups of up to 6.5 times, about 1.4 times more than EAGLE-2, and 1.38 times higher throughput at a batch size of 64 in SGLang. EAGLE-style drafting is supported by vLLM, SGLang and TensorRT-LLM.

Self-speculative decoding: LayerSkip

LayerSkip uses the early layers of the target model as the draft. The model is trained with layer dropout, which skips the later layers more often, and with an early exit loss, which trains every layer to make good predictions through the shared output head. At inference, the draft exits after the first \(E\) layers, and the verification runs only the remaining layers, reusing the computation and the cache of the first \(E\). There is no second model to store, and the reported speedups reach up to about 2 times on summarization and coding tasks.

Multi-token prediction

Multi-token prediction (MTP) trains the model itself, during pre-training, to predict several future tokens:

  • Gloeckle et al. (Meta, 2024) add \(n\) independent output heads on a shared trunk, each predicting one of the next \(n\) tokens. Besides improving sample efficiency, notably on code, the extra heads serve as drafts for self-speculative decoding, making the inference of a model trained to predict 4 tokens up to 3 times faster.
  • DeepSeek-V3 uses sequential MTP modules: each module combines the representation of the previous depth with the embedding of the next token, and runs one transformer block, so every prediction stays conditioned on the previous ones. MTP is an extra training objective, and at inference the modules can either be discarded or used as drafts: the second predicted token is accepted 85-90% of the time, giving 1.8 times more tokens per second.

Medusa and EAGLE add drafting to an existing model, while MTP builds it in during pre-training.

Summary of speculative decoding methods

Method Source of the drafts Extra training Extra memory Reported speedup
Draft model a separate small model none, or distilling a draft a second model about 2-3x
Prompt lookup n-grams of the context none none large on copy-heavy tasks
Medusa extra heads the heads (or joint fine-tuning) small 2.2-3.6x
EAGLE-3 a feature-level draft layer the draft layer small up to 6.5x
LayerSkip early layers of the model special fine-tuning none up to about 2x
Multi-token prediction MTP heads or modules pre-training small 1.8x (DeepSeek-V3) to 3x

The reported speedups come from the respective papers, with different models, tasks and hardware, and they shrink as the batch size grows.

Quantization

Quantization stores numbers with fewer bits. Weights take less memory, fewer bytes are read at every decode step (which is what limits decode), and low-precision tensor cores perform more operations per second. The cost is a loss of precision, which must stay small enough not to hurt the quality of the model.

Number formats

A floating point number has a sign bit, exponent bits that set its range, and mantissa (fraction) bits that set its precision. The picture below compares the three most common formats: FP32, FP16 and BF16. BF16 (brain floating point) keeps the 8 exponent bits of FP32 and only 7 mantissa bits, so it covers the same range as FP32 with the memory and speed of a 16-bit format. FP16 has more precision but a much smaller range (5 exponent bits), so large values can overflow. This is why BF16 is the default format of LLMs today.

Recent GPUs added smaller formats:

Format Bits Sign / exponent / mantissa bits Largest value Typical use
FP32 32 1 / 8 / 23 about 3.4 × 10³⁸ reference precision
FP16 16 1 / 5 / 10 65,504 weights and activations
BF16 16 1 / 8 / 7 about 3.4 × 10³⁸ default for LLM weights and activations
FP8 E4M3 8 1 / 4 / 3 448 weights, activations and KV cache (Hopper, Ada and newer GPUs)
FP8 E5M2 8 1 / 5 / 2 57,344 values that need more range than precision
FP4 E2M1 4 1 / 2 / 1 6 weights and activations with micro-scaling (Blackwell GPUs)
INT8 8 integers from -128 to 127 127 weights, activations and KV cache
INT4 4 integers from -8 to 7 7 weights, with group-wise scales

Four bits represent only 16 values: FP4 E2M1 can only represent \(\pm\{0, 0.5, 1, 1.5, 2, 3, 4, 6\}\). Such small formats only work with fine-grained scaling. In the micro-scaling (MX) formats, every block of 32 consecutive values shares an 8-bit power-of-two scale (MXFP8, MXFP6 and MXFP4), and NVIDIA’s NVFP4 uses blocks of 16 values with an FP8 scale, plus one FP32 scale per tensor. Blackwell tensor cores apply these scales in hardware. OpenAI’s gpt-oss models, for example, ship their mixture-of-experts weights in MXFP4.

How quantization works

To quantize a tensor \(x\) to \(b\)-bit integers, we choose a scale \(s\) (and optionally a zero-point \(z\)), then round and clip:

\[x_q = \text{clip}\left(\text{round}\left(\frac{x}{s}\right) + z, \; q_\text{min}, \; q_\text{max} \right), \quad\quad \hat{x} = s \, (x_q - z)\]

where \(\hat{x}\) is the dequantized approximation of \(x\). In symmetric (absmax) quantization, \(z = 0\) and \(s = \max \lvert x \rvert / q_\text{max}\), e.g., with \(q_\text{max} = 127\) for INT8. Asymmetric quantization maps the range \([\min x, \max x]\) onto the full integer range, which fits skewed distributions better.

The granularity of the scales matters as much as the number of bits. A single scale per tensor is cheap but fragile. One scale per output channel (a row of the weight matrix), per group of, e.g., 128 consecutive weights, or per token for the activations, follows the data much more closely, at the cost of storing more scales. Group-wise scales are the standard for 4-bit weights.

The main enemy is outliers: a single large value stretches the scale, and all the small values collapse onto a few quantization levels. Weights are well-behaved and easy to quantize. Activations are not: LLM.int8() showed that large transformers develop a few outlier feature dimensions, with values much larger than the rest, that break naive 8-bit quantization of the activations.

Quantization can be applied after training with a small calibration dataset (post-training quantization, PTQ), which is fast and is what we do here, or simulated during training so that the model learns to tolerate it (quantization-aware training, QAT), which gives the best results at 4 bits and below. Finally, there are four things to quantize: the weights, the activations, the KV cache and the attention computation.

Weight-only quantization

Storing the weights in 4 or 8 bits while computing in 16 bits (e.g., W4A16: 4-bit weights, 16-bit activations) speeds up decode almost in proportion to the bytes saved, because decode is memory-bound: 4-bit weights are read four times faster than BF16 ones. The kernel must dequantize the weights on the fly, in registers, right before the multiplication; dequantizing the whole matrix to memory first would cancel the gain. Prefill, which is compute-bound, barely benefits.

Rounding each weight to the nearest level (round-to-nearest) works well at 8 bits, but loses too much accuracy at 4 bits. Two methods made 4-bit weights practical:

  • GPTQ quantizes a weight matrix one column at a time and, after each column, updates the columns not yet quantized to compensate for the error just introduced. The update uses second-order information: the inverse of the Hessian \(H = 2XX^T\) of the layer’s reconstruction error, computed from calibration inputs \(X\). GPTQ quantizes 175B-parameter models to 3-4 bits in about four GPU hours.
  • AWQ (Activation-aware Weight Quantization) observes that a small fraction of the weights matters much more than the rest: those that multiply large activations. Instead of keeping them in high precision, it multiplies the salient input channels of the weights by a factor before quantization (and divides the corresponding activations by the same factor, which can be folded into the previous operation), so that they are rounded more accurately. The factors are searched on calibration data.

Other common formats are NF4, from QLoRA, whose levels are placed for normally-distributed weights, and the GGUF formats of llama.cpp, for CPUs and edge devices.

To see how it works, we write our own quantizer: symmetric INT8 quantization with one scale per output channel, and INT4 quantization with one scale per group of 64 weights, packing two 4-bit weights per byte. It quantizes the linear layers of the transformer blocks of GPTlite, and we compare the memory, the time per token and the validation loss against the FP32 and BF16 models. With one scale per output channel, the INT8 layer can apply the scale after the matrix multiplication, while the INT4 layer must dequantize the weights first. In plain PyTorch, both convert the weights back to floating point in memory before each multiplication, which saves memory but not time: actual speedups require kernels that dequantize in registers, such as those of torchao or Marlin, or a compiler that fuses the conversion into the multiplication.

The following code implements our own weight-only INT8 and INT4 quantization of GPTlite. The full code is in main_quantization.py:

Show code
class QuantizedLinear(nn.Module):
  """ Linear layer with symmetric (absmax) weight-only quantization: INT8 with one scale per output channel,
      or INT4 with one scale per group of group_size consecutive weights of a row, packing two weights per byte.
      The weights are dequantized on the fly, and the computation runs in the precision of the input """

  def __init__(self, linear, bits, group_size=64):
    super().__init__()
    assert bits in (4, 8), "only 8 and 4 bits are supported"
    W = linear.weight.data  # [out_features, in_features]
    self.bits, (self.out_features, self.in_features) = bits, W.shape
    self.bias = None if linear.bias is None else nn.Parameter(linear.bias.data.clone(), requires_grad=False)
    q_max = 2 ** (bits - 1) - 1  # 127 for INT8, 7 for INT4
    if bits == 8:
      scale = W.abs().amax(dim=1, keepdim=True) / q_max  # [out_features, 1]
    else:
      self.group_size = group_size if self.in_features % group_size == 0 else self.in_features
      W = W.view(self.out_features, -1, self.group_size)   # [out_features, n_groups, group_size]
      scale = W.abs().amax(dim=2, keepdim=True) / q_max    # [out_features, n_groups, 1]
    scale = scale.clamp(min=1e-8)
    q = torch.clamp(torch.round(W / scale), -q_max - 1, q_max).to(torch.int8)
    if bits == 4:  # shift the values from [-8, 7] to [0, 15] and pack two of them per byte
      q = (q + 8).to(torch.uint8).view(self.out_features, -1)
      q = q[:, 0::2] | (q[:, 1::2] << 4)
    self.register_buffer('qweight', q)
    self.register_buffer('scale', scale)

  def dequantize(self):
    """ Weight matrix in the precision of the scales """
    if self.bits == 8:
      return self.qweight.to(self.scale.dtype) * self.scale
    low, high = (self.qweight & 0x0F).to(torch.int8) - 8, (self.qweight >> 4).to(torch.int8) - 8
    q = torch.stack([low, high], dim=-1).view(self.out_features, -1, self.group_size)  # unpack and interleave
    return (q.to(self.scale.dtype) * self.scale).view(self.out_features, self.in_features)

  def forward(self, x):
    if self.bits == 8:
      # with one scale per output channel, the scale can be applied after the matrix multiplication
      out = F.linear(x, self.qweight.to(x.dtype)) * self.scale.view(-1).to(x.dtype)
      return out if self.bias is None else out + self.bias
    return F.linear(x, self.dequantize().to(x.dtype), self.bias)

def quantize(model, bits, group_size=64):
  """ Replaces the linear layers of the transformer blocks by quantized ones. The embeddings and the
      output layer stay in full precision """
  linears = [(parent, name, child) for block in model.blocks for parent in block.modules()
             for name, child in parent.named_children() if isinstance(child, nn.Linear)]
  for parent, name, child in linears:
    setattr(parent, name, QuantizedLinear(child, bits, group_size))
  return model

Weight and activation quantization

To make the matrix multiplication itself faster, both of its inputs must be in low precision, so that the GPU can use its INT8, FP8 or FP4 tensor cores, which are 2 to 4 times faster than the BF16 ones. This also speeds up the compute-bound prefill and large-batch decode. The difficulty is the outliers of the activations:

  • LLM.int8() keeps the few outlier feature dimensions in 16 bits and runs everything else in INT8.
  • SmoothQuant moves the difficulty from the activations to the weights. Dividing the \(j\)-th channel of the activations by a factor \(s_j\), and multiplying the \(j\)-th input channel of the weights by the same factor, does not change the output of the layer but tames the outliers. With \(s_j = \max \lvert X_j \rvert^\alpha / \max \lvert W_j \rvert^{1-\alpha}\) and typically \(\alpha = 0.5\), both the weights and the activations become easy to quantize to INT8 (W8A8).
  • FP8 weights and activations are the most common choice on Hopper and Blackwell GPUs: with per-channel or per-token scales, they are nearly lossless for most LLMs.
  • Rotations (QuaRot, SpinQuant) multiply the weights and activations by orthogonal matrices, such as Hadamard matrices. Since \(RR^T = I\), the output does not change, but the rotation spreads each outlier across all channels. This makes 4-bit weights, activations and KV cache possible (W4A4KV4).
  • FP4 (NVFP4 and MXFP4) on Blackwell GPUs uses the micro-scaled formats described above, for weights and activations.

KV cache quantization

At long contexts and large batches, the KV cache can be larger than the model weights. Storing it in FP8 or INT8 halves its size and the bytes read at each decode step, with almost no loss of quality; most serving engines enable it with a single configuration flag. More aggressive methods go down to 2 bits: KIVI quantizes the keys per channel, since their outliers live in fixed channels, and the values per token, and keeps the most recent tokens in full precision. Quantized attention kernels, such as SageAttention, apply the same ideas to the attention computation itself (see part 1).

Which one to use

A reasonable order: start in BF16; quantize the KV cache to FP8 for long contexts; use FP8 weights and activations on Hopper or newer GPUs; use 4-bit weights (AWQ or GPTQ) when memory is the constraint or the batch size is small; and use FP4 on Blackwell GPUs, ideally with a model trained or fine-tuned for it (QAT). Always measure the quality as well as the speed: the perplexity on held-out data, and the tasks you care about.

Compilers and kernels

Why kernels matter

Every PyTorch operation launches one or more GPU kernels. Each launch costs a few microseconds of CPU and driver time, and each kernel reads its inputs from GPU memory and writes its outputs back. For large models and batches, these costs hide behind the math. For small models and batches, like GPTlite, they dominate: a decode step launches hundreds of small kernels, and the GPU spends much of its time idle, waiting for the CPU to launch the next kernel, or for the memory round-trips between kernels.

Kernel fusion and torch.compile

Kernel fusion merges consecutive operations into a single kernel, so that intermediate results stay in registers or on-chip memory instead of going through GPU memory. Typical candidates are the chains of element-wise operations after a matrix multiplication (bias, activation function, dropout, residual addition) and the normalization layers. FlashAttention is an extreme example of fusion.

torch.compile fuses kernels automatically: TorchDynamo captures the graph of the model from the Python bytecode, and TorchInductor generates fused kernels, written in Triton for GPUs. It is a one-line change: model = torch.compile(model). A change in the shapes of the inputs triggers a recompilation, which is one more reason to use a static KV cache during generation.

CUDA graphs

Even after fusion, a decode step launches many kernels from Python, one after the other. A CUDA graph records the whole sequence of kernel launches once, and replays it with a single launch, removing almost all the CPU overhead. The recording fixes the shapes and memory addresses of all tensors, so the decode step must use static shapes: a preallocated KV cache of maximum length, with attention over the full buffer and a mask for the positions not yet filled. torch.compile(model, mode="reduce-overhead") uses CUDA graphs automatically. Combining these techniques with quantization and speculative decoding, the PyTorch team’s gpt-fast made the decoding of Llama-7B almost 10 times faster, in plain PyTorch. Going one step further, recent work fuses a whole forward pass into a single persistent megakernel, removing the gaps between kernels altogether for low-latency decoding.

The following code implements a static KV cache for GPTlite, and compiles its decode step with torch.compile, which records and replays CUDA graphs on a GPU. The full code is in main_cuda_graphs.py, which checks that the outputs are the same as with the dynamic KV cache:

Show code
class GPTlite_StaticCache(nn.Module):
  """ GPTlite with a static KV cache: buffers preallocated for max_seqlen tokens, where every step writes the
      keys and values of the new tokens in place. Attention runs over the whole buffer, with a mask for the
      positions not filled yet, so all tensors have the same shape at every decode step """

  def __init__(self, model, batch_size, max_seqlen):
    super().__init__()
    self.model = model
    n_heads, d_head = model.blocks[0].mha.n_heads, model.blocks[0].mha.d_head
    shape = (len(model.blocks), batch_size, n_heads, max_seqlen, d_head)
    weight = model.token_embedding.weight
    self.register_buffer('k_cache', torch.zeros(shape, dtype=weight.dtype, device=weight.device))
    self.register_buffer('v_cache', torch.zeros(shape, dtype=weight.dtype, device=weight.device))
    self.register_buffer('positions', torch.arange(max_seqlen, device=weight.device))

  def forward(self, tokens, input_pos):
    """ tokens: [B, T] new tokens, input_pos: [T] their positions. Returns the logits [B, T, vocab_size] """
    model, (B, T) = self.model, tokens.shape
    x = model.token_embedding(tokens) + model.position_embedding(input_pos)
    mask = input_pos[:, None] >= self.positions[None, :]  # [T, max_seqlen]: attend to the positions up to itself
    for i, block in enumerate(model.blocks):
      mha, H, D = block.mha, block.mha.n_heads, block.mha.d_head
      h = block.ln1(x)
      q = mha.query_proj(h).view(B, T, H, D).transpose(1, 2)  # [B, H, T, D]
      k = mha.key_proj(h).view(B, T, H, D).transpose(1, 2)
      v = mha.value_proj(h).view(B, T, H, D).transpose(1, 2)
      self.k_cache[i, :, :, input_pos] = k  # write the new keys and values in place
      self.v_cache[i, :, :, input_pos] = v
      out = F.scaled_dot_product_attention(q, self.k_cache[i], self.v_cache[i], attn_mask=mask)
      x = x + mha.out_proj(out.transpose(1, 2).reshape(B, T, H * D))
      x = x + block.ffwd(block.ln2(x))
    return model.fc_out(model.ln(x))

def generate_static(model, prompt, n_tokens, decode_step):
  """ Greedy generation with the static cache: the prompt is processed by the model (prefill), and every
      new token by decode_step, which can be the same model or a compiled version of it """
  P = prompt.size(1)
  logits = model(prompt, torch.arange(P, device=prompt.device))
  tokens = [torch.argmax(logits[:, -1], dim=-1, keepdim=True)]
  for pos in range(P, P + n_tokens - 1):
    logits = decode_step(tokens[-1], torch.tensor([pos], device=prompt.device))
    tokens.append(torch.argmax(logits[:, -1], dim=-1, keepdim=True))
  return torch.cat([prompt] + tokens, dim=1)

# compile the decode step: with mode="reduce-overhead", torch.compile records and replays CUDA graphs
model_static = GPTlite_StaticCache(model, batch_size, seqlen)
decode_compiled = torch.compile(model_static, mode="reduce-overhead", fullgraph=True)
tokens = generate_static(model_static, prompt, n_tokens, decode_compiled)

Inference engines

Production deployments rarely hand-write these optimizations: inference engines combine most of the techniques in this post.

  • vLLM and SGLang are open-source serving engines with continuous batching, PagedAttention or RadixAttention, prefix caching, speculative decoding and quantization.
  • TensorRT-LLM is NVIDIA’s engine, with graph optimizations, kernel fusion, in-flight batching, paged KV caches, and FP8 and FP4 quantization.
  • DeepSpeed-FastGen and DeepSpeed-Inference are DeepSpeed’s serving and inference engines.
  • llama.cpp runs models on CPUs and edge devices, with its own quantized formats (GGUF).
  • ONNX Runtime runs models exported to the portable ONNX format on many kinds of hardware; it is covered in a separate post.

Architectures built for fast inference

Some of the largest gains come from the architecture of the model itself. These choices are made before pre-training, so we describe them without implementing them in GPTlite.

Mixture of experts

A mixture-of-experts (MoE) layer replaces the MLP of a block with many smaller MLPs (experts) and a router that sends each token to the top few experts. Only a fraction of the parameters is used per token: Mixtral 8x7B has 47B parameters but uses 13B per token, and DeepSeek-V3 uses 37B of its 671B. The compute per token follows the active parameters, but the memory follows the total parameters, since all experts must be loaded. At small batch sizes, each decode step reads only the weights of the selected experts, which makes MoE models fast; at large batch sizes, most experts are used at every step. Serving large MoE models relies on expert parallelism and on balancing the load between experts.

Hybrid attention models

Models that mix a few full attention layers with many linear attention or state space layers, such as Kimi Linear, Qwen3-Next, Jamba and MiniMax-01, have a much smaller KV cache and a lower cost per token at long contexts (see the linear attention section of part 1).

Diffusion language models

Diffusion language models do not generate text from left to right. They start from a fully masked sequence and denoise it over a number of steps, predicting many tokens in parallel at every step, as in LLaDA. Since the number of steps can be much smaller than the number of tokens, they can generate faster than autoregressive models, and commercial diffusion LMs such as Inception Labs’ Mercury and Google’s Gemini Diffusion advertise very high generation speeds. They are still an active research area: their bidirectional attention makes KV caching harder, and their quality and controllability are still catching up with autoregressive models.

Summary

The table below summarizes the techniques of both posts.

Technique What it speeds up Lossless? Needs training? When it helps most
KV cache decode yes no always
Prefix caching prefill (time to first token) yes no shared prompts, multi-turn chat
Semantic caching whole requests no no many similar queries
Continuous batching, PagedAttention throughput yes no serving many users
Chunked prefill, disaggregation latency stability, throughput yes no large deployments
FlashAttention, FlashDecoding attention yes no long contexts
MQA, GQA, MLA KV cache size and reads no (changes the model) yes long contexts, large batches
Sparse attention (NSA, DSA) attention no yes very long contexts
Linear and hybrid attention attention, KV cache no (changes the model) yes (pre-training) very long contexts
Speculative decoding decode latency yes (with exact acceptance) depends on the method small batches, latency-sensitive applications
Quantization memory, bandwidth, compute no (small loss) no (PTQ) or yes (QAT) memory-bound decode, large models
Distillation and pruning everything (smaller model) no yes when a smaller model is good enough
Kernel fusion, torch.compile, CUDA graphs overheads yes no small models, small batches
Mixture of experts compute per token no (changes the model) yes (pre-training) large-scale serving

The following code benchmarks the methods of both posts on GPTlite, measuring the time to first token, the time per output token, the throughput, the peak memory and the validation loss. It reuses the code of each method, generates a single sequence, as in latency-sensitive serving, and loads the checkpoints created by main_distillation.py and main_gqa.py. The full code is in main_benchmark.py:

Show code
def benchmark(generate_fn, prompt, n_tokens):
  """ Time to first token (generating a single token), time per output token, throughput and peak GPU memory """
  generate_fn(prompt, n_tokens)  # warm-up, which also compiles the compiled models
  if torch.cuda.is_available():
    torch.cuda.reset_peak_memory_stats()
  with Timer() as first:
    generate_fn(prompt, 1)
  with Timer() as total:
    generate_fn(prompt, n_tokens)
  memory = f"{torch.cuda.max_memory_allocated() / 2**20:.1f}" if torch.cuda.is_available() else "-"
  time_per_token = (total.elapsed - first.elapsed) / (n_tokens - 1)
  return first.elapsed, time_per_token, prompt.numel() // prompt.size(1) * n_tokens / total.elapsed, memory

# name: (generation function, model used to compute the validation loss)
methods = {
  "No KV cache": (lambda p, n: generate(model, p, n, seqlen), model),
  "KV cache": (lambda p, n: generate_kvcache(model_kvcache, p, n, seqlen), model_kvcache),
  "Speculative decoding (distilled draft)": (lambda p, n: generate_speculative(model_kvcache, p, n, DraftModel(model_distilled), n_drafts)[0], model_kvcache),
  # ...
}
for name, (generate_fn, model_obj) in methods.items():
  ttft, tpot, throughput, memory = benchmark(generate_fn, prompt, n_tokens)
  loss = validation_loss(model_obj, valid_data)

Further reading