KV-Cache: The Hidden Tax

At 128K context, the cache alone fills an 80 GB GPU — room for exactly one user.

node
intermediate
Discover that KV-cache memory — not model weights, not compute — determines how many users you can serve concurrently. Sweep batch size and context length to find the real OOM boundary.

The Question

You deploy Llama-3 8B on an H100. The model weights take 16 GB. You have 64 GB left. Surely you can serve dozens of users concurrently?

Not if they have long contexts. Every active user requires a KV-cache that grows linearly with sequence length. At 128K context, a single user’s cache can consume the entire remaining memory. This tutorial shows you exactly where the real memory wall lives and how to push it back.

NotePrerequisites

Complete Tutorial 1: The Memory Wall and Tutorial 2: Two Phases, One Request. You should understand memory-bound vs. compute-bound regimes and the two-phase LLM serving model.

NoteWhat You Will Learn
  • Calculate the KV-cache size for any model, sequence length, and batch size
  • Identify the OOM boundary where KV-cache exhausts GPU memory
  • Explain why context length — not model size — is the binding memory constraint in serving
  • Compare static batching vs. paged attention for maximizing concurrent users
TipBackground: What Is the KV-Cache?

During LLM decoding, every attention layer stores Key and Value matrices for all tokens generated so far. If you have studied data structures, this is memoization applied to the attention mechanism: store computed results instead of recomputing them. The names come from a database-style lookup: the Query is what you search for, the Key is what you match against, and the Value is what you retrieve. Without this cache, the model would need to recompute attention over the entire context at every step — quadratic cost. The KV-cache trades memory for compute:

Factor Effect on KV-Cache
More layers Linear growth (one K + one V per layer)
Longer context Linear growth (one entry per token)
More users (batch) Linear growth (independent cache per user)
Lower precision Proportional reduction (INT8 = half of FP16)

The formula: KV-cache = 2 x layers x kv_heads x head_dim x seq_len x batch x bytes_per_element. At short contexts this is negligible. At long contexts it dominates everything.

Note on GQA (Grouped Query Attention): Modern architectures like Llama-3 use GQA, where kv_heads < num_heads. Llama-3 8B has 32 attention heads but only 8 KV-heads, reducing KV-cache by 4× compared to standard multi-head attention. Using num_heads instead of kv_heads in the formula is a common source of 4× overestimates.


1. Setup

import mlsysim
from mlsysim.solvers import ServingModel

2. Single-User Baseline: Where Does the Memory Go?

Let’s start with a single user at a modest 2K context and see how memory breaks down:

from mlsysim.solvers import ServingModel

model = mlsysim.Models.Language.Llama3_8B
hardware = mlsysim.Hardware.Cloud.H100
solver = ServingModel()

# Single user, 2K context — the easy case
r = solver.solve(model=model, hardware=hardware, seq_len=2048, batch_size=1, precision="fp16")

from mlsysim.show import table, info

info("Memory Breakdown",
     Model_weights=r.model_weights_size,
     KV_cache_1_user=r.kv_cache_size,
     Total_memory=r.total_memory_required,
     Memory_utilization=f"{r.memory_utilization:.1%}",
     KV_as_pct_of_total=f"{r.kv_cache_size / r.total_memory_required * 100:.1f}%")
── Memory Breakdown ────────────────────────
Model weights:       16.06 GB
KV cache 1 user:     0.268 GB
Total memory:        16.33 GB
Memory utilization:  19.0%
KV as pct of total:  1.6 dimensionless%

At 2K context with one user, the KV-cache is tiny — a rounding error compared to the model weights. This is why many engineers assume memory pressure comes from model size. They are about to be surprised.


3. Batch Size Sweep: The Concurrency Wall

Now let’s add users. Each concurrent user needs their own KV-cache. Watch memory utilization climb:

rows = []
for batch in [1, 4, 8, 16, 32, 64, 128]:
    r = solver.solve(
        model=model, hardware=hardware,
        seq_len=2048, batch_size=batch, precision="fp16"
    )
    rows.append([batch, r.kv_cache_size, r.total_memory_required,
                 f"{r.memory_utilization:.1%}", "OK" if r.feasible else "OOM"])

table(["Batch", "KV-Cache", "Total", "Util", "Feasible"], rows)
Batch  KV-Cache     Total   Util  Feasible
──────────────────────────────────────────
1      0.268 GB  16.33 GB  19.0%        OK
4       1.07 GB  17.13 GB  19.9%        OK
8       2.15 GB  18.21 GB  21.2%        OK
16      4.29 GB  20.35 GB  23.7%        OK
32      8.59 GB  24.65 GB  28.7%        OK
64     17.18 GB  33.24 GB  38.7%        OK
128    34.36 GB  50.42 GB  58.7%        OK

