Chuyển tới nội dung chính

KV cache, hay vì sao sinh chữ là bài toán bộ nhớ

Sáu tham số sửa đượcKế toán từng byteTính lại thật

KV cache, hay vì sao sinh chữ là bài toán bộ nhớ

Ai cũng nghĩ chạy một mô hình lớn là chuyện thiếu sức tính. Thực tế thứ chặn bạn trước tiên là bộ nhớ, và phần lớn nó không nằm ở trọng số.

Bạn đã học self-attention: để tính đầu ra cho một token, mô hình cần vector Q của token đó nhân với vector K của mọi token trước nó, rồi lấy trọng số đó để trộn các vector V. Khi sinh chữ, mô hình làm việc này lặp lại: sinh token thứ 1000 thì phải chú ý về 999 token cũ, sinh token thứ 1001 thì phải chú ý về 1000 token cũ.

Nếu mỗi bước lại tính KV cho toàn bộ quá khứ thì cùng một phép nhân bị làm lại hàng nghìn lần. Nên không ai làm vậy: tính một lần rồi giữ lại. Chỗ giữ đó gọi là KV cache. Đổi lại, cái giá không còn là thời gian mà là bộ nhớ, và nó là loại bộ nhớ phình ra theo từng token bạn sinh thêm. Trọng số thì nằm im một chỗ, nạp lúc khởi động là xong. Bộ đệm thì mỗi token một ít, mãi cho tới khi hết chỗ.

Bảng tính dưới đây là toàn bộ hoá đơn đó. Bạn đặt hình dạng mô hình và cách phục vụ, nó tính lại từng thành phần.

Hoá đơn bộ nhớ của KV cache · sinh chữ là bài toán bộ nhớ
Byte mỗi token 128,0 KiBTổng KV cache 16,0 GiBSo với trọng số 107,4%
Cấu hình mẫu:
Llama 3.1 8B Instruct: 32 đầu truy vấn, 8 đầu KV. Nguồn: config.json công bố trên Hugging Face, đọc ngày 27/07/2026. Số tham số làm tròn về 8 tỉ. Lĩnh vực này đổi nhanh nên số có thể đã lạc hậu.
Cấu hình✎ sửa được
số lớp
mỗi lớp giữ một bộ K và V riêng của nó
số đầu KV
MHA thì bằng số đầu truy vấn, GQA thì nhỏ hơn
chiều mỗi đầu
độ dài một vector K hoặc V, thường 64 hoặc 128
độ dài ngữ cảnh
số token đang nằm trong bộ đệm
cỡ lô
số chuỗi phục vụ cùng lúc, mỗi chuỗi một bộ đệm riêng
số tham số (tỉ)
chỉ dùng để so bộ đệm với trọng số
byte mỗi token2 × 32 × 8 × 128 × 2 = 131.072 byte
tổng KV cache131.072 × 131.072 × 1 = 17.179.869.184 byte
trọng số mô hình8.000.000.000 × 2 = 16.000.000.000 byte
128,0 KiB
byte mỗi token
131.072 byte
16,0 GiB
tổng KV cache
1 chuỗi
107,4%
bộ đệm so với trọng số
bộ đệm nặng hơn
14,9 GiB
trọng số mô hình
hằng số, không đổi theo token
74.537
ngữ cảnh dài nhất còn vừa
ngân sách 9,1 GiB
30,9 GiB
trọng số cộng bộ đệm
thẻ có 24,0 GiB
KV cache theo độ dài ngữ cảnh, cắt ngang bởi ngân sách của thẻ
không còn vừangân sách bộ đệm 9,1 GiB128k token · 16,0 GiB032k64k96k128k0 byte4,4 GiB8,8 GiB13,2 GiB17,6 GiBđộ dài ngữ cảnh (token)KV cache của cả lô
Ngữ cảnh 131.072 token không vừa: bộ đệm cần 16,0 GiB mà ngân sách chỉ có 9,1 GiB. Dài nhất còn vừa là 74.537 token.
Ngữ cảnhKV cache cả lôVừa 9,1 GiB?độ lớn
4.096512,0 MiBvừa
8.1921,0 GiBvừa
32.7684,0 GiBvừa
131.07216,0 GiBkhông
524.28864,0 GiBkhông
1.048.576128,0 GiBkhông

Vì sao công thức bắt đầu bằng số 2

Mỗi token, ở mỗi lớp, mỗi đầu chú ý cần giữ hai vector: một K và một V. Không phải một, cũng không phải ba. Chỉ vậy thôi, và đó là tất cả ý nghĩa của số 2 đứng đầu công thức:

byte mỗi token = 2 × số lớp × số đầu KV × chiều mỗi đầu × byte mỗi phần tử

Rồi nhân tiếp với hai thứ không thuộc về mô hình mà thuộc về cách bạn dùng nó:

tổng KV cache = byte mỗi token × độ dài ngữ cảnh × cỡ lô

