وقتی یک مدل ۶۴ سر پرسش دارد و فقط ۸ سر کلید و مقدار، کش کلید-مقدار هر توکن از ۲۵۶۰ کیلوبایت به ۳۲۰ کیلوبایت می‌افتد؛ یعنی ۸ برابر. این همان کاری است که توجه گروهی می‌کند و دلیل آن است که Llama 3 هفتاد میلیاردی روی هشت کارت ۸۰ گیگابایتی سرو می‌شود. در این نوشته config.json هفت مدل واقعی را می‌خوانیم و با پایتون خالص نشان می‌دهیم هر سر پرسش کدام کلید را می‌خواند. لحظه‌ی خواندن: ۲۹ سپتامبر ۲۰۲۶.

مسئله: کش کلید-مقدار چرا از خود مدل مهم‌تر است

یک ترنسفورمر در حال تولید توکن، دو کار متفاوت انجام می‌دهد: ضرب‌های ریاضی، و خواندن داده از حافظه. در هر گام تولید، کل کلید و مقدار توکن‌های قبلی باید بارگذاری شود تا توجه محاسبه شود. در زمینه‌ی بلند، همین خواندن گلوگاه است، نه ضرب‌ها.

کش کلید-مقدار این خواندن را از محاسبه‌ی دوباره جدا می‌کند: کلید و مقدار هر توکن یک بار ساخته و ذخیره می‌شود. بهایش این است که حافظه‌ی لازم با طول زمینه رشد خطی دارد، در حالی که وزن‌های مدل ثابت‌اند. پس حافظه‌ی قابل استفاده را زمینه‌ی بلند و تعداد درخواست‌های هم‌زمان می‌خورند.

اندازه‌ی کش به ازای هر توکن برابر است با دو ضرب در تعداد لایه، ضرب در تعداد سرهای کلید-مقدار، ضرب در بعد هر سر، ضرب در دو بایت. ضریب دو یعنی هم کلید و هم مقدار ذخیره می‌شوند، و دو بایت اندازه‌ی هر عدد در دقت bf16 و fp16 است. تنها پارامتری که یک معماری می‌تواند آزادانه کم کند، تعداد سرهای کلید-مقدار است.

مقاله‌ی Shazeer در ۶ نوامبر ۲۰۱۹ اشتراک کلید و مقدار میان همه‌ی سرها را پیشنهاد کرد، ولی کیفیت افت کرد. دو سال بعد مقاله‌ی Ainslie و همکاران راهی میانی پیشنهاد کرد.

توجه گروهی چیست: یک پیچ تنظیم بین دو انتها

در توجه چندسر، هر سر پرسش کلید و مقدار خودش را دارد. اگر همه‌ی سرها یک جفت کلید-مقدار را به اشتراک بگذارند، به آن توجه چندپرسشی می‌گویند: کش به اندازه‌ی یک سر کوچک می‌شود، ولی کیفیت افت می‌کند. توجه گروهی دقیقا وسط این دو می‌ایستد: سرهای پرسش را به گروه‌های مساوی تقسیم می‌کند و هر گروه یک جفت کلید-مقدار می‌گیرد.

اندازه‌ی هر گروه را با نسبت تعداد سرهای پرسش بر تعداد سرهای کلید-مقدار می‌سنجند. اگر هر سر پرسش گروه خودش را داشته باشد، به توجه چندسر برمی‌گردیم؛ اگر همه در یک گروه باشند، به توجه چندپرسشی. هر چه تعداد گروه‌ها بیشتر باشد، کش بزرگ‌تر و کیفیت به چندسر نزدیک‌تر می‌شود.

نکته‌ای که در بیشتر توضیح‌های این معماری گم می‌شود این است: آنچه میان سرهای یک گروه به اشتراک گذاشته می‌شود، بردارهای ذخیره‌شده‌ی کلید و مقدار است، نه خودِ وزن‌های توجه. چون هر سر پرسش بردار خودش را دارد، نتیجه‌ی نهایی هر سر متفاوت می‌ماند و همین توضیح می‌دهد چرا این فشرده‌سازی کم‌هزینه بوده است.

