Systems & ScaleintermediateGPU~40 minColab
The KV Cache Is the Real Memory Bill
ذاكرة المفاتيح والقيم: فاتورة الذاكرة الحقيقية
Memory is three bills, and only one grows with every token served
Someone asks whether a GPU can serve a model, and the answer comes from the parameter count: a billion parameters in is two gigabytes, so it fits. That arithmetic prices the model and forgets the conversations. Every sequence being generated keeps a key vector and a value vector for every it has seen, in every layer, and that store grows with both context length and the number of users. Here you build a Llama-shaped decoder, measure its three memory bills separately, and find where the cache overtakes the weights: about 14,353 tokens in flight on this model, and 25,711 on Llama-2-7B by the same formula once you have checked it against the GPU.
The goal
Measure weights, activations and KV cache separately; confirm the cache matches its formula to within one percent; watch it pass the weights inside an ordinary context sweep; and show which levers move which bill — grouped-query attention shrinks the cache, FlashAttention removes activations and leaves the cache exactly where it was.
Colab opens a read-only copy. Save a copy to Drive to keep your edits.
The notebook needs a keyboard — best opened on a desktop.
The papers behind this
Ask how much memory a language model needs and you will usually hear its parameter count multiplied by two bytes. That is the price of loading the model. It is not the price of using it. During , every conversation the model is serving carries its own growing store of keys and values, the KV cache, and at long contexts or large batches that store becomes the largest thing in . This workshop measures it directly rather than asserting it.
Check that the runtime has a GPU; every number below is read from its memory allocator.
import azimuth_nb as azimuth
env = azimuth.setup(SLUG, lang=LANG, profile=PROFILE)The KV Cache Is the Real Memory Bill
Tesla T4 · 14.6 GB · 12.7 GB RAM · PyTorch 2.11.0+cu130
profile: free
ready · seed=17, d_model=2048, n_layers=16, n_heads=16, ffn=5504, vocab=32000, batch=8, probe_ctx=2048, contexts=[512, 1024, 2048, 4096], flash_batch=4, flash_ctx=2048, kv_variants=[16, 4, 1], gqa_kv=4, decode_batch=8, decode_contexts=[256, 1024, 4096], decode_warmup=3, decode_steps=20, plan_ctx=4096, budget_gib=10, headroom_gib=0.5, project_batch=8, project_ctx=4096
code · 9ef034b8e11cd93fIn a model writes one token at a time, and each new token must pay to every token before it. In that means comparing the new token's query against the key of every earlier token, then mixing their values. The keys and values of earlier tokens never change, so recomputing them at every step would be wasted work. The model keeps them instead. That saving is real, and so is its cost: two vectors per token, per layer, for as long as the conversation lasts.
Compare the two numbers at the end: the weights are a fixed size, while the cache is quoted per token. The whole workshop is the question of how many tokens it takes for the second to catch the first.
import gc
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
assert torch.cuda.is_available(), "This workshop measures GPU memory: switch the runtime to a GPU."
_tv = tuple(int(p) for p in torch.__version__.split("+")[0].split(".")[:2])
assert _tv >= (2, 1), (
f"torch {torch.__version__} is too old for fused attention; use Colab's default runtime."
)
DEVICE = "cuda"
DTYPE = torch.float16
GIB = 1024**3
cfg = env.cfg
D_MODEL = cfg["d_model"]
N_LAYERS = cfg["n_layers"]
N_HEADS = cfg["n_heads"]
D_HEAD = D_MODEL // N_HEADS
FFN = cfg["ffn"]
VOCAB = cfg["vocab"]
def say(en, ar):
print(ar if env.lang == "ar" else en)
def free():
gc.collect()
torch.cuda.empty_cache()
def linear(n_in, n_out):
return nn.Linear(n_in, n_out, bias=False, device=DEVICE, dtype=DTYPE)
class RMSNorm(nn.Module):
def __init__(self, dim):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim, device=DEVICE, dtype=DTYPE))
def forward(self, x):
xf = x.float()
xf = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + 1e-6)
return xf.to(x.dtype) * self.weight
class Attention(nn.Module):
# n_kv_heads == N_HEADS is multi-head attention (MHA),
# 1 < n_kv_heads < N_HEADS is grouped-query (GQA), n_kv_heads == 1 is multi-query (MQA).
# Rotary position embeddings are left out: they change the values in K, not its size.
def __init__(self, n_kv_heads):
super().__init__()
self.n_kv = n_kv_heads
self.group = N_HEADS // n_kv_heads
self.q = linear(D_MODEL, N_HEADS * D_HEAD)
self.k = linear(D_MODEL, n_kv_heads * D_HEAD)
self.v = linear(D_MODEL, n_kv_heads * D_HEAD)
self.o = linear(N_HEADS * D_HEAD, D_MODEL)
def prefill(self, x, naive):
B, T, _ = x.shape
q = self.q(x).view(B, T, N_HEADS, D_HEAD).transpose(1, 2)
k = self.k(x).view(B, T, self.n_kv, D_HEAD).transpose(1, 2).contiguous()
v = self.v(x).view(B, T, self.n_kv, D_HEAD).transpose(1, 2).contiguous()
# The cache keeps k and v with n_kv heads. The expanded copy below is
# only borrowed for this one computation and is freed with it.
kk, vv = k, v
if self.group > 1:
kk = k.repeat_interleave(self.group, dim=1)
vv = v.repeat_interleave(self.group, dim=1)
if naive:
# Textbook attention: the full T x T score matrix exists in memory.
scores = (q @ kk.transpose(-2, -1)) / math.sqrt(D_HEAD)
mask = torch.ones(T, T, dtype=torch.bool, device=DEVICE).triu(1)
scores = scores.masked_fill(mask, float("-inf"))
out = scores.softmax(dim=-1) @ vv
else:
# Fused kernel (FlashAttention or memory-efficient, whichever the GPU
# supports): the score matrix is computed in tiles and never stored.
out = F.scaled_dot_product_attention(q, kk, vv, is_causal=True)
out = out.transpose(1, 2).reshape(B, T, N_HEADS * D_HEAD)
return self.o(out), (k, v)
def decode(self, x, kv, pos):
# One new token per sequence, attending over a cache of fixed length.
# A fixed length stands in for a growing one: the step reads every
# cached position either way.
B = x.shape[0]
k_cache, v_cache = kv
q = self.q(x).view(B, self.n_kv, self.group, D_HEAD)
slot = pos % k_cache.shape[2]
k_cache[:, :, slot] = self.k(x).view(B, self.n_kv, D_HEAD)
v_cache[:, :, slot] = self.v(x).view(B, self.n_kv, D_HEAD)
# Grouped form: each cached head is read once and shared by its group
# of query heads. Nothing is expanded.
scores = (q @ k_cache.transpose(-2, -1)) / math.sqrt(D_HEAD)
out = scores.softmax(dim=-1) @ v_cache
return self.o(out.reshape(B, 1, N_HEADS * D_HEAD))
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.gate = linear(D_MODEL, FFN)
self.up = linear(D_MODEL, FFN)
self.down = linear(FFN, D_MODEL)
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))
class Block(nn.Module):
def __init__(self, n_kv_heads):
super().__init__()
self.norm1 = RMSNorm(D_MODEL)
self.attn = Attention(n_kv_heads)
self.norm2 = RMSNorm(D_MODEL)
self.mlp = MLP()
def prefill(self, x, naive):
a, kv = self.attn.prefill(self.norm1(x), naive)
x = x + a
return x + self.mlp(self.norm2(x)), kv
def decode(self, x, kv, pos):
x = x + self.attn.decode(self.norm1(x), kv, pos)
return x + self.mlp(self.norm2(x))
class Decoder(nn.Module):
def __init__(self, n_kv_heads):
super().__init__()
self.n_kv = n_kv_heads
self.embed = nn.Embedding(VOCAB, D_MODEL, device=DEVICE, dtype=DTYPE)
self.blocks = nn.ModuleList(Block(n_kv_heads) for _ in range(N_LAYERS))
self.norm = RMSNorm(D_MODEL)
self.head = linear(D_MODEL, VOCAB)
def prefill(self, ids, naive=False):
x = self.embed(ids)
cache = []
for block in self.blocks:
x, kv = block.prefill(x, naive)
cache.append(kv)
# Logits for the last position only: that is all generation needs.
return cache, self.head(self.norm(x[:, -1:]))
def decode(self, ids, cache, pos):
x = self.embed(ids)
for block, kv in zip(self.blocks, cache):
x = block.decode(x, kv, pos)
return self.head(self.norm(x))
def filled_cache(self, batch, length):
# A cache of the right shape for timing; its contents do not matter.
shape = (batch, self.n_kv, length, D_HEAD)
return [
(
torch.randn(shape, device=DEVICE, dtype=DTYPE),
torch.randn(shape, device=DEVICE, dtype=DTYPE),
)
for _ in range(N_LAYERS)
]
def build(n_kv_heads):
torch.manual_seed(cfg["seed"])
model = Decoder(n_kv_heads)
with torch.no_grad():
for p in model.parameters():
if p.dim() > 1:
p.normal_(0.0, 0.02)
return model.eval().requires_grad_(False)
def weight_bytes(model):
return sum(p.numel() * p.element_size() for p in model.parameters())
def predicted_cache_bytes(n_layers, n_kv, d_head, batch, ctx, bytes_per_value=2):
# 2 tensors (K and V) x layers x kv heads x head size x bytes, per token.
return 2 * n_layers * n_kv * d_head * bytes_per_value * batch * ctx
free()
_before = torch.cuda.memory_allocated()
model = build(N_HEADS)
weights_measured = torch.cuda.memory_allocated() - _before
weights_gib = round(weight_bytes(model) / GIB, 2)
n_params_m = round(sum(p.numel() for p in model.parameters()) / 1e6)
kv_kib_per_token = round(predicted_cache_bytes(N_LAYERS, N_HEADS, D_HEAD, 1, 1) / 1024)
say(
f"{n_params_m:,.0f}M parameters, {N_LAYERS} layers, {N_HEADS} heads of size {D_HEAD}\n"
f"weights in fp16: {weights_gib:.2f} GiB (allocator saw {weights_measured / GIB:.2f} GiB)\n"
f"cache per token, by the formula: {kv_kib_per_token:.0f} KiB",
f"عدد المعاملات: {n_params_m:,.0f} مليون. الطبقات: {N_LAYERS}. رؤوس الانتباه: {N_HEADS}، وحجم كلٍّ منها {D_HEAD}\n"
f"الأوزان بدقة fp16: {weights_gib:.2f} GiB (رصدها المُخصِّص {weights_measured / GIB:.2f} GiB)\n"
f"الذاكرة المؤقتة لكل رمز، حسب المعادلة: {kv_kib_per_token:.0f} KiB",
)941M parameters, 16 layers, 16 heads of size 128
weights in fp16: 1.75 GiB (allocator saw 1.78 GiB)
cache per token, by the formula: 128 KiBThe model's during a forward pass has three parts. The weights are allocated once and stay. Activations are the intermediate tensors of the pass itself; they exist while it runs and are freed when it ends, so they show up only as a peak. The cache is what the pass leaves behind. The measurement below separates them by reading the allocator three times: before the pass, at its peak, and after it.
One moment of serving: 8 sequences that have just read a prompt of 2048 tokens each. Look at how close the cache already is to the weights.
import pandas as pd
def measure_prefill(model, batch, ctx, naive=False):
"""Run one prefill and split memory into what stays and what passed through."""
gen = torch.Generator(device=DEVICE).manual_seed(cfg["seed"])
ids = torch.randint(0, VOCAB, (batch, ctx), device=DEVICE, generator=gen)
free()
torch.cuda.synchronize()
before = torch.cuda.memory_allocated()
torch.cuda.reset_peak_memory_stats()
with torch.inference_mode():
cache, logits = model.prefill(ids, naive=naive)
torch.cuda.synchronize()
after = torch.cuda.memory_allocated()
peak = torch.cuda.max_memory_allocated()
result = {
# What the forward pass leaves behind, minus the last-token logits.
"cache": after - before - logits.untyped_storage().nbytes(),
# What existed only while the pass ran.
"act": peak - after,
}
del cache, logits, ids
free()
return result
BATCH = cfg["batch"]
PROBE_CTX = cfg["probe_ctx"]
probe = measure_prefill(model, BATCH, PROBE_CTX)
probe_weights_gib = weights_gib
probe_cache_gib = round(probe["cache"] / GIB, 2)
probe_act_gib = round(probe["act"] / GIB, 2)
say(
f"{BATCH} sequences x {PROBE_CTX} tokens, fp16",
f"الدفعة: {BATCH} × {PROBE_CTX} (تسلسلات × رموز)، بدقة fp16",
)
if env.lang == "ar":
probe_table = pd.DataFrame(
{
"الفاتورة": ["الأوزان", "التفعيلات (ذروة عابرة)", "ذاكرة المفاتيح والقيم"],
"GiB": [probe_weights_gib, probe_act_gib, probe_cache_gib],
}
)
else:
probe_table = pd.DataFrame(
{
"bill": ["weights", "activations (transient peak)", "KV cache"],
"GiB": [probe_weights_gib, probe_act_gib, probe_cache_gib],
}
)
probe_table.round(2)8 sequences x 2048 tokens, fp16
bill GiB
0 weights 1.75
1 activations (transient peak) 0.75
2 KV cache 2.01The grey base never moves. Watch the orange band, and the dotted line where the cache would equal the weights.
import glob
import matplotlib.pyplot as plt
from matplotlib import font_manager
if env.lang == "ar":
# matplotlib draws Arabic as isolated letters in logical order unless the
# text is reshaped and reordered, and its default font has no Arabic glyphs.
import arabic_reshaper
from bidi.algorithm import get_display
_fonts = sorted(glob.glob("/usr/share/fonts/**/NotoSansArabic-Regular.ttf", recursive=True))
if _fonts:
font_manager.fontManager.addfont(_fonts[0])
_family = font_manager.FontProperties(fname=_fonts[0]).get_name()
plt.rcParams["font.family"] = [_family, "DejaVu Sans"]
else:
print("تحذير: لم يُعثر على خط عربي؛ ستظهر النصوص العربية في الرسوم مربعات فارغة.")
def ar(en, ar_text):
if env.lang == "ar":
return get_display(arabic_reshaper.reshape(ar_text))
return en
rows = []
for ctx in cfg["contexts"]:
m = measure_prefill(model, BATCH, ctx)
rows.append({"ctx": ctx, "cache": m["cache"], "act": m["act"]})
wb = weight_bytes(model)
sweep_max_ctx = rows[-1]["ctx"]
cache_gib_max = round(rows[-1]["cache"] / GIB, 2)
act_gib_max = round(rows[-1]["act"] / GIB, 2)
cache_over_weights = round(rows[-1]["cache"] / wb, 2)
# Measured bytes per token in flight, from the largest point.
cache_bytes_per_token = rows[-1]["cache"] / (BATCH * sweep_max_ctx)
crossover_tokens = round(wb / cache_bytes_per_token)
crossover_ctx = round(crossover_tokens / BATCH)
fig, ax = plt.subplots(figsize=(7.5, 4.2))
xs = range(len(rows))
w = [wb / GIB] * len(rows)
c = [r["cache"] / GIB for r in rows]
a = [r["act"] / GIB for r in rows]
ax.bar(xs, w, color="#8a94a6", label=ar("weights", "الأوزان"))
ax.bar(xs, c, bottom=w, color="#d9822b", label=ar("KV cache", "ذاكرة المفاتيح والقيم"))
ax.bar(
xs,
a,
bottom=[wi + ci for wi, ci in zip(w, c)],
color="#5b8def",
alpha=0.55,
label=ar("activations (transient)", "التفعيلات (عابرة)"),
)
ax.axhline(2 * wb / GIB, color="#333", lw=0.8, ls=":")
ax.text(
len(rows) - 0.5,
2 * wb / GIB,
ar(" cache = weights", " الذاكرة المؤقتة = الأوزان"),
va="bottom",
ha="right",
fontsize=8,
)
ax.set_xticks(list(xs), [str(r["ctx"]) for r in rows])
ax.set_xlabel(ar(f"context length (batch of {BATCH})", f"طول السياق (دفعة من {BATCH})"))
ax.set_ylabel("GiB")
ax.set_ylim(bottom=0)
ax.set_title(ar("Where the memory goes during prefill", "أين تذهب الذاكرة أثناء معالجة الموجّه"))
ax.legend(loc="upper left", frameon=False)
fig.tight_layout()
plt.show()
say(
f"At {sweep_max_ctx} tokens the cache is {cache_over_weights:.2f}x the weights.\n"
f"It matches the weights at {crossover_tokens:,.0f} tokens in flight "
f"({crossover_ctx:,.0f} tokens each across {BATCH} sequences).",
f"عند سياق طوله {sweep_max_ctx}، حجم الذاكرة المؤقتة = {cache_over_weights:.2f} × حجم الأوزان.\n"
f"نقطة التعادل مع الأوزان: {crossover_tokens:,.0f} رمزاً قيد المعالجة معاً، "
f"أي سياق طوله {crossover_ctx:,.0f} لكل تسلسل في دفعة من {BATCH}.",
)
At 4096 tokens the cache is 2.28x the weights.
It matches the weights at 14,353 tokens in flight (1,794 tokens each across 8 sequences).At a context of 4,096 tokens the cache is 4 GiB, 2.280 times the weights. Nothing about that context is unusual: it is a typical served to 8 users at once. The two bills are equal at about 14,353 tokens in flight, which across a of 8 is a of only about 1,794 each. The cache depends on the product of batch and context, so doubling the users costs the same as doubling the context. Activations, meanwhile, peaked at 1.510 GiB; they are real, but they are paid once per pass and then returned.
The first line should read 1.0000: the allocator and the arithmetic agree to the byte. Then compare the last two rows of the table, the same 70B model with and without shared key heads.
predicted = [predicted_cache_bytes(N_LAYERS, N_HEADS, D_HEAD, BATCH, r["ctx"]) for r in rows]
formula_ratio = round(sum(r["cache"] for r in rows) / sum(predicted), 4)
say(
f"measured cache / formula, summed over the sweep: {formula_ratio:.4f}",
f"الذاكرة المقيسة ÷ المعادلة، مجموعةً على كل نقاط المسح: {formula_ratio:.4f}",
)
# The formula now has a measurement behind it, so it can be pointed at models
# this GPU could never load. Shapes are from the Llama 2 paper.
LLAMA2 = [
# English name, Arabic name, parameters, layers, kv heads, head size
("Llama-2-7B (MHA)", "Llama-2-7B (MHA)", 6.74e9, 32, 32, 128),
("Llama-2-13B (MHA)", "Llama-2-13B (MHA)", 13.0e9, 40, 40, 128),
("Llama-2-70B (GQA, 8 kv heads)", "Llama-2-70B (GQA، 8 رؤوس KV)", 69.0e9, 80, 8, 128),
("Llama-2-70B if it were MHA", "Llama-2-70B لو كان MHA", 69.0e9, 80, 64, 128),
]
LLAMA_CTX = cfg["project_ctx"]
LLAMA_BATCH = cfg["project_batch"]
proj = []
for name_en, name_ar, params, layers, kv, dh in LLAMA2:
name = name_ar if env.lang == "ar" else name_en
per_token = predicted_cache_bytes(layers, kv, dh, 1, 1)
weights = params * 2
proj.append(
{
"name": name,
"kib_per_token": per_token / 1024,
"weights_gib": weights / GIB,
"crossover_tokens": weights / per_token,
"cache_gib_b8": per_token * LLAMA_BATCH * LLAMA_CTX / GIB,
}
)
llama7b_crossover_tokens = round(proj[0]["crossover_tokens"])
llama7b_cache_over_weights_b8 = round(proj[0]["cache_gib_b8"] / proj[0]["weights_gib"], 2)
llama70b_crossover_tokens = round(proj[2]["crossover_tokens"])
llama70b_mha_crossover_tokens = round(proj[3]["crossover_tokens"])
if env.lang == "ar":
cols = {
"name": "النموذج",
"kib_per_token": "KiB لكل رمز",
"weights_gib": "الأوزان GiB",
"crossover_tokens": "رموز التعادل",
"cache_gib_b8": f"الذاكرة المؤقتة GiB ({LLAMA_BATCH} × {LLAMA_CTX})",
}
else:
cols = {
"name": "model",
"kib_per_token": "KiB per token",
"weights_gib": "weights GiB",
"crossover_tokens": "tokens to match weights",
"cache_gib_b8": f"cache GiB at {LLAMA_BATCH} x {LLAMA_CTX}",
}
proj_table = pd.DataFrame(proj).rename(columns=cols).round(1)
proj_tablemeasured cache / formula, summed over the sweep: 1.0000
model KiB per token weights GiB \
0 Llama-2-7B (MHA) 512.0 12.6
1 Llama-2-13B (MHA) 800.0 24.2
2 Llama-2-70B (GQA, 8 kv heads) 320.0 128.5
3 Llama-2-70B if it were MHA 2560.0 128.5
tokens to match weights cache GiB at 8 x 4096
0 25711.1 16.0
1 31738.3 25.0
2 421142.6 10.0
3 52642.8 80.0A measured ratio of 1 means the formula is not an estimate: it is the exact size of the cache. That is what lets us apply it to models this GPU could never load. Llama-2-7B matches its own weights at 25,711 tokens in flight, so 8 users at its full 4096-token context hold 1.270 times its weights in cache. Llama-2-70B is ten times larger, yet it holds out until 421,143 tokens. The reason is in the last row: with a key head for every query head it would cross at 52,643. The difference is one design decision about attention, not the size of the model.
Two techniques are often filed together as "making attention cheaper on memory", and they act on different bills. changes how attention is computed: the score matrix, one entry per pair of tokens, is processed in tiles and never stored whole. That is an activation. Grouped-query attention changes what is stored: several query heads share one key head and one value head, so there are fewer keys and values to keep. That is the cache. The next two cells pull each lever on its own, and a third tries a lever that applies to either: storing fewer bits.
Read the two columns separately. One changes by a large factor. The other should not change at all.
FB, FC = cfg["flash_batch"], cfg["flash_ctx"]
naive = measure_prefill(model, FB, FC, naive=True)
fused = measure_prefill(model, FB, FC, naive=False)
naive_act_gib = round(naive["act"] / GIB, 2)
fused_act_gib = round(fused["act"] / GIB, 2)
act_reduction = round(naive["act"] / fused["act"], 1)
cache_naive_gib = round(naive["cache"] / GIB, 2)
cache_fused_gib = round(fused["cache"] / GIB, 2)
score_matrix_gib = round(FB * N_HEADS * FC * FC * 2 / GIB, 2)
if env.lang == "ar":
flash_table = pd.DataFrame(
{
"الانتباه": ["تقليدي (مصفوفة كاملة)", "مدمج، على شكل كتل"],
"التفعيلات GiB": [naive_act_gib, fused_act_gib],
"الذاكرة المؤقتة GiB": [cache_naive_gib, cache_fused_gib],
}
)
else:
flash_table = pd.DataFrame(
{
"attention": ["naive (full matrix)", "fused, tiled kernel"],
"activations GiB": [naive_act_gib, fused_act_gib],
"KV cache GiB": [cache_naive_gib, cache_fused_gib],
}
)
say(
f"{FB} x {FC} tokens. One layer's score matrix alone: {score_matrix_gib:.2f} GiB",
f"الدفعة: {FB} × {FC} (تسلسلات × رموز). مصفوفة الدرجات لطبقة واحدة وحدها: {score_matrix_gib:.2f} GiB",
)
flash_table.round(2)4 x 2048 tokens. One layer's score matrix alone: 0.50 GiB
attention activations GiB KV cache GiB
0 naive (full matrix) 1.13 1.0
1 fused, tiled kernel 0.38 1.0Three models that differ only in how many key heads they keep: multi-head, grouped, and multi-query. Compare how far the cache falls with how little the weights move.
del model
free()
variants = []
for kv in cfg["kv_variants"]:
m_kv = build(kv)
res = measure_prefill(m_kv, BATCH, PROBE_CTX)
variants.append({"kv": kv, "weights": weight_bytes(m_kv), "cache": res["cache"]})
del m_kv
free()
by_kv = {v["kv"]: v for v in variants}
mha = by_kv[N_HEADS]
gqa_ratio = round(mha["cache"] / by_kv[cfg["gqa_kv"]]["cache"], 2)
mqa_ratio = round(mha["cache"] / by_kv[1]["cache"], 2)
gqa_weight_change_pct = round(
100 * (mha["weights"] - by_kv[cfg["gqa_kv"]]["weights"]) / mha["weights"], 1
)
if env.lang == "ar":
gqa_table = pd.DataFrame(
{
"رؤوس KV": [v["kv"] for v in variants],
"الأوزان GiB": [v["weights"] / GIB for v in variants],
"الذاكرة المؤقتة GiB": [v["cache"] / GIB for v in variants],
"تقلّص الذاكرة المؤقتة": [mha["cache"] / v["cache"] for v in variants],
}
)
else:
gqa_table = pd.DataFrame(
{
"kv heads": [v["kv"] for v in variants],
"weights GiB": [v["weights"] / GIB for v in variants],
"KV cache GiB": [v["cache"] / GIB for v in variants],
"cache shrink vs MHA": [mha["cache"] / v["cache"] for v in variants],
}
)
gqa_table.round(2)kv heads weights GiB KV cache GiB cache shrink vs MHA
0 16 1.75 2.00 1.0
1 4 1.56 0.50 4.0
2 1 1.52 0.12 16.0The ratio will sit a little above one half: each block of values carries its own scale. Note the error, and remember these weights are random.
def quantize(x):
# Symmetric int8, one fp16 scale per (sequence, head, position).
scale = x.abs().amax(dim=-1, keepdim=True).float().clamp(min=1e-8) / 127
q = (x.float() / scale).round().clamp(-127, 127).to(torch.int8)
return q, scale.to(DTYPE)
model = build(N_HEADS)
gen = torch.Generator(device=DEVICE).manual_seed(cfg["seed"])
ids = torch.randint(0, VOCAB, (BATCH, PROBE_CTX), device=DEVICE, generator=gen)
free()
base = torch.cuda.memory_allocated()
with torch.inference_mode():
cache, logits = model.prefill(ids)
base += logits.untyped_storage().nbytes()
fp16_bytes = torch.cuda.memory_allocated() - base
err_sq, ref_sq = 0.0, 0.0
for i in range(len(cache)):
k, v = cache[i]
qk, sk = quantize(k)
qv, sv = quantize(v)
for orig, qx, sx in ((k, qk, sk), (v, qv, sv)):
err_sq += (qx.float() * sx.float() - orig.float()).pow(2).sum().item()
ref_sq += orig.float().pow(2).sum().item()
cache[i] = (qk, sk, qv, sv)
del k, v, qk, sk, qv, sv, orig, qx, sx
free()
int8_bytes = torch.cuda.memory_allocated() - base
int8_ratio = round(int8_bytes / fp16_bytes, 3)
int8_err_pct = round(100 * math.sqrt(err_sq / ref_sq), 2)
say(
f"cache in fp16: {fp16_bytes / GIB:.2f} GiB in int8 + scales: {int8_bytes / GIB:.2f} GiB "
f"ratio {int8_ratio:.3f}\nrelative reconstruction error: {int8_err_pct:.2f}%",
f"الذاكرة المؤقتة بدقة fp16: {fp16_bytes / GIB:.2f} GiB وبدقة int8 مع معاملات التحجيم: "
f"{int8_bytes / GIB:.2f} GiB النسبة {int8_ratio:.3f}\nخطأ إعادة البناء النسبي: {int8_err_pct:.2f}%",
)
del cache, logits, ids, model
free()cache in fp16: 2.00 GiB in int8 + scales: 1.02 GiB ratio 0.508
relative reconstruction error: 0.65%The fused kernel cut activations by a factor of 3 and left the cache at 1 GiB, the same as the naive version. Grouped-query attention with 4 key heads shrank the cache by a factor of 4, and multi-query by 16, while the weights moved by only 10.700%. to int8 is a third, independent lever: it kept 0.508 of the fp16 cache at 0.650% error. That error is a floor, not a forecast. Random weights produce well-behaved keys, while trained models often have a few large channels that a per-token scale handles worse. The levers multiply: grouped heads and int8 together hold the same conversations in a small fraction of the original cache.
The cache also costs time. Generating one token does very little arithmetic per byte: every weight is read once to process one token per sequence, and every cached key and value is read once to attend over it. The speed of a decoding step is therefore set by , not by compute. As the cache grows, each step has more bytes to read.
The model is the same at every point; only the cache gets longer. Watch how the two lines separate as context grows.
def time_decode(model, batch, ctx):
cache = model.filled_cache(batch, ctx)
ids = torch.zeros(batch, 1, dtype=torch.long, device=DEVICE)
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
with torch.inference_mode():
for step in range(cfg["decode_warmup"]):
model.decode(ids, cache, step)
torch.cuda.synchronize()
start.record()
for step in range(cfg["decode_steps"]):
model.decode(ids, cache, step)
end.record()
torch.cuda.synchronize()
ms = start.elapsed_time(end) / cfg["decode_steps"]
cache_b = sum(k.numel() * k.element_size() * 2 for k, _ in cache)
del cache
free()
return ms, cache_b
DB = cfg["decode_batch"]
timings = {}
for kv in (N_HEADS, cfg["gqa_kv"]):
m_kv = build(kv)
wbytes = weight_bytes(m_kv)
timings[kv] = []
for ctx in cfg["decode_contexts"]:
ms, cache_b = time_decode(m_kv, DB, ctx)
timings[kv].append({"ctx": ctx, "ms": ms, "bytes": wbytes + cache_b})
del m_kv
free()
mha_t, gqa_t = timings[N_HEADS], timings[cfg["gqa_kv"]]
decode_ms_short = round(mha_t[0]["ms"], 1)
decode_ms_long = round(mha_t[-1]["ms"], 1)
decode_slowdown = round(mha_t[-1]["ms"] / mha_t[0]["ms"], 1)
decode_gqa_ms_long = round(gqa_t[-1]["ms"], 1)
# Bytes that must cross from memory to the cores per step, over the step time.
decode_gbps = round(mha_t[-1]["bytes"] / (mha_t[-1]["ms"] / 1000) / 1e9)
fig, ax = plt.subplots(figsize=(7.5, 4.0))
ctxs = [t["ctx"] for t in mha_t]
ax.plot(
ctxs,
[t["ms"] for t in mha_t],
"o-",
color="#d9822b",
label=ar(f"MHA, {N_HEADS} kv heads", f"MHA، رؤوس KV: {N_HEADS}"),
)
ax.plot(
ctxs,
[t["ms"] for t in gqa_t],
"s-",
color="#2b9d8f",
label=ar(f"GQA, {cfg['gqa_kv']} kv heads", f"GQA، رؤوس KV: {cfg['gqa_kv']}"),
)
ax.set_xscale("log", base=2)
ax.set_xticks(ctxs, [str(x) for x in ctxs])
ax.set_ylim(bottom=0)
ax.set_xlabel(
ar(f"cached tokens per sequence (batch of {DB})", f"الرموز المخزّنة لكل تسلسل (دفعة من {DB})")
)
ax.set_ylabel(ar("ms per generated token", "ميلي ثانية لكل رمز مولَّد"))
ax.set_title(ar("Every new token reads the whole cache", "كل رمز جديد يقرأ الذاكرة المؤقتة كاملة"))
ax.legend(frameon=False)
fig.tight_layout()
plt.show()
say(
f"MHA step: {decode_ms_short:.1f} ms at {ctxs[0]} tokens, {decode_ms_long:.1f} ms at {ctxs[-1]} "
f"({decode_slowdown:.1f}x). GQA at {ctxs[-1]}: {decode_gqa_ms_long:.1f} ms.\n"
f"effective read rate at the longest context: {decode_gbps:.0f} GB/s",
f"زمن الخطوة مع MHA: {decode_ms_short:.1f} ms عند سياق {ctxs[0]}، و{decode_ms_long:.1f} ms عند سياق {ctxs[-1]} "
f"(أي {decode_slowdown:.1f} × أبطأ). ومع GQA عند سياق {ctxs[-1]}: {decode_gqa_ms_long:.1f} ms.\n"
f"معدّل القراءة الفعلي عند أطول سياق: {decode_gbps:.0f} GB/s",
)
MHA step: 11.4 ms at 256 tokens, 40.4 ms at 4096 (3.6x). GQA at 4096: 17.6 ms.
effective read rate at the longest context: 153 GB/sWith full multi-head attention a step took 11.400 ms at the shortest context and 40.400 ms at the longest, a slowdown of 3.600 times, with no change to the model. The effective read rate of 153 GB/s is the number to compare with the card's rated bandwidth: if they are close, the step is waiting on memory, and the only way to speed it up is to read fewer bytes. Grouped-query attention does exactly that, which is why it took 17.600 ms at the same context. A smaller cache improves per token, and because more sequences fit at once, it also improves .
Exercise
You have 10 GiB of and want to serve conversations of 4096 tokens. Before running, predict how many fit with 16 key heads. Run it, then set KV_HEADS to 4 and to 1. The cell does not trust its own plan: it allocates every cache and generates a token, so the measured peak tells you whether the plan was honest. How much of the budget is left over for the weights in each case?
# YOUR TURN — size a server.
# Predict the number first, then run. Change KV_HEADS to 8, 4, 2 or 1.
KV_HEADS = N_HEADS
CONTEXT = cfg["plan_ctx"]
assert N_HEADS % KV_HEADS == 0, f"KV_HEADS must divide {N_HEADS}"
budget = cfg["budget_gib"] * GIB
headroom = cfg["headroom_gib"] * GIB
free()
base = torch.cuda.memory_allocated()
model = build(KV_HEADS)
per_seq = predicted_cache_bytes(N_LAYERS, KV_HEADS, D_HEAD, 1, CONTEXT)
plan_batch = int((budget - weight_bytes(model) - headroom) // per_seq)
say(
f"budget {cfg['budget_gib']} GiB, weights {weight_bytes(model) / GIB:.2f} GiB, "
f"{per_seq / 2**20:.0f} MiB of cache per sequence at {CONTEXT} tokens\n"
f"-> {plan_batch} sequences at once",
f"الميزانية {cfg['budget_gib']} GiB، والأوزان {weight_bytes(model) / GIB:.2f} GiB، "
f"والذاكرة المؤقتة لكل تسلسل بطول {CONTEXT}: {per_seq / 2**20:.0f} MiB\n"
f"← عدد التسلسلات التي تتسع لها البطاقة معاً: {plan_batch}",
)
if plan_batch >= 1:
# Prove the plan: hold that many full caches and generate one token.
torch.cuda.reset_peak_memory_stats()
cache = model.filled_cache(plan_batch, CONTEXT)
ids = torch.zeros(plan_batch, 1, dtype=torch.long, device=DEVICE)
with torch.inference_mode():
model.decode(ids, cache, CONTEXT - 1)
torch.cuda.synchronize()
plan_peak_gib = (torch.cuda.max_memory_allocated() - base) / GIB
say(
f"measured peak {plan_peak_gib:.2f} GiB of {cfg['budget_gib']} GiB",
f"الذروة المقيسة {plan_peak_gib:.2f} GiB من أصل {cfg['budget_gib']} GiB",
)
del cache, ids
del model
free()A hint is available in the notebook — env.hint(1)
The formula check comes first because the other two rest on it: a cache that does not match its formula means the measurement, not the model, is wrong.
formula_ok = env.check("cache-matches-formula", formula_ratio)
overtake_ok = env.check("cache-overtakes-weights", cache_over_weights)
gqa_ok = env.check("gqa-shrinks-cache", gqa_ratio)✓ Measured KV cache over the formula's prediction, summed across the sweep: 1 (needs ≥ 0.99 and ≤ 1.01)
✓ KV cache over weights at the largest point of the sweep: 2.28 (needs ≥ 1.5)
✓ MHA cache over GQA cache, same batch and context: 4 (needs ≥ 3.8)A parameter count prices the model; it says nothing about the work the model is doing. are written in parameters, but a serving bill is written in tokens in flight, and that is the bill this workshop measured. When you size hardware, start from the formula you checked here — 2 × layers × key heads × head size × bytes, per token — and multiply by the users and context you expect. Then the design choices of modern models read differently: grouped heads, fused attention, quantized caches and shorter contexts are each a way of paying one of these bills.
receipt = env.receipt()Workshop complete.
Completion code: AZ-██████████
Paste it on the workshop's page on Azimuth to record it.Terms in this workshop
- Attentionآلية الانتباه
- Self-Attentionالانتباه الذاتي
- Transformerالمحوِّل
- Autoregressive Generationالتوليد الارتجاعي
- Context Windowنافذة السياق
- Sequence Lengthطول المتتالية
- Tokenوحدة لغوية (رمز)
- Memory Footprintالبصمة الذاكرية
- GPU Memoryذاكرة المعالج الرسومي
- VRAMذاكرة المعالج الرسومي
- Inferenceالاستدلال
- Latencyزمن الاستجابة (التأخير البيني)
- Throughputمعدل التدفق والإنتاجية
- Batch Sizeحجم الدفعة الحسابية
- Quantizationالتكميم