Chia sáu thừa số này thành hai nhóm thì mọi thứ sáng ra. Số 2 là hằng số của cơ chế. Số lớp, số đầu KVchiều mỗi đầu là hình dạng kiến trúc, chốt xong từ lúc huấn luyện nên bạn không đổi được nữa. Còn byte mỗi phần tử, độ dài ngữ cảnhcỡ lô là quyết định lúc chạy, tức ba núm bạn xoay được ngay hôm nay mà không cần huấn luyện lại gì cả. Ba núm đó là toàn bộ dư địa của người triển khai.

Chú ý là không có Q trong danh sách: vector Q chỉ dùng đúng ở bước sinh ra nó rồi bỏ, không ai cần nó nữa nên không lưu.

Cũng chú ý số đầu ở đây là số đầu giữ K và V, không phải số đầu truy vấn. Với multi-head attention cổ điển thì hai con số bằng nhau. Với grouped query attention, nhiều đầu truy vấn dùng chung một bộ KV, nên số đầu KV nhỏ hơn hẳn. Đó là chỗ tiết kiệm lớn nhất trong bảng, và là chủ đề của bài kế tiếp.

Đọc lại con số ở trạng thái mặc định

Bảng mở ra với hình dạng của Llama 3.1 8B: 32 lớp, 8 đầu KV, chiều mỗi đầu 128, bộ đệm FP16 tức 2 byte một số.

2 × 32 × 8 × 128 × 2 = 131.072 byte

Tức 128 KiB cho mỗi token. Nghe nhỏ. Nhưng mô hình này quảng cáo ngữ cảnh 131.072 token, và nhân lên thì:

131.072 × 131.072 × 1 = 17.179.869.184 byte

Đúng 16,0 GiB. Trong khi trọng số của nó, 8 tỉ tham số ở FP16, chỉ là 8.000.000.000 × 2 = 16.000.000.000 byte, tức 14,9 GiB. Ô tỉ lệ ghi 107,4%: ở ngữ cảnh tối đa của chính nó, bộ đệm đã nặng hơn cả mô hình.

Đây là điều đáng dừng lại một chút. Bạn tải về một tệp trọng số 15 GiB và tưởng đó là chi phí. Không phải. Nó chỉ là phần không đổi. Phần đổi được thì lớn hơn, và nó do bạn quyết định.

Kéo thanh độ dài ngữ cảnh xuống 8.192 token: tỉ lệ tụt còn 6,7%, bộ đệm chỉ còn 1,0 GiB. Kéo dần lên, tỉ lệ lật qua mốc 100% ở khoảng 122.070 token. Trọng số không nhích một byte trong suốt lúc đó. Chỉ có mẫu số đứng yên còn tử số leo, tuyến tính, không có chỗ nào bẻ cong.

Điểm cắt: nơi ngữ cảnh không còn vừa nữa

Trên biểu đồ có một đường ngang. Nó không phải VRAM của thẻ, nó là VRAM trừ đi trọng số, vì trọng số đã ngồi trong bộ nhớ trước khi token đầu tiên xuất hiện. Với thẻ 24 GiB và mô hình 8B ở FP16, ngân sách còn lại cho bộ đệm là:

25.769.803.776 trừ 16.000.000.000 = 9.769.803.776 byte

tức 9,1 GiB. Chia cho 131.072 byte mỗi token được 74.537,6875, lấy phần nguyên là 74.537 token. Đó là chỗ biểu đồ đổi màu, và bên phải vạch đó là vùng không còn vừa. Con số này chỉ bằng 56,9% ngữ cảnh mà mô hình quảng cáo. Nói cách khác: thẻ chạy được mô hình, nhưng không chạy được mô hình ở ngữ cảnh của nó.

Nếu bạn quên trừ trọng số thì phép chia cho ra 196.608 token, cao hơn thực tế gần ba lần. Lỗi này rất hay gặp trong các bảng ước lượng chép tay.

Đổi thẻ sang 16 GiB để thấy chỗ khó chịu hơn: trọng số 14,9 GiB vẫn vừa, nhưng ngân sách còn lại chỉ 1,1 GiB, và ngữ cảnh dài nhất còn vừa tụt xuống 9.001 token. Mô hình vào được thẻ mà gần như không còn chỗ để làm việc. Đây là lý do câu "mô hình này chạy được trên card X" gần như không có nghĩa gì nếu không kèm độ dài ngữ cảnh và cỡ lô.

Ba cách làm hoá đơn nhỏ lại, và giá của từng cách

Giảm số đầu KV. Bấm preset Mô hình 7B kiểu MHA (minh hoạ): cùng 32 lớp, cùng chiều 128, nhưng 32 đầu KV thay vì 8. Byte mỗi token nhảy từ 128 KiB lên 512 KiB, gấp bốn lần, và ở ngữ cảnh chỉ 32.768 token nó đã ngốn 16,0 GiB với tỉ lệ 122,7%. Bốn đầu truy vấn dùng chung một bộ KV là cắt bộ đệm đi bốn lần, đúng bằng tỉ số nhóm. Giá phải trả là gì thì bảng này không nói được, vì phải huấn luyện mới biết.

