الأنظمة والتوسّعintermediateمعالج رسوميات~40 دقيقةColab

ذاكرة المفاتيح والقيم: فاتورة الذاكرة الحقيقية

The KV Cache Is the Real Memory Bill

للذاكرة ثلاث فواتير، وواحدة منها فقط تكبر مع كل رمز يُخدَم

يسألك أحدهم: هل تكفي هذه البطاقة الرسومية لتشغيل النموذج؟ وغالباً يأتي الجواب من عدد المعاملات وحده. مليار معامل بدقة fp16 يعني غيغابايتين تقريباً، إذن البطاقة تكفي. هذا الحساب صحيح، لكنه يقدّر كلفة تحميل النموذج فقط، وينسى المحادثات التي سيخدمها. فكل تسلسل قيد التوليد يحتفظ، في كل طبقة، بمتجه مفتاح ومتجه قيمة لكل رمز مرّ عليه. وهذا المخزون يكبر كلما طال السياق، ويكبر كلما زاد عدد المستخدمين. في هذه الورشة تبني نموذجاً بأبعاد Llama، وتقيس فواتير ذاكرته الثلاث كلاً على حدة، ثم تبحث عن اللحظة التي تتجاوز فيها الذاكرة المؤقتة حجم الأوزان. في هذا النموذج تحدث عندما يبلغ عدد الرموز قيد المعالجة نحو 14,353. وفي Llama-2-7B تحدث عند 25,711، بالمعادلة نفسها بعد أن تتحقق منها على البطاقة.

الهدف

أن تقيس الأوزان والتفعيلات والذاكرة المؤقتة كلاً على حدة، وتتأكد أن الذاكرة المؤقتة تطابق معادلتها بفارق لا يتجاوز واحداً بالمئة. ثم أن تراها تتجاوز الأوزان ضمن أطوال سياق عادية تماماً. وأخيراً أن تبيّن أيّ أداة تخفّض أيّ فاتورة: الانتباه المجمّع يقلّص الذاكرة المؤقتة، أما <bdi dir="ltr">FlashAttention</bdi> فيزيل التفعيلات ويترك الذاكرة المؤقتة على حالها.

يفتح Colab نسخة للقراءة فقط. احفظ نسخة في Drive للاحتفاظ بتعديلاتك.

يحتاج الدفتر إلى لوحة مفاتيح — يُفضَّل فتحه على حاسوب مكتبي.

الأوراق وراء هذه الورشة

اسأل أيّ مهندس كم ذاكرةً يحتاج نموذج لغوي، وغالباً سيضرب عدد معاملاته في بايتين. هذا الرقم هو كلفة تحميل النموذج. لكنه ليس كلفة استخدامه. فأثناء تحمل كل محادثة يخدمها النموذج مخزوناً خاصاً بها من المفاتيح والقيم، وهذا المخزون يكبر مع كل رمز. هذا المخزون هو ذاكرة المفاتيح والقيم المؤقتة، وسنسمّيها اختصاراً «الذاكرة المؤقتة». ومع السياقات الطويلة أو الدفعات الكبيرة تصبح أكبر ما في . في هذه الورشة لن نكتفي بقول ذلك، بل سنقيسه.

تأكد أن بيئة التشغيل فيها معالج رسومي، فكل رقم في هذه الورشة نقرؤه من مُخصِّص ذاكرته.

import azimuth_nb as azimuth

env = azimuth.setup(SLUG, lang=LANG, profile=PROFILE)
شيفرة الورشة
ذاكرة المفاتيح والقيم: فاتورة الذاكرة الحقيقية
Tesla T4 · 14.6 غ.ب · ذاكرة 12.7 غ.ب · PyTorch 2.11.0+cu130
الملف: ⁦free⁩
جاهز · ⁦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⁩
الشيفرة · ⁦9ef034b8e11cd93f⁩

في يكتب النموذج واحداً في كل خطوة. وكل رمز جديد يحتاج إلى إلى كل الرموز التي سبقته. في يعني هذا أن نقارن استعلام الرمز الجديد بمفتاح كل رمز سابق، ثم نمزج قيم تلك الرموز. لكن مفاتيح الرموز السابقة وقيمها لا تتغير من خطوة إلى أخرى. فلو أعدنا حسابها في كل خطوة لأهدرنا عملاً كاملاً بلا فائدة. لذلك يحتفظ بها النموذج. هذا التوفير حقيقي، وكلفته حقيقية أيضاً: متجهان لكل رمز، في كل طبقة، ما دامت المحادثة مستمرة.

