All articles LLM Optimization

KV Cache Compression & Context Window Extension: YaRN, RoPE & PagedAttention

Managing gigabyte KV caches for 1M+ token contexts: vLLM PagedAttention, Rotary Position Embeddings (RoPE), YaRN frequency scaling, and Grouped-Query Attention.

Extending Large Language Model context windows from 4,024 to 1,000,000+ tokens introduces massive Key-Value (KV) cache memory bottlenecks. For a 70B parameter model processing 128k context lengths, storing FP16 KV cache tensors requires over 16GB of VRAM per single request. PagedAttention (vLLM), Grouped-Query Attention (GQA), and YaRN (Yet Another RoPE Extension) solve this memory wall through virtual memory pagination and position embedding frequency scaling.

KV cache memory fragmentation limits batch sizes; PagedAttention partitions KV memory into virtual pages with zero waste.

The KV cache memory bottleneck

During autoregressive decoding, Key and Value vectors for past tokens are cached to avoid recomputing them. The total KV cache memory requirement per token is:

Memory = 2 * (2 * num_layers * num_kv_heads * head_dim) bytes per token
For Llama 3 70B (80 layers, 8 KV heads, 128 head_dim):
Memory per token = 2 * (2 * 80 * 8 * 128) = 327,680 bytes = 327.68 KB / token!

For 128,000 tokens: 128,000 * 327.68 KB = 41.9 GB VRAM per single sequence!

Grouped-Query Attention (GQA) vs Multi-Head Attention (MHA)

Standard Multi-Head Attention assigns 1 Key/Value head for every Query head (e.g. 64 Q heads, 64 KV heads). Grouped-Query Attention (Ainslie et al., 2023) groups multiple Query heads to share a single KV head (e.g. 64 Q heads sharing 8 KV heads):

Multi-Head Attention (MHA):     64 Q heads <---> 64 KV heads  (1:1 ratio - Heavy KV Memory)
Multi-Query Attention (MQA):    64 Q heads <---> 1  KV head   (64:1 ratio - Minimal Memory, lower quality)
Grouped-Query Attention (GQA):  64 Q heads <---> 8  KV heads   (8:1 ratio - Optimal balance!)

GQA reduces KV cache memory footprint by 8x with virtually identical model accuracy.

PagedAttention & virtual memory block tables

Inspired by virtual memory paging in operating systems, PagedAttention (Kwon et al., 2023) partitions KV cache into fixed-size physical blocks (e.g. 16 tokens per block) allocated dynamically on demand:

Logical KV Cache (Tokens 0..31):
+-------------------------+-------------------------+
| Block 0 (Tokens 0..15)  | Block 1 (Tokens 16..31) |
+-------------------------+-------------------------+
            |                         |
            v (Block Table Mapping)   v
+-------------------------+-------------------------+
| Physical Page #47       | Physical Page #12       |
+-------------------------+-------------------------+

This eliminates external memory fragmentation, reducing memory waste from ~60% down to under 1%!

Rotary Position Embedding (RoPE) & YaRN frequency scaling

To extend context windows beyond pre-training limits without full fine-tuning, YaRN (Peng et al., 2023) rescales Rotary Position Embedding (RoPE) frequencies by interpolating high frequencies and extrapolating low frequencies:

RoPE Transformation for 2D sub-vector (x1, x2) at position m:
R_theta,m * [x1, x2]^T = [ x1 * cos(m * theta) - x2 * sin(m * theta),
                           x1 * sin(m * theta) + x2 * cos(m * theta) ]

Python implementation of RoPE rotary embedding transformation

import torch

def apply_rotary_pos_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
    """
    Applies Rotary Position Embeddings (RoPE) to Query / Key tensors.
    x shape: (batch_size, seq_len, num_heads, head_dim)
    """
    # Split head dimension in half
    d = x.shape[-1]
    x1 = x[..., :d//2]
    x2 = x[..., d//2:]
    
    # Rotate sub-vectors
    rotated_x = torch.cat((-x2, x1), dim=-1)
    
    # Apply Euler rotation
    return (x * cos) + (rotated_x * sin)

# Example execution
B, T, H, D = 2, 1024, 8, 128
q = torch.randn(B, T, H, D)

freqs = 1.0 / (10000 ** (torch.arange(0, D, 2).float() / D))
t = torch.arange(T).float()
freqs_matrix = torch.outer(t, freqs)
emb = torch.cat((freqs_matrix, freqs_matrix), dim=-1)

cos_emb = emb.cos().view(1, T, 1, D)
sin_emb = emb.sin().view(1, T, 1, D)

q_rotated = apply_rotary_pos_emb(q, cos_emb, sin_emb)
print("Rotated Query Tensor Shape:", q_rotated.shape)

Context window scaling checklist

  1. Deploy Grouped-Query Attention (GQA): 8x reduction in KV cache size compared to standard MHA.
  2. Use vLLM with PagedAttention: Virtual memory pagination eliminates memory fragmentation and enables 2x-4x larger batch sizes.
  3. Apply YaRN Frequency Scaling: Smoothly extend pre-trained 4k models to 128k context windows with minimal fine-tuning tokens.
← Back to all articles