وقتی یک مدل ۶۴ سر پرسش دارد و فقط ۸ سر کلید و مقدار، کش کلید-مقدار هر توکن از ۲۵۶۰ کیلوبایت به ۳۲۰ کیلوبایت میافتد؛ یعنی ۸ برابر. این همان کاری است که توجه گروهی میکند و دلیل آن است که 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 حافظهی هر درخواست را حساب کنید. در هفت مدلی که اینجا خواندیم، پنج مدل گروهی و دو مدل چندسر بودند، و هر پنج مدل گروهی یا ۴ برابر یا ۸ برابر کش کمتری به ازای هر توکن میخواستند.
منابع
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints — arXiv:2305.13245، ۲۲ مه ۲۰۲۳
- Fast Transformer Decoding: One Write-Head is All You Need — arXiv:1911.02150، ۶ نوامبر ۲۰۱۹
- تابع repeat_kv در پیادهسازی لایما
- config.json مدل Llama 3 70B
- config.json مدل Llama 3 8B
- config.json مدل Llama 2 70B
- config.json مدل Llama 2 7B
- config.json مدل Mistral 7B v0.3
- config.json مدل Falcon 7B
- config.json مدل Qwen 2.5 72B
دیدگاهها
۰ موردهنوز دیدگاهی ثبت نشده. اولین نفر باشید.