Caching attention keys and values avoids repeating earlier projections, but one wrong position or mask can change the next-token distribution silently.
Autoregressive KV cache position and mask parity
Separate prefill from decode
A causal decoder first processes the supplied prefix. During this prefill, each position may attend only to allowed earlier positions. Afterward, each new token contributes a key and value at its own absolute position; the next query attends to all permitted cached positions. A cache stores per-layer key and value tensors, not merely the final logits. It should be associated with one exact prefix, model revision and attention configuration. Attention masks define which cached positions are visible.
Keep position and padding rules aligned
For a padded batch, the visible-token count and the physical tensor index may differ. An incremental position ID must match the model’s training and full-forward convention. Reusing a prefix cache for a different prefix, or appending at an off-by-one position, can produce fluent-looking but numerically different outputs. Check each batch member separately when lengths differ. Sliding-window and other bounded-attention layers also evict old states by a model-specific rule; simply concatenating every past key can violate that rule. Record cache length and visible-token count at each step.
Prove one-head arithmetic first
The code below compares attention over a four-position key/value array with attention over a three-position cache followed by one appended position. Both calls use the same final query and the same keys and values; exact numerical parity is expected. This is an algebra check for one attention operation, not a whole-transformer cache implementation. A full parity test must compare logits from an uncached full-prefix forward against incremental forwards across several lengths, batch padding patterns and device precisions.
Count the memory trade
Without caching, repeated decoding projects earlier tokens again and revisits the growing prefix. A cache avoids much repeated projection work, yet keys and values occupy memory proportional to layer count, retained sequence length, head count and head width. Beam search may multiply that memory by the number of live hypotheses, and cached rows must be reordered when beams are selected. Decoding policies determine how many hypotheses survive each step. Measure first-token latency, later-token latency and peak cache memory separately.
Test cache invalidation
Invalidate a cache when the model revision, tokenizer, prefix text, prefix mask, position convention or context-window policy changes. If two user requests share a prefix, sharing a mutable cache object can leak one request’s appended tokens into another; copy or isolate state before decoding. Compare fixed-prefix logits after serialization and reload, and deliberately alter one prefix token to ensure cache reuse is rejected. The project turns these checks into a release test for event sequences.
Implementation
import math
import torch
from torch.nn import functional as functional
torch.manual_seed(47)
projected_keys = torch.randn(1, 4, 8)
projected_values = torch.randn(1, 4, 8)
last_query = torch.randn(1, 1, 8)
def attended_value(query, keys, values):
scores = query @ keys.transpose(-2, -1) / math.sqrt(keys.shape[-1])
return functional.softmax(scores, dim=-1) @ values
full_result = attended_value(last_query, projected_keys, projected_values)
cached_keys = projected_keys[:, :3, :]
cached_values = projected_values[:, :3, :]
updated_keys = torch.cat((cached_keys, projected_keys[:, 3:4, :]), dim=1)
updated_values = torch.cat((cached_values, projected_values[:, 3:4, :]), dim=1)
cached_result = attended_value(last_query, updated_keys, updated_values)
torch.testing.assert_close(full_result, cached_result)
assert updated_keys.shape == (1, 4, 8)Performance and operating cost
For L retained tokens, H attention heads, width d per head and N layers, storing keys and values requires O(2NLHd) elements per request before allocator overhead. A decode step still compares its query with up to L keys, so the attention portion is not constant time. Cache use removes repeated earlier key/value projection and much full-prefix recomputation; the actual latency gain depends on model architecture, kernel and memory pressure. The simple code recomputes only one attention output and does not benchmark a transformer.
Common Mistakes
- Do not reuse a mutable cache across different request prefixes.
- Do not assume padded tensor index equals the next visible position.
- Do not compare cached and full runs with different masks or position IDs.