At 2K context, you can fit many users. The KV-cache per user is small enough that batch size scales comfortably. But this picture changes dramatically when we extend the context.


4. Context Length Sweep: The Real Memory Wall

Fix batch size at 8 users and sweep context length from 512 tokens to 128K. This is where the hidden tax reveals itself:

rows = []
for ctx in [512, 2048, 4096, 8192, 16384, 32768, 65536, 131072]:
    r = solver.solve(
        model=model, hardware=hardware,
        seq_len=ctx, batch_size=8, precision="fp16"
    )
    rows.append([ctx, r.kv_cache_size, r.model_weights_size,
                 r.total_memory_required, f"{r.memory_utilization:.1%}",
                 "OK" if r.feasible else "OOM"])

table(["Context", "KV-Cache", "Weights", "Total", "Util", "Status"], rows)
Context  KV-Cache   Weights     Total    Util  Status
─────────────────────────────────────────────────────
512      0.537 GB  16.06 GB  16.60 GB   19.3%      OK
2048      2.15 GB  16.06 GB  18.21 GB   21.2%      OK
4096      4.29 GB  16.06 GB  20.35 GB   23.7%      OK
8192      8.59 GB  16.06 GB  24.65 GB   28.7%      OK
16384    17.18 GB  16.06 GB  33.24 GB   38.7%      OK
32768    34.36 GB  16.06 GB  50.42 GB   58.7%      OK
65536    68.72 GB  16.06 GB  84.78 GB   98.7%      OK
131072   137.4 GB  16.06 GB  153.5 GB  178.7%     OOM
ImportantKey Insight

KV-cache grows linearly with sequence length and batch size. It is the hidden memory consumer that determines your maximum concurrent users — not model size, not compute, but cache state. At 2K context, the cache is negligible. At 128K context, a single user’s cache can exceed the model weights. The same 80 GB GPU that serves 64 users at short context can serve exactly one user at long context. The “context length” on the model card is not a feature — it is a memory bill.

Now let’s see what happens when we try to serve even a single user at 128K:

# Single user at 128K context — the extreme case
r_long = solver.solve(
    model=model, hardware=hardware,
    seq_len=131072, batch_size=1, precision="fp16"
)

info("Single User @ 128K Context",
     Context="131,072 tokens (128K)",
     Model_weights=r_long.model_weights_size,
     KV_cache=r_long.kv_cache_size,
     Total=r_long.total_memory_required,
     Feasible=str(r_long.feasible),
     KV_as_pct_of_total=f"{r_long.kv_cache_size / r_long.total_memory_required * 100:.0f}%")
── Single User @ 128K Context ──────────────
Context:             131,072 tokens (128K)
Model weights:       16.06 GB
KV cache:            17.18 GB
Total:               33.24 GB
Feasible:            True
KV as pct of total:  52 dimensionless%

5. Paged Attention: Pushing Back the Wall

So the KV-cache fills memory fast, and at long contexts you hit OOM with just a handful of users. Is the only option to buy more memory? No — the allocation strategy itself is wasting space. Most sequences do not actually use the maximum context length, yet static batching reserves memory for the worst case.

Static batching allocates contiguous memory for the maximum sequence length, wasting space on incomplete sequences. PagedAttention (from vLLM) allocates KV-cache in small, fixed-size pages — exactly like how an operating system uses virtual memory paging to avoid physical memory fragmentation. Just as the OS maps virtual pages to physical frames on demand, PagedAttention maps KV-cache blocks to GPU memory on demand, eliminating fragmentation and fitting more concurrent requests.

Let’s prove it with an experiment: two allocators competing for the same memory budget. We fix the worst-case context at 8K tokens, but draw request lengths from a long-tailed distribution — most requests are far shorter than the worst case, which is exactly the workload that punishes static allocation.

Set up the KV math. Reuse the model and GPU from the Silicon Zoo and compute the cache cost of a single token:

import random
from mlsysim import Q_
from mlsysim.show import table

model = mlsysim.Models.Language.Llama3_8B
hardware = mlsysim.Hardware.Cloud.H100