بگذارید این را اندازه بگیریم. دو فایل کوچک، هیچ بسته‌ی بیرونی نمی‌خواهند.

# gqa_part2.py را ذخیره و اجرا کنید:
$ python3 gqa_part2.py

def softmax(xs):
    m = max(xs)
    e = [pow(2.718281828459045, x - m) for x in xs]
    t = sum(e)
    return [v / t for v in e]


n_q, n_kv, d = 4, 2, 3        # چهار سر پرسش، دو سر K/V
G = n_q // n_kv
q = [[0.10 * (h + 1) + 0.01 * j for j in range(d)] for h in range(n_q)]
k_stored = [[0.20, 0.05, 0.30], [0.11, 0.42, 0.07]]   # یک K برای هر گروه

print(f"== n_q={n_q} query heads, n_kv={n_kv} K/V heads, G={G} ==")
print(f"  K vectors actually stored: {len(k_stored)}  (MHA would store {n_q})")
attn = []
for h in range(n_q):
    g = h // G             # هر سر پرسش یک K مشترک می‌خواند
    s = [sum(qi * ki for qi, ki in zip(q[h], k)) for k in k_stored]
    p = softmax([x / (d ** 0.5) for x in s])   # همان مقیاس جذر بعد سر
    attn.append(p)
    print(f"  q head {h} reads K index {g}  attn={[round(x, 4) for x in p]}")
print(f"  attention rows still differ, because Q differs: {attn[0] != attn[1]}")
== n_q=4 query heads, n_kv=2 K/V heads, G=2 ==
  K vectors actually stored: 2  (MHA would store 4)
  q head 0 reads K index 0  attn=[0.4994, 0.5006]
  q head 1 reads K index 0  attn=[0.4987, 0.5013]
  q head 2 reads K index 1  attn=[0.498, 0.502]
  q head 3 reads K index 1  attn=[0.4972, 0.5028]
  attention rows still differ, because Q differs: True

چهار سطر میانی هسته‌ی کار را نشان می‌دهند. سر صفر و سر یک هر دو از کلید شماره‌ی صفر می‌خوانند، و سر دو و سر سه هر دو از کلید شماره‌ی یک. یعنی چهار سر پرسش، دو بردار کلید می‌خواهند نه چهار تا. سطر آخر به همان اندازه مهم است: وزن‌های توجه دو سر نخست یکسان نیستند، چون هر سر پرسش خودش را دارد. اگر این سطر صفر بود، فشرده‌سازی تنوع سرها را از بین می‌برد.

حساب واقعی روی هفت مدل

حالا که مکانیزم روشن است، سرهای واقعی مدل‌ها را از فایل منتشرشده‌ی خودشان می‌خوانیم. هیچ عددی در دو جدول بعدی از حافظه نیامده است.

# gqa_part1.py را ذخیره و اجرا کنید؛ خروجی همان دو ستون جدول است:
$ python3 gqa_part1.py

#!/usr/bin/env python3
"""gqa_part1.py: the table's per-token column, without reprinting the model's
own config.json in the output.  python3 gqa_part1.py
"""
import json, os, urllib.request

CACHE = ".gqa-cache"
MODELS = [
    ("NousResearch/Meta-Llama-3-70B-Instruct", "Llama 3 70B"),
    ("NousResearch/Meta-Llama-3-8B-Instruct",  "Llama 3 8B"),
    ("NousResearch/Llama-2-70b-hf",            "Llama 2 70B"),
    ("NousResearch/Llama-2-7b-hf",             "Llama 2 7B"),
    ("tiiuae/falcon-7b",                       "Falcon 7B"),
]