قارن الرقمين في آخر المخرجات. الأوزان حجمها ثابت، أما الذاكرة المؤقتة فحجمها مذكور لكل رمز. والورشة كلها تدور حول سؤال واحد: كم رمزاً نحتاج حتى يلحق الرقم الثاني بالأول؟

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",
)
شيفرة الورشة
عدد المعاملات: 941 مليون. الطبقات: 16. رؤوس الانتباه: 16، وحجم كلٍّ منها 128
الأوزان بدقة fp16: 1.75 GiB (رصدها المُخصِّص 1.78 GiB)
الذاكرة المؤقتة لكل رمز، حسب المعادلة: 128 KiB

تتكوّن للنموذج أثناء التمرير الأمامي من ثلاثة أجزاء. الجزء الأول هو الأوزان: تُحجز مرة واحدة وتبقى. الجزء الثاني هو التفعيلات، أي الموترات الوسيطة التي يحتاجها التمرير نفسه. هذه توجد أثناء التمرير وتُحرَّر بعد انتهائه، لذلك لا نراها إلا على شكل ذروة. الجزء الثالث هو الذاكرة المؤقتة، وهي ما يتركه التمرير خلفه بعد أن ينتهي. وللفصل بين الثلاثة نقرأ المُخصِّص ثلاث مرات: قبل التمرير، وعند ذروته، وبعده.

هذه لحظة واحدة من الخدمة: دفعة حجمها 8، قرأ كل تسلسل فيها للتو موجّهاً طوله 2048. انظر كم صارت الذاكرة المؤقتة قريبة من الأوزان منذ الآن.

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 × 2048 (تسلسلات × رموز)، بدقة fp16
الفاتورة   GiB
0                 الأوزان  1.75
1  التفعيلات (ذروة عابرة)  0.75
2   ذاكرة المفاتيح والقيم  2.01

القاعدة الرمادية لا تتحرك أبداً. راقب الشريط البرتقالي، والخط المنقّط الذي تتساوى عنده الذاكرة المؤقتة مع الأوزان.

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}.",
)
شيفرة الورشة
عند سياق طوله 4096، حجم الذاكرة المؤقتة = 2.28 × حجم الأوزان.
نقطة التعادل مع الأوزان: 14,353 رمزاً قيد المعالجة معاً، أي سياق طوله 1,794 لكل تسلسل في دفعة من 8.

عند سياق طوله 4,096 بلغت الذاكرة المؤقتة 4 GiB. هذا يساوي 2.280 ضعف الأوزان. ولا شيء غريب في هذا السياق: إنه عادية، تُخدَم لدفعة من المستخدمين حجمها 8 في وقت واحد. تتساوى الفاتورتان عندما يبلغ عدد الرموز قيد المعالجة نحو 14,353. ومع حجمها 8، يعني هذا لا يتجاوز نحو 1,794 لكل واحد منها. السبب أن الذاكرة المؤقتة تعتمد على حاصل ضرب حجم الدفعة في طول السياق. لذلك فمضاعفة عدد المستخدمين تكلّف تماماً ما تكلّفه مضاعفة طول السياق. أما التفعيلات فبلغت ذروتها 1.510 GiB. هي كلفة حقيقية، لكنها تُدفع مرة في كل تمرير، ثم تُستعاد.

يُفترض أن يُظهر السطر الأول القيمة 1.0000، أي أن المُخصِّص والحساب يتفقان حتى آخر بايت. بعد ذلك قارن آخر صفّين في الجدول: إنهما النموذج الأكبر نفسه، مرة برؤوس مفاتيح مشتركة ومرة بدونها.

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_table
شيفرة الورشة
الذاكرة المقيسة ÷ المعادلة، مجموعةً على كل نقاط المسح: 1.0000
النموذج  KiB لكل رمز  الأوزان GiB  رموز التعادل  \
0              Llama-2-7B (MHA)        512.0         12.6       25711.1   
1             Llama-2-13B (MHA)        800.0         24.2       31738.3   
2  Llama-2-70B (GQA، 8 رؤوس KV)        320.0        128.5      421142.6   
3        Llama-2-70B لو كان MHA       2560.0        128.5       52642.8   

   الذاكرة المؤقتة GiB (8 × 4096)  
0                            16.0  
1                            25.0  
2                            10.0  
3                            80.0

نسبة مقيسة قدرها 1 تعني أن المعادلة ليست تقديراً تقريبياً. إنها الحجم الدقيق للذاكرة المؤقتة. وهذا بالضبط ما يسمح لنا بتطبيقها على نماذج لا تستطيع هذه البطاقة تحميلها أصلاً. نبدأ بـ Llama-2-7B: تتساوى ذاكرته المؤقتة مع أوزانه عندما يبلغ عدد الرموز قيد المعالجة 25,711. فإذا خدم 8 مستخدمين بسياقه الكامل البالغ 4096، حملت ذاكرته المؤقتة 1.270 ضعف أوزانه. أما Llama-2-70B فأكبر منه بعشر مرات، ومع ذلك لا يصل إلى نقطة التعادل قبل 421,143. السبب في الصف الأخير من الجدول. لو كان لكل رأس استعلام رأس مفاتيح خاص به، لوصل إلى التعادل عند 52,643 فقط. إذن الفرق لا يأتي من حجم النموذج، بل من قرار تصميمي واحد في طريقة الانتباه.