# KV cache bytes per token (FP16): 2 (K+V) x layers x kv_heads x head_dim x 2 bytes
head_dim = model.hidden_dim // model.heads
bytes_per_token = 2 * model.layers * model.kv_heads * head_dim * 2
max_seq = 8192  # worst-case context length we must support
kv_budget = hardware.memory.capacity - model.size_in_bytes()  # HBM left after weights
budget_tokens = int(kv_budget.to("byte").magnitude // bytes_per_token)

The two allocators. Static contiguous reserves a full max_seq slot per request, whether the request needs it or not. Paged hands out fixed-size blocks on demand — a request only pays for the blocks it actually fills:

def paged_blocks(tokens, page):
    """Whole blocks needed for `tokens` tokens (ceiling division)."""
    return -(-tokens // page)


def sample_lengths(seed, n, mean=2048, cap=max_seq):
    """n request lengths from a long-tailed distribution (avg well below the cap)."""
    rng = random.Random(seed)
    return [max(1, min(cap, int(rng.expovariate(1 / mean)))) for _ in range(n)]


class StaticAllocator:
    """Static contiguous KV: every request reserves a full max_seq slot up front."""

    def __init__(self, budget_tokens, max_seq):
        self.budget = budget_tokens
        self.max_seq = max_seq
        self.reserved = 0
        self.used_tokens = 0
        self.admitted = []

    def admit(self, tokens):
        if self.reserved + self.max_seq <= self.budget:
            self.reserved += self.max_seq
            self.used_tokens += tokens
            self.admitted.append(tokens)
            return True
        return False


class PagedAllocator:
    """Paged KV: fixed-size blocks handed out on demand from a free-block pool."""

    def __init__(self, budget_tokens, page):
        self.page = page
        self.free_blocks = budget_tokens // page
        self.used_tokens = 0
        self.admitted = []

    def admit(self, tokens):
        need = paged_blocks(tokens, self.page)
        if need <= self.free_blocks:
            self.free_blocks -= need
            self.used_tokens += tokens
            self.admitted.append(tokens)
            return True
        return False

Same memory budget, two strategies. Feed both allocators the same stream of requests in arrival order until the next request no longer fits, and watch where the memory goes:

def fill(alloc, lengths):
    """Admit requests in arrival order; stop at the first one that does not fit."""
    for t in lengths:
        if not alloc.admit(t):
            break
    return alloc


lens = sample_lengths(seed=42, n=10_000)

static = fill(StaticAllocator(budget_tokens, max_seq), lens)
paged16 = fill(PagedAllocator(budget_tokens, 16), lens)
paged64 = fill(PagedAllocator(budget_tokens, 64), lens)
paged256 = fill(PagedAllocator(budget_tokens, 256), lens)


def kv_tokens(alloc):
    if isinstance(alloc, StaticAllocator):
        return alloc.reserved
    return sum(paged_blocks(t, alloc.page) * alloc.page for t in alloc.admitted)


def internal_frag_pct(alloc):
    return 1 - alloc.used_tokens / kv_tokens(alloc)


def decode_throughput(alloc):
    # Memory-bound decode (Tutorial 2): weights + used KV streamed once per token
    step_bytes = model.size_in_bytes().magnitude + alloc.used_tokens * bytes_per_token
    step_sec = step_bytes / hardware.memory.bandwidth.to("byte/s").magnitude
    return len(alloc.admitted) / step_sec


rows = []
for label, alloc in [(f"Static (reserve {max_seq // 1024}K)", static), ("Paged (16 tok)", paged16),
                     ("Paged (64 tok)", paged64), ("Paged (256 tok)", paged256)]:
    rows.append([label, len(alloc.admitted),
                 f"{Q_(kv_tokens(alloc) * bytes_per_token, 'byte').to('GiB'):~.1f}",
                 f"{internal_frag_pct(alloc):.1%}",
                 f"{decode_throughput(alloc):,.0f} t/s"])

table(["System", "Max Users", "KV Allocated", "Int Frag", "Throughput"], rows)
System               Max Users  KV Allocated  Int Frag  Throughput
──────────────────────────────────────────────────────────────────
Static (reserve 8K)         65      65.0 GiB     77.8%   6,894 t/s
Paged (16 tok)             259      64.8 GiB      0.4%  10,162 t/s
Paged (64 tok)             255      64.4 GiB      1.5%  10,152 t/s
Paged (256 tok)            245      64.8 GiB      6.0%  10,075 t/s
ImportantKey Insight

The same 65 GiB KV budget serves 65 users at 77.8% internal fragmentation with static reservation — or 259 users at 0.4% with 16-token pages. Static pays for the worst case (8K) on every request even though the average request is ~4× shorter; paged pays for what is used. The extra concurrency buys 1.47× decode throughput under the memory-bound model from Tutorial 2.

The second kind of waste: external fragmentation. The static allocator above never suffers from this, because every freed slot is a full max_seq and fits any request. That is exactly the trade it makes, paying in internal waste instead. A contiguous allocator that hands out exact sizes avoids the reservation waste but fragments under churn: requests finish out of order, leaving holes too small for the next long request. Paged allocation never does — any free block works, so external fragmentation is ~0:

class FirstFitAllocator:
    """Contiguous exact-size allocation: variable-size chunks, first-fit, holes merge on free."""

    def __init__(self, budget_tokens):
        self.budget = budget_tokens
        self.holes = [(0, budget_tokens)]
        self.allocated = {}
        self._next_id = 0

    def admit(self, tokens):
        for i, (s, sz) in enumerate(self.holes):
            if sz >= tokens:
                self.holes.pop(i)
                if sz > tokens:
                    self.holes.insert(i, (s + tokens, sz - tokens))
                self.allocated[self._next_id] = (s, tokens)
                self._next_id += 1
                return True
        return False

    def free(self, rid):
        s, sz = self.allocated.pop(rid)
        self.holes.append((s, sz))
        self.holes.sort()
        merged = []
        for s, sz in self.holes:
            if merged and merged[-1][0] + merged[-1][1] == s:
                merged[-1] = (merged[-1][0], merged[-1][1] + sz)
            else:
                merged.append((s, sz))
        self.holes = merged

    def external_frag_pct(self):
        free = sum(sz for _, sz in self.holes)
        largest = max((sz for _, sz in self.holes), default=0)
        return 0.0 if free == 0 else 1 - largest / free


class PagedPool:
    def __init__(self, budget_tokens, page):
        self.page = page
        self.free_blocks = budget_tokens // page
        self.admitted = []

    def admit(self, tokens):
        need = paged_blocks(tokens, self.page)
        if need <= self.free_blocks:
            self.free_blocks -= need
            self.admitted.append(tokens)
            return True
        return False

    def free(self, tokens):
        self.free_blocks += paged_blocks(tokens, self.page)


rng = random.Random(777)
first_fit = FirstFitAllocator(budget_tokens)
paged = PagedPool(budget_tokens, 16)

# Fill both pools in arrival order until each is full, then complete every other
# request (finish out of order — the churn that fragments contiguous memory).
# Both pools admit a prefix of the same wave, so index i is the same request in each.
wave = [rng.randint(256, max_seq) for _ in range(400)]
fill(first_fit, wave)
fill(paged, wave)

n_done = min(len(first_fit.allocated), len(paged.admitted))
for i in range(0, n_done, 2):
    first_fit.free(i)
    paged.free(paged.admitted[i])

# A fresh wave of worst-case requests arrives — who can still admit them?
new_reqs = [max_seq] * 10
ff_adm = sum(first_fit.admit(t) for t in new_reqs)
pg_adm = sum(paged.admit(t) for t in new_reqs)
free_ff = int(sum(sz for _, sz in first_fit.holes))
table(["System", "External Frag", "Free Tokens", f"New {max_seq // 1024}K Req. Admitted"],
      [["Contiguous (first-fit)", f"{first_fit.external_frag_pct():.1%}", free_ff, f"{ff_adm}/10"],
       ["Paged (16 tok)", "~0%", int(paged.free_blocks * 16), f"{pg_adm}/10"]])
System                  External Frag  Free Tokens  New 8K Req. Admitted
────────────────────────────────────────────────────────────────────────
Contiguous (first-fit)          97.1%       268270                  1/10
Paged (16 tok)                    ~0%       194080                 10/10

Page size: the fragmentation knob. Page size trades waste against bookkeeping. Small pages leave a smaller unused tail on each sequence; large pages mean fewer block-table entries per sequence but waste more. Sweep the page size and watch both:

rows = []
for page in [8, 16, 64, 256, 1024]:
    a = fill(PagedAllocator(budget_tokens, page), lens)
    n = len(a.admitted)
    blocks_per_req = sum(paged_blocks(t, page) for t in a.admitted) / n
    rows.append([page, n, f"{internal_frag_pct(a):.1%}", f"{blocks_per_req:.1f}",
                 paged_blocks(max_seq, page), f"{kv_tokens(a) / budget_tokens:.1%}"])

table(["Page Size", "Max Users", "Int Frag", "Avg Blocks/Req",
       f"Blocks/{max_seq // 1024}K Seq", "Pool Used"], rows)
Page Size  Max Users  Int Frag  Avg Blocks/Req  Blocks/8K Seq  Pool Used
────────────────────────────────────────────────────────────────────────
8                259      0.2%           255.7           1024      99.4%
16               259      0.4%           128.1            512      99.6%
64               255      1.5%            32.3            128      99.0%
256              245      6.0%             8.5             32      99.6%
1024             209     22.0%             2.5              8      99.2%

16-token pages (vLLM’s default) sit at the sweet spot: 0.4% waste, with about 128 blocks for an average request and 512 for a full 8K sequence. Going to 1024-token pages shrinks the block table 64-fold (8 entries for a full 8K sequence) but burns 22.0% of the cache as slack and serves 50 fewer users, the wrong trade for long-tail workloads.

The third trick: sharing. Because blocks are first-class, sequences can share them. Eight requests that begin with the same 1K-token system prompt (or eight beam-search branches that fork from one prefix) can point at the same blocks via reference counts — copy-on-write — instead of storing the prefix eight times:

prefix, gen, n_req = 1024, 512, 8

rows = []
static_kv = n_req * max_seq
no_share = n_req * paged_blocks(prefix + gen, 16) * 16
shared = paged_blocks(prefix, 16) * 16 + n_req * paged_blocks(gen, 16) * 16
for label, kv in [("Static (reserve max)", static_kv),
                  ("Paged, no sharing", no_share),
                  ("Paged + shared prefix", shared)]:
    rows.append([label, f"{Q_(kv * bytes_per_token, 'byte').to('MiB'):~.1f}",
                 f"{kv / static_kv:.1%}"])

table(["Scheme", "KV Allocated", "vs Static"], rows)
Scheme                 KV Allocated  vs Static
──────────────────────────────────────────────
Static (reserve max)     8192.0 MiB     100.0%
Paged, no sharing        1536.0 MiB      18.8%
Paged + shared prefix     640.0 MiB       7.8%

Sharing the prefix cuts peak KV from 1536 MiB to 640 MiB — a 58% saving that grows with the number of concurrent requests sharing the prompt. Production systems call this prefix caching; it is why shared system prompts and few-shot examples are so cheap to serve.

Paged attention therefore attacks the memory wall from three directions at once: it eliminates reservation waste (~78% → <1% internal fragmentation), survives churn with ~0 external fragmentation, and shares common prefixes across requests. This is why vLLM and TensorRT-LLM default to paged KV-cache management in production.


Your Turn

CautionExercises

Exercise 1: Predict before you compute. Llama-3 70B has 80 layers (vs. 32 for the 8B model) and 8 KV-heads with 128 head_dim. Before running any code, predict: at seq_len=4096 and FP16, what batch size will cause OOM on an 80 GB H100? Write your prediction, then sweep batch sizes with mlsysim.Models.Language.Llama3_70B to find the actual limit. How close were you?

Exercise 2: Maximum users at 128K context. Using the H200 (141 GB HBM3e), calculate the maximum number of concurrent users you can serve with Llama-3 8B at 128K context in FP16. Then try INT8. How many additional users does quantization buy you?

Exercise 3: Paged vs. static at long context. Change max_seq in Section 5 to 32768 and re-run the cells. How many users do Static (reserve 32K) and Paged (16 tok) serve now, and why does raising the worst case hurt one far more than the other? Then compare Paged (16 tok) with Paged (256 tok). Does page size matter more when requests get longer, or less? Test your prediction by passing mean=16384 to sample_lengths. (Hint: page slack is at most one page per request, a fixed cost that shrinks as a share of a longer request.)

Self-check: If a model has 32 layers, 8 KV-heads, 128 head_dim, and uses FP16 (2 bytes), how many bytes does the KV-cache consume per token per user? (Answer: 2 x 32 x 8 x 128 x 2 = 131,072 bytes = 128 KB per token.)


Key Takeaways

TipSummary
  • KV-cache size scales linearly with layers, KV-heads, sequence length, and batch size
  • At short context, cache is negligible — model weights dominate and you can serve many users
  • At long context, cache dominates — a single 128K user’s cache can exceed model weights
  • The OOM boundary depends on context length x batch size, not just model size
  • Paged attention reduces fragmentation, fitting more concurrent requests in the same memory and sharing common prefixes across requests

Next Steps

  • Quantization: Not a Free Lunch — Learn when reducing precision shrinks the KV-cache effectively vs. when it doesn’t help
  • Two Phases, One Request — Revisit the prefill/decode split now that you understand the cache pressure
  • Where to Invest — Use sensitivity analysis to quantify whether more memory or more bandwidth helps more
  • Silicon Zoo — Compare HBM capacity across H100, H200, MI300X, and see which GPUs tolerate long context
Back to top