Giảm số byte mỗi phần tử. Đổi kiểu số bộ đệm từ FP16 sang INT8: byte mỗi token thành 64 KiB, tổng thành 8,0 GiB, tỉ lệ về 53,7%, và ngữ cảnh dài nhất còn vừa trên thẻ 24 GiB nhảy lên 149.075 token. Lúc này ngữ cảnh 131.072 đã vừa. Lưu ý FP8 và INT8 cùng là 1 byte nên cùng cho một hoá đơn: chúng khác nhau ở chỗ biểu diễn được cái gì, không khác nhau ở chỗ chiếm bao nhiêu. Giá phải trả là sai số lượng tử hoá, và bảng này cũng không đo được nó.

Giảm cỡ lô. Đây là chỗ trực giác hay sai nhất. Cỡ lô nhân thẳng vào tổng: đặt cỡ lô 8 ở cấu hình mặc định thì bộ đệm thành 128,0 GiB, tỉ lệ 859,0%. Trọng số vẫn dùng chung cho cả tám chuỗi, còn bộ đệm thì mỗi chuỗi một bộ riêng. Nên khi ai đó khoe thông lượng cao nhờ cỡ lô lớn, câu hỏi tiếp theo luôn là bộ đệm nằm đâu.

Cái bảng này không nói gì

Đây là chỗ phải nói thẳng, vì một hoá đơn tính đúng vẫn có thể bị đọc quá xa.

Đây là kế toán cho suy luận có bộ đệm, không phải bộ nhớ lúc huấn luyện. Huấn luyện là hoá đơn khác hẳn: phải giữ activation cho lượt truyền ngược, thêm gradient, thêm trạng thái của bộ tối ưu, và một mô hình vừa suy luận thoải mái có thể không huấn luyện nổi trên cùng cái thẻ. Không con số nào ở trang này áp cho huấn luyện được.

Nó đếm đúng các tensor KV, không đếm phần dư. Một máy chủ thật còn tốn đệm của bộ cấp phát, phân mảnh, vùng làm việc tạm, ngữ cảnh CUDA. Nên số đo thực tế luôn lớn hơn con số ở đây một chút. Coi nó là sàn, không phải dự báo.

Nó giả định mọi lớp lưu bộ đệm theo cùng một cách. Có kiến trúc không như vậy: cửa sổ trượt chỉ giữ một đoạn gần đây thay vì cả quá khứ, còn multi-head latent attention nén KV xuống một vector tiềm ẩn nên công thức nhân đơn giản ở trên không áp được. Cả hai là bài riêng trong chương này.

Số trong preset có ngày và có thể đã lạc hậu. Bốn preset mang tên mô hình thật lấy hình dạng từ bản cấu hình mà mỗi dự án công bố, ngày đọc ghi ngay trên giao diện; số tham số đã làm tròn về con số mà mô hình được gọi tên. Preset thứ năm ghi rõ là bịa ra để minh hoạ. Lĩnh vực này đổi rất nhanh, nên nếu bạn có số mới hơn thì gõ thẳng vào, cả sáu ô đều sửa được. Cổng kiểm số của bài canh phép tính, tức là với cấu hình như vậy thì con số phải là như vậy; nó không canh và không thể canh việc mô hình ngoài kia có đúng hình dạng đó hay không.

Điều rút ra

Trọng số là hằng số, KV cache thì không: nó bằng 2 × số lớp × số đầu KV × chiều mỗi đầu × byte mỗi phần tử cho mỗi token, rồi nhân với độ dài ngữ cảnh và cỡ lô. Vì mọi thừa số vào tuyến tính nên không có chỗ nào cho một mẹo nhỏ tạo phép màu: muốn ngữ cảnh dài thì phải bớt một thừa số nào đó đi thật. Và vì trọng số chiếm VRAM trước, ngân sách cho bộ đệm luôn nhỏ hơn con số ghi trên vỏ thẻ. Gần như mọi cơ chế trong các bài sau của chương này là một cách trả giá cho đúng một dòng nhân đó.

Câu hỏi tự kiểm0/3 đúngchưa trả lời
  1. 1Một mô hình có 32 lớp, 8 đầu KV, chiều mỗi đầu 128, bộ đệm lưu ở FP16. Mỗi token tốn bao nhiêu byte trong KV cache?
  2. 2Ở cấu hình mặc định, tỉ lệ bộ đệm so với trọng số là 107,4%. Bạn đổi kiểu số bộ đệm từ FP16 sang INT8 và không sửa gì khác. Tỉ lệ đó thành bao nhiêu?
  3. 3Thẻ 24 GiB, mô hình 8 tỉ tham số ở FP16, mỗi token tốn 128 KiB. Vì sao ngữ cảnh dài nhất còn vừa là 74.537 token, không phải 196.608 token?