def fetch(repo):
    # یک بار دانلود، بعد از روی دیسک خوانده می‌شود
    os.makedirs(CACHE, exist_ok=True)
    path = os.path.join(CACHE, repo.replace("/", "_") + ".json")
    if os.path.exists(path):
        return json.load(open(path))
    req = urllib.request.Request(
        "https://huggingface.co/" + repo + "/raw/main/config.json",
        headers={"User-Agent": "Mozilla/5.0"})
    raw = urllib.request.urlopen(req, timeout=45).read()
    open(path, "wb").write(raw)
    return json.loads(raw)


def kv_bytes_per_token(layers, n_kv, head_dim, elem=2):
    # ضریب 2 یعنی یک تانسور برای K و یکی برای V
    return 2 * layers * n_kv * head_dim * elem


print("== KV cache per token, read from each config.json ==")
for repo, label in MODELS:
    c = fetch(repo)
    nq = c["num_attention_heads"]
    nkv = c.get("num_key_value_heads", nq)
    L = c["num_hidden_layers"]
    hd = c.get("head_dim") or c["hidden_size"] // nq
    per = kv_bytes_per_token(L, nkv, hd)
    mha = kv_bytes_per_token(L, nq, hd)
    print(f"  {label:12} nq={nq:3} nkv={nkv:3} L={L:3} hd={hd:4}"
          f"  {per/1024:7.1f} KiB/token   (MHA {mha/1024:7.1f} KiB)")

== KV cache per token, read from each config.json ==
  Llama 3 70B  nq= 64 nkv=  8 L= 80 hd= 128    320.0 KiB/token   (MHA  2560.0 KiB)
  Llama 3 8B   nq= 32 nkv=  8 L= 32 hd= 128    128.0 KiB/token   (MHA   512.0 KiB)
  Llama 2 70B  nq= 64 nkv=  8 L= 80 hd= 128    320.0 KiB/token   (MHA  2560.0 KiB)
  Llama 2 7B   nq= 32 nkv= 32 L= 32 hd= 128    512.0 KiB/token   (MHA   512.0 KiB)
  Falcon 7B    nq= 71 nkv= 71 L= 32 hd=  64    568.0 KiB/token   (MHA   568.0 KiB)
مدلنوعسر پرسشسر کلید-مقدارگروهلایهکش هر توکناگر چندسر بود
Llama 2 7Bچندسر۳۲۳۲۱۳۲۵۱۲ کیلوبایت۵۱۲ کیلوبایت
Llama 2 70Bگروهی۶۴۸۸۸۰۳۲۰ کیلوبایت۲٬۵۶۰ کیلوبایت
Llama 3 8Bگروهی۳۲۸۴۳۲۱۲۸ کیلوبایت۵۱۲ کیلوبایت
Llama 3 70Bگروهی۶۴۸۸۸۰۳۲۰ کیلوبایت۲٬۵۶۰ کیلوبایت
Mistral 7B v0.3گروهی۳۲۸۴۳۲۱۲۸ کیلوبایت۵۱۲ کیلوبایت
Qwen 2.5 72Bگروهی۶۴۸۸۸۰۳۲۰ کیلوبایت۲٬۵۶۰ کیلوبایت
Falcon 7Bچندسر۷۱۷۱۱۳۲۵۶۸ کیلوبایت۵۶۸ کیلوبایت

روش حساب در تابع kv_bytes_per_token بالا نوشته شده. برای Llama 3 هفتاد میلیاردی: ۲ ضرب در ۸۰ ضرب در ۸ ضرب در ۱۲۸ ضرب در ۲ که برابر است با ۳۲۷٬۶۸۰ بایت، یعنی همان ۳۲۰ کیلوبایت جدول بالا.

سطر نخست و آخر جدول، داستان مهاجرت را در خود دارند. Llama 2 هفت میلیاردی هنوز چندسر کامل است و ۵۱۲ کیلوبایت برای هر توکن می‌خواهد. نسخه‌ی هفتاد میلیاردی همان خانواده گروهی است و ۳۲۰ کیلوبایت. یعنی این گروه‌بندی در یک نسل وارد این خانواده شد، نه با انتشار یک مقاله، بلکه وقتی هزینه‌ی حافظه دیگر قابل تحمل نبود. Falcon 7B هم الگو را نپذیرفت و بهایش را پرداخت.

