هذا المقال يتطلب مستوى: محترف
لو خدمتك بتولّد نص من نموذج لغة، وأول توكن بياخد وقت محسوس وباقي التوكنات بتطلع بسرعة، ده مش عشوائي. ده الـ KV Cache شغّال. المقال ده هيوريك بالظبط ليه بيحصل كده، إزاي توفّر تكلفة حسابية بترتيب O(n) بدل O(n²)، وإيه الثمن اللي بتدفعه في الذاكرة.
الـ KV Cache: ليه أول توكن بطيء وباقي التوكنات بتطير
المشكلة باختصار
نموذج اللغة بيولّد توكن واحد في المرة. عشان يطلّع التوكن الجديد، الانتباه الذاتي (Self-Attention) محتاج يبص على كل التوكنات اللي قبله. لو كل خطوة أعادت حساب المفاتيح (Keys) والقيم (Values) لكل التوكنات السابقة من الأول، تكون بتكرّر نفس الشغل ملايين المرات في نص طويل. النتيجة: التكلفة بتكبر تربيعيًا مع طول النص، والـ GPU بيتخنق من غير سبب حقيقي.
الفكرة ببساطة: المذيع اللي بيقرأ نشرة
تخيّل مذيع بيقرأ نشرة أخبار جملة ورا جملة. عشان ينطق الكلمة الجاية صح، هو محتاج يفتكر اللي قاله قبل كده. فيه طريقتين. الأولى الغبية: قبل كل كلمة جديدة، يرجع يقرأ النشرة كلها من أول سطر لغاية مكانه، وبعدين ينطق الكلمة الواحدة الجديدة. الثانية العاقلة: هو فاكر خلاص كل اللي فات، فبيضيف الكلمة الجديدة على طول من غير ما يعيد قراءة أي حاجة.
الطريقة الأولى هي التوليد بدون KV Cache. الثانية هي التوليد بالـ KV Cache. الفرق مش في الجودة، الفرق في إنك بتوفّر إعادة قراءة كل اللي فات في كل خطوة.
المفهوم علميًا
في طبقة الانتباه، كل توكن بيتحوّل لثلاث متجهات: Query و Key و Value. التوكن الجديد بيقارن الـ Query بتاعه بكل الـ Keys السابقة عشان يحسب أوزان الانتباه، وبعدين يجمّع الـ Values حسب الأوزان دي. الملاحظة المفتاحية: الـ Keys والـ Values بتاعة التوكنات القديمة ما بتتغيّرش لما نضيف توكن جديد. فبدل ما نحسبها تاني كل مرة، نحسبها مرة واحدة ونخزّنها في ذاكرة الـ GPU. ده هو الـ KV Cache.
عشان كده الاستدلال بينقسم لمرحلتين. الأولى Prefill: بنحسب K و V لكل توكنات الـ prompt دفعة واحدة، وده بياخد الوقت الأطول (أول توكن). الثانية Decode: كل توكن جديد بيحسب K و V بتاعه هو بس، بيضيفهم للـ cache، ويقرأ الباقي جاهز. ده اللي بيحصل فعلاً لما تحس إن أول توكن تقيل وباقي التوكنات خفيفة.
القياس بالكود
الكود ده بيقيس الفرق فعليًا على أي نموذج من Hugging Face. جرّبه بنفسك:
import time, torch
from transformers import AutoModelForCausalLM, AutoTokenizer
name = "meta-llama/Llama-2-7b-hf"
tok = AutoTokenizer.from_pretrained(name)
model = AutoModelForCausalLM.from_pretrained(name, torch_dtype=torch.float16).cuda()
prompt = tok("اشرح لي فكرة الـ KV Cache باختصار:", return_tensors="pt").to("cuda")
for use_cache in (True, False):
torch.cuda.synchronize(); t0 = time.time()
model.generate(**prompt, max_new_tokens=200, use_cache=use_cache)
torch.cuda.synchronize()
print(f"use_cache={use_cache}: {time.time()-t0:.2f}s")