كثيراً ما تُوضع تقنيتان تحت عنوان واحد: «جعل الانتباه أخفّ على الذاكرة». لكن كل واحدة منهما تخفّض فاتورة مختلفة. التقنية الأولى هي ، وهي تغيّر طريقة حساب الانتباه. فمصفوفة الدرجات، التي فيها خانة لكل زوج من الرموز، تُعالَج على شكل كتل صغيرة ولا تُخزَّن كاملة أبداً. وهذه المصفوفة تفعيل، لا جزء من الذاكرة المؤقتة. التقنية الثانية هي الانتباه المجمّع، وهي تغيّر ما نخزّنه. فعدة رؤوس استعلام تتشارك رأس مفاتيح واحداً ورأس قيم واحداً، فيقلّ عدد المفاتيح والقيم التي يجب حفظها. وهذا يمسّ الذاكرة المؤقتة مباشرة. الخليتان التاليتان تجرّبان كل تقنية وحدها. ثم تجرّب خلية ثالثة أداة تنفع مع الاثنتين: التخزين بعدد أقل من البتات.

اقرأ العمودين كلاً على حدة. أحدهما يتغيّر بفارق كبير، والآخر يُفترض ألا يتغيّر إطلاقاً.

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 × 2048 (تسلسلات × رموز). مصفوفة الدرجات لطبقة واحدة وحدها: 0.50 GiB
الانتباه  التفعيلات GiB  الذاكرة المؤقتة GiB
0  تقليدي (مصفوفة كاملة)           1.13                  1.0
1      مدمج، على شكل كتل           0.38                  1.0

ثلاثة نماذج لا تختلف إلا في عدد رؤوس المفاتيح التي تحتفظ بها: الانتباه متعدد الرؤوس، والانتباه المجمّع، وانتباه الاستعلامات المتعددة. قارن بين مقدار انخفاض الذاكرة المؤقتة ومقدار تغيّر الأوزان، وهو تغيّر صغير.

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  الأوزان GiB  الذاكرة المؤقتة GiB  تقلّص الذاكرة المؤقتة
0       16         1.75                 2.00                    1.0
1        4         1.56                 0.50                    4.0
2        1         1.52                 0.12                   16.0

ستكون النسبة أعلى من النصف بقليل، لأن كل مجموعة من القيم تحمل معامل تحجيم خاصاً بها. لاحظ مقدار الخطأ، وتذكّر أن أوزان هذا النموذج عشوائية.

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()
شيفرة الورشة
الذاكرة المؤقتة بدقة fp16: 2.00 GiB   وبدقة int8 مع معاملات التحجيم: 1.02 GiB   النسبة 0.508
خطأ إعادة البناء النسبي: 0.65%

خفّضت النواة المدمجة التفعيلات بعامل قدره 3. لكنها تركت الذاكرة المؤقتة عند 1 GiB، وهو الرقم نفسه في النسخة التقليدية. أما الانتباه المجمّع بـ4 رؤوس مفاتيح فقلّص الذاكرة المؤقتة بعامل قدره 4. وقلّصها انتباه الاستعلامات المتعددة بعامل قدره 16. وفي الحالتين لم تتغيّر الأوزان إلا بنحو 10.700٪. ثم يأتي أداةً ثالثة مستقلة عن الأوليين. فالتخزين بصيغة int8 أبقى 0.508 من حجم الذاكرة المؤقتة بصيغة ، بخطأ قدره 0.650٪. لكن هذا الخطأ حدّ أدنى، وليس توقعاً لما سيحدث في نموذج حقيقي. فالأوزان العشوائية تنتج مفاتيح منتظمة القيم. أما النماذج المدرَّبة فكثيراً ما تحوي قنوات قليلة بقيم كبيرة، ومعامل التحجيم لكل رمز يتعامل معها بدقة أقل. وأخيراً، هذه الأدوات تتضاعف آثارها معاً. فالرؤوس المجمّعة مع int8 تحفظ المحادثات نفسها في جزء صغير من الذاكرة المؤقتة الأصلية.

للذاكرة المؤقتة كلفة زمنية أيضاً. فتوليد رمز واحد يتطلب حساباً قليلاً جداً مقارنةً بعدد البايتات التي يلمسها. كل وزن يُقرأ مرة واحدة لمعالجة رمز واحد فقط من كل تسلسل. وكل مفتاح وقيمة في الذاكرة المؤقتة يُقرآن مرة واحدة ليتم الانتباه إليهما. لذلك لا تحدّد قدرةُ الحساب سرعةَ خطوة التوليد. ما يحدّدها هو ، أي سرعة نقل البايتات. وكلما كبرت الذاكرة المؤقتة، زادت البايتات التي يجب قراءتها في كل خطوة.

النموذج هو نفسه عند كل نقطة، والذي يطول هو الذاكرة المؤقتة وحدها. راقب كيف يبتعد الخطّان أحدهما عن الآخر كلما طال السياق.

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: 12.2 ms عند سياق 256، و42.1 ms عند سياق 4096 (أي 3.5 × أبطأ). ومع GQA عند سياق 4096: 18.7 ms.
معدّل القراءة الفعلي عند أطول سياق: 147 GB/s

مع الانتباه متعدد الرؤوس الكامل استغرقت الخطوة 12.200 ms عند أقصر سياق، و42.100 ms عند أطوله. أي أن الخطوة أبطأ بـ3.500 مرة، مع أن النموذج لم يتغيّر. الرقم الذي يستحق المقارنة هنا هو معدّل القراءة الفعلي: 147 GB/s. قارنه بعرض النطاق المعلن للبطاقة. إذا كان الرقمان متقاربين، فالخطوة تقضي وقتها في انتظار الذاكرة. وعندها لا طريق لتسريعها إلا قراءة بايتات أقل. وهذا بالضبط ما يفعله الانتباه المجمّع، ولذلك استغرقت خطوته 18.700 ms فقط عند السياق نفسه. إذن الذاكرة المؤقتة الأصغر تحسّن لكل رمز. ولأنها تسمح بخدمة تسلسلات أكثر في وقت واحد، فهي تحسّن أيضاً.

تمرين

لديك 10 GiB من ، وتريد خدمة محادثات طول كلٍّ منها 4096. قبل التشغيل، توقّع عدد المحادثات التي تتسع لها البطاقة عندما يكون عدد رؤوس المفاتيح 16. شغّل الخلية، ثم غيّر KV_HEADS إلى 4، ثم إلى 1. الخلية لا تثق بخطتها: فهي تحجز الذاكرة المؤقتة لكل التسلسلات فعلاً وتولّد رمزاً واحداً. لذلك تخبرك الذروة المقيسة إن كانت الخطة صادقة. وفي كل حالة، ما نسبة الميزانية التي تبقى للأوزان؟

# 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()
شيفرة الورشة

يتوفّر تلميح في الدفتر — env.hint(1)

نبدأ بفحص المعادلة لأن الفحصين الآخرين يعتمدان عليه. فإذا لم تطابق الذاكرة المؤقتة معادلتها، فالخطأ في طريقة القياس، لا في النموذج.

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)
شيفرة الورشة
✓ الذاكرة المؤقتة المقيسة مقسومةً على تنبؤ المعادلة، مجموعةً على كل نقاط المسح: 1 (المطلوب ≥ 0.99 و≤ 1.01)
✓ الذاكرة المؤقتة مقسومةً على الأوزان عند أكبر نقطة في المسح: 2.28 (المطلوب ≥ 1.5)
✓ ذاكرة الانتباه متعدد الرؤوس مقسومةً على ذاكرة الانتباه المجمّع، بالدفعة والسياق نفسيهما: 4 (المطلوب ≥ 3.8)

عدد المعاملات يقدّر كلفة النموذج نفسه، لكنه لا يقول شيئاً عن العمل الذي يؤديه. تُكتب بلغة المعاملات. أما فاتورة الخدمة فتُكتب بعدد الرموز قيد المعالجة، وهذه هي الفاتورة التي قسناها في هذه الورشة. فعندما تخطّط لعتاد، ابدأ من المعادلة التي تحققت منها هنا: 2 × عدد الطبقات × عدد رؤوس المفاتيح × حجم الرأس × عدد البايتات، لكل رمز. ثم اضربها في عدد المستخدمين وطول السياق الذي تتوقعه. بعد ذلك ستقرأ قرارات تصميم النماذج الحديثة بطريقة مختلفة. فالرؤوس المجمّعة، والانتباه المدمج، والذاكرة المؤقتة المكمّمة، والسياقات الأقصر، كلٌّ منها طريقة لدفع واحدة من هذه الفواتير.

receipt = env.receipt()
شيفرة الورشة
اكتملت الورشة.

رمز الإتمام: ⁦AZ-██████████⁩
الصقه في صفحة الورشة على أزيموث لتسجيل إتمامها.
آخر تحقّق: 2026-10-02 · Tesla T4 · PyTorch 2.11.0+cu130 · Python 3.13.15 · e4dc5b2

مصطلحات هذه الورشة