حالا همان ضریب را روی طول زمینه بگذارید. برای Llama 3 هفتاد میلیاردی، هر درخواست جداگانه چقدر کش لازم دارد؟

طول زمینهکش در توجه چندسرکش در توجه گروهیآزادشده
۸٬۱۹۲ توکن۲۰ گیگابایت۲٫۵ گیگابایت۱۷٫۵ گیگابایت
۳۲٬۷۶۸ توکن۸۰ گیگابایت۱۰ گیگابایت۷۰ گیگابایت
۱۳۱٬۰۷۲ توکن۳۲۰ گیگابایت۴۰ گیگابایت۲۸۰ گیگابایت

سطر دوم همان چیزی است که در عمل رخ می‌دهد: یک کارت ۸۰ گیگابایتی با ۸۰ گیگابایت کش، جایی برای نگه‌داری وزن‌های ۱۴۰ گیگابایتی مدل باقی نمی‌گذارد. اگر این مدل چندسر بود، یک درخواست با زمینه‌ی ۳۲ هزار توکن به تنهایی کل حافظه را پر می‌کرد.

در کد واقعی، یک نکته‌ی عملی هست. در پیاده‌سازی لایما روی شاخه‌ی main، تابع repeat_kv بردارهای کلید و مقدار را از تعداد سر کلید-مقدار به تعداد سر پرسش بسط می‌دهد تا ضرب ماتریسی شکل یکنواخت بگیرد. خودِ تابع می‌گوید این کار معادل torch.repeat_interleave است.

پس دو نکته را از هم جدا کنید. کش ذخیره‌شده کوچک است، چون فقط سرهای کلید-مقدار واقعی نگه‌داری می‌شوند؛ ولی در محاسبه، بردارها موقتا بزرگ می‌شوند. اگر کارت گرافیکی شما کم‌حافظه است و خطای کمبود حافظه می‌دهد، همین بسط موقت می‌تواند یکی از علت‌ها باشد.

در مقاله‌ی GQA آمده که تبدیل یک checkpoint چندسر به حالت گروهی تنها ۵ درصد از توان آموزش اولیه را می‌خواهد. یعنی تبدیل ارزان است، نه این‌که کیفیت رایگان به دست می‌آید.

اگر تازه با معماری ترنسفورمر آشنا شده‌اید، نوشته‌ی ما درباره‌ی سازوکار توجه توضیح می‌دهد چه چیزی در هر لایه اتفاق می‌افتد و نوشته‌ی RMSNorm سراغ تکه‌ی دیگری از همان بلوک رفته است. توجه گروهی سومین تکه است.

جمع‌بندی

قاعده‌ی عملی که از این نوشته بیرون می‌آید این است: پیش از انتخاب مدل برای زمینه‌ی بلند، تعداد سر کلید-مقدار را از config.json همان مدل بخوانید و با تابع kv_bytes_per_token حافظه‌ی هر درخواست را حساب کنید. در هفت مدلی که اینجا خواندیم، پنج مدل گروهی و دو مدل چندسر بودند، و هر پنج مدل گروهی یا ۴ برابر یا ۸ برابر کش کمتری به ازای هر توکن می‌خواستند.

منابع

  1. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints — arXiv:2305.13245، ۲۲ مه ۲۰۲۳
  2. Fast Transformer Decoding: One Write-Head is All You Need — arXiv:1911.02150، ۶ نوامبر ۲۰۱۹
  3. تابع repeat_kv در پیاده‌سازی لایما
  4. config.json مدل Llama 3 70B
  5. config.json مدل Llama 3 8B
  6. config.json مدل Llama 2 70B
  7. config.json مدل Llama 2 7B
  8. config.json مدل Mistral 7B v0.3
  9. config.json مدل Falcon 7B
  10. config.json مدل Qwen 2.5 72B