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

MLA và chuyện nén bộ đệm bằng một chiếu xuống chiều thấp

Bảy tham số sửa đượcKế toán chính xác từng byteCa vàng sai số bằng 0

MLA và chuyện nén bộ đệm bằng một chiếu xuống chiều thấp

Nếu bộ đệm KV là thứ ăn hết VRAM khi ngữ cảnh dài ra, thì tại sao không lưu một bản nén của nó rồi dựng lại khi cần? Đó đúng là ý tưởng của Multi-head Latent Attention, và cái giá phải trả tính được bằng số.

Ở bài KV cache bạn đã thấy vì sao sinh chữ là bài toán bộ nhớ chứ không phải bài toán tính: mỗi token mới phải nhớ K và V của mọi token cũ, ở mọi lớp, cho mọi đầu. Ở bài MHA tới MQA tới GQA bạn đã thấy cách chữa đầu tiên, tức là bớt số đầu dùng để lưu K và V, cực đoan nhất là một đầu duy nhất. Cách đó rẻ và hiệu quả nhưng nó chỉ có vài mức: bạn chọn được 8 nhóm hay 4 nhóm, chứ không có núm để vặn liên tục.

Bài này là cách chữa thứ hai, và nó tới từ một chỗ bạn đã học rồi: ma trận biến hình. Nếu một vector d chiều nhân với một ma trận r hàng d cột thì ta được một vector r chiều, và nếu r nhỏ hơn d thì ta vừa nén nó lại. Multi-head Latent Attention trong họ DeepSeek làm đúng thế với bộ đệm: thay vì lưu K và V đầy đủ cho từng đầu ở từng lớp, nó chiếu xuống một vector tiềm ẩn chiều thấp, lưu đúng vector đó, rồi khi cần tính chú ý thì chiếu lên lại để dựng K và V. Núm ở đây là số r, tức chiều tiềm ẩn, và nó vặn được liên tục.

MLA · nén bộ đệm KV bằng một chiếu xuống chiều thấp
Nén ×56.89 · bộ đệm 67.50 KiB mỗi token
DeepSeek-V2 (236B)số của mô hình có thật
Kiểm tra chéo được: 60 lớp × (512 + 64) = 34.560 số mỗi token, đúng con số 34,56K mà báo cáo kỹ thuật của DeepSeek-V2 nêu.
Số kiến trúc của hai preset DeepSeek lấy từ tệp config.json phát hành kèm mô hình, đối chiếu ngày 27/07/2026: sáu trường của V3 khớp đúng num_hidden_layers 61, hidden_size 7168, num_attention_heads 128, v_head_dim 128, kv_lora_rank 512 và qk_rope_head_dim 64. Lĩnh vực này đổi rất nhanh nên hãy coi đây là ảnh chụp một thời điểm: số có thể đã lạc hậu, và bạn nên tự đối chiếu lại với nguồn. Mọi tham số dưới đây sửa được, nên có số mới hơn thì gõ vào là tính lại ngay. Độ dài ngữ cảnh và kiểu số cố tình không nằm trong preset: đó là hai thứ của bạn.
1 · Bộ đệm mỗi token: đầy đủ so với nén
Đại lượngMHA đầy đủMLATỉ lệ
Số phần tử mỗi token mỗi lớp2 × 128 × 128 = 32.768512 + 64 = 576×56.89
Byte mỗi token mỗi lớp65.5361.152×56.89
Byte mỗi token, cả 60 lớp3.75 MiB67.50 KiB×56.89
Cả ngữ cảnh 4.096 token15.00 GiB270.00 MiBtiết kiệm 14.74 GiB
hệ số nén = 2 × n_h × d_h / (d_c + d_rope) = 2 × 128 × 128 / (512 + 64) = 56.8889bỏ số hạng RoPE thì xấp xỉ 2 × n_h × d_h / d_c = 64.0000
2 · Cái giá phải trả: tham số và phép tính thêm
Ma trậnCông thứcTham số mỗi lớpCả 60 lớp
MHA: W_KW_V2 × 5120 × 128 × 128167.772.16010.066.329.600
MLA chiếu xuống: W_DKV5120 × 5122.621.440157.286.400
MLA chiếu lên: W_UKW_UV2 × 512 × 128 × 12816.777.2161.006.632.960
Khoá RoPE tách riêng: W_KR5120 × 64327.68019.660.800
Chênh lệch MLA trừ MHAtổng MLA − tổng MHA-148.045.824-8.882.749.440
Với cấu hình này bốn ma trận của MLA lại ít hơn 8.882.749.440 tham số so với hai ma trận W_KW_V của MHA đầy đủ, vì d_c nhỏ hơn d_model khá nhiều. Đừng đọc thành nén là miễn phí: cái thật sự đắt lên là phép tính mỗi bước sinh. Mỗi token, mỗi lớp phải nhân lại vector tiềm ẩn để dựng K và V, tốn 512 × 2 × 128 × 128 = 16.777.216 phép nhân cộng, tức 1.006.632.960 cho cả 60 lớp, là việc MHA không phải làm. Kéo d_c lên gần d_model thì dòng chênh lệch sẽ đổi dấu.
3 · Phép chiếu thật sự làm gì, tính bằng số thật
Vector tiềm ẩn được lưuz0 = 10.960z1 = 6.027z2 = 1.115
8.0001234567
vector gốc x vector dựng lại x_hat
Số lưu: 8 chiều xuống 3 chiều×2.667
Độ dài vector gốc ||x||12.609520
Sai số dựng lại ||x - x_hat||1.145214405
Sai số tương đối9.0821%
Đang giữ 3 trong 8 chiều, tức bộ đệm nhỏ đi 2.667 lần, và cái giá là sai số dựng lại 9.082%. Kéo r lên tới 8 thì sai số về đúng 0.
Chỉ sốxx_hatHiệu
08.0007.3460.654440
16.0006.594-0.593851
25.0005.336-0.335708
34.0003.9480.052301
43.0002.7720.228030
52.0001.9880.012485
62.0001.5830.417074
71.0001.435-0.434771
Sai số theo từng chiều tiềm ẩn, cùng vector và cùng ma trận
rNén được||x - x_hat||Sai số tương đốiđộ lớn
1×8.0006.23498195749.447%
2×4.0001.59851007812.677%
3×2.6671.1452144059.082%
4×2.0000.8468028666.716%
5×1.6000.7694641606.102%
6×1.3330.3379135452.680%
7×1.1430.3284874392.605%
8r = d×1.0000.0000000000.000%
Sim này không đo được cái gì. Ma trận chiếu ở phần 3 do tôi đặt cố định theo một công thức đóng, còn mô hình thật học cả hai ma trận xuống và lên từ dữ liệu. Ma trận học được biết những hướng nào thật sự xuất hiện trong K và V nên sai số dựng lại của mô hình thật nhỏ hơn nhiều con số bạn thấy ở đây, và không có cách nào suy nó ra từ sim. Thêm nữa, phép chiếu lên ở đây là giả nghịch đảo của phép chiếu xuống, tức cách dựng lại tốt nhất khi ta coi mọi hướng trong không gian quan trọng như nhau, còn mô hình thật học hai ma trận riêng biệt. Kết luận: phần kế toán bộ nhớ ở mục 1 và 2 chính xác, phần chất lượng ở mục 3 chỉ là minh hoạ cơ chế.
Hai đơn giản hoá nữa cần nói thẳng. Thứ nhất, trong bản gốc phần mang thông tin vị trí được tách riêng và không nén cùng vector tiềm ẩn, nên công thức kế toán có số hạng + d_rope: đó là lý do ô Chiều RoPE d_rope tồn tại và hệ số nén thật luôn nhỏ hơn con số xấp xỉ 2 × n_h × d_h / d_c. Thứ hai, bản gốc còn chia khoá thành phần nén và phần vị trí với hai bề rộng khác nhau, và có thể gộp W_UK vào ma trận truy vấn để không bao giờ dựng K ra; ở đây tôi dùng một bề rộng d_h cho cả K và V và không gộp. Đó là đơn giản hoá của bài này, không phải cách DeepSeek cài.

Con số đáng nhìn nhất khi mở trang

Cấu hình mặc định là hình dạng của DeepSeek-V2: 60 lớp, 128 đầu, mỗi đầu 128 chiều, chiều tiềm ẩn 512, phần RoPE tách riêng 64, kiểu số bf16.

Nhìn dòng đầu bảng thứ nhất. MHA đầy đủ phải lưu 2 × 128 × 128 bằng 32.768 số cho mỗi token ở mỗi lớp: nhân 2 vì có cả K và V. MLA lưu 512 + 64 bằng 576 số. Cùng một token, cùng một lớp, một bên 32.768 con số, một bên 576. Nhân lên cả 60 lớp và đổi sang byte thì thành 3,75 MiB mỗi token so với 67,50 KiB mỗi token. Với ngữ cảnh 4096 token, đó là khoảng cách giữa 15,00 GiB270,00 MiB, tiết kiệm 14,74 GiB. Con số 15 GiB kia mới là điều đáng dừng lại: nó chỉ là bộ đệm, chưa tính một byte trọng số nào, và nó lớn hơn cả bộ nhớ của phần lớn card tiêu dùng.

Có một cách kiểm tra chéo dễ chịu ở đây. 60 × 576 bằng 34.560 số cho mỗi token, và đó đúng là con số 34,56K mà báo cáo kỹ thuật của DeepSeek-V2 nêu ra. Nếu bảng của bạn cho số khác thì một trong hai bên sai, và ta biết chỗ để đi tìm.

Vì sao hệ số nén xấp xỉ hai lần số đầu nhân chiều mỗi đầu chia chiều tiềm ẩn

Đây là điểm cần nhớ của cả bài, và nó ra từ một phép chia. Số lượng lưu của MHA cho mỗi token mỗi lớp là 2 × n_h × d_h. Số lượng lưu của MLA là d_c. Chia hai bên cho nhau:

hệ số nén xấp xỉ 2 × n_h × d_h / d_c

Với cấu hình mặc định, đó là 2 × 128 × 128 / 512 bằng đúng 64. Chú ý cả số lớp và kiểu số đều triệt tiêu trong phép chia này: cả hai bên đều nhân với số lớp, cả hai bên đều nhân với số byte mỗi phần tử. Cho nên chuyển từ bf16 sang fp8 làm cả hai cột nhỏ đi một nửa mà hệ số nén không nhích một li. Bạn gạt ô kiểu số rồi nhìn cột tỉ lệ để tự thấy.

Nhưng ô tỉ lệ thật lại ghi 56,89 chứ không phải 64, và chênh lệch đó không phải lỗi làm tròn. Nó là số hạng d_rope.

Một chi tiết thật: phần mang vị trí không nén cùng

Trong bản gốc của MLA, phần khoá mang thông tin vị trí (RoPE) được tách riêng ra và không nén cùng vector tiềm ẩn. Lý do rất kỹ thuật: RoPE quay vector khoá theo vị trí của token, và phép quay đó không giao hoán được với phép chiếu lên, nên nếu nén chung thì mỗi lần dựng lại sẽ phải quay lại từ đầu và mất hết lợi thế. Cách chữa của DeepSeek là để một phần nhỏ của khoá đi đường riêng, không nén, và ghép vào lúc tính chú ý.

Hệ quả là công thức kế toán có một số hạng cộng thêm:

hệ số nén = 2 × n_h × d_h / (d_c + d_rope)

Với mặc định là 2 × 128 × 128 / (512 + 64) bằng 56,8889. Số hạng d_rope luôn làm hệ số nén nhỏ hơn con số xấp xỉ, chứ không bao giờ lớn hơn. Muốn thấy nó biến mất thì đặt ô Chiều RoPE d_rope về 0: hai con số 56,8889 và 64 lập tức trùng nhau. Cấu hình mẫu tròn số cũng để d_rope bằng 0 chính vì thế, và ở đó hệ số nén là 2 × 4 × 16 / 32 bằng đúng 4, nhân chia được trên giấy.

Nén không miễn phí, và cái giá không nằm ở chỗ bạn tưởng

Bảng thứ hai là chỗ dễ bị kể sai nhất, nên hãy đọc con số thật thay vì nghe một câu chuyện gọn gàng.

MHA đầy đủ cần hai ma trận cho khoá và giá trị, tổng 2 × d_model × n_h × d_h tham số mỗi lớp. MLA cần bốn ma trận: một chiếu xuống W_DKV cỡ d_model × d_c, hai chiếu lên W_UKW_UV mỗi cái d_c × n_h × d_h, và một W_KR cỡ d_model × d_rope cho phần vị trí tách riêng. Bốn ma trận nghe nhiều hơn hai, nên câu chuyện thường được kể là "MLA đổi tham số lấy bộ nhớ".

Ở cấu hình mặc định thì câu chuyện đó sai. Tổng của MHA là 10.066.329.600 tham số, tổng của MLA là 1.183.580.160, chênh lệch âm 8.882.749.440. MLA vừa nhẹ bộ đệm hơn 56 lần, vừa ít tham số hơn khoảng 8,5 lần. Lý do đơn giản: DeepSeek-V2 có n_h × d_h bằng 16.384, gấp hơn ba lần d_model bằng 5120, nên một MHA đầy đủ trên ngân sách đầu đó sẽ khổng lồ, và chính vì nó khổng lồ nên người ta mới đi tìm cách khác.

Vậy cái giá thật nằm ở đâu? Ở phép tính mỗi bước sinh. Mỗi token, mỗi lớp, MLA phải nhân vector tiềm ẩn với hai ma trận chiếu lên để dựng lại K và V, tốn 512 × 2 × 128 × 128 bằng 16.777.216 phép nhân cộng, tức 1.006.632.960 cho cả 60 lớp, cho mỗi một token. MHA không phải làm việc đó: nó đọc thẳng từ bộ đệm. Đây là kiểu đánh đổi rất hiện đại, đổi băng thông bộ nhớ lấy phép tính, và nó đáng vì trên GPU hiện nay phép tính thì rẻ còn băng thông thì đắt.

Và dấu của dòng chênh lệch đổi được. Chọn cấu hình mẫu tròn số rồi kéo ô Chiều tiềm ẩn d_c từ 32 lên 43: dòng chênh lệch nhảy từ âm 8.192 sang dương 256. Nghĩa là "MLA tốn thêm tham số" không phải một định luật, nó là một phát biểu về cấu hình cụ thể.

Phép chiếu thật sự làm gì, và ca vàng

Phần 3 của sim không nói về bộ đệm nữa, nó nói về chuyện gì xảy ra với con số khi ta nén rồi dựng lại. Vector vào mặc định có 8 chiều là 8, 6, 5, 4, 3, 2, 2, 1, độ dài 12,609520. Nén xuống 3 chiều, tức chỉ lưu 3 con số thay vì 8, rồi dựng lại về 8 chiều: sai số dựng lại là 1,145214405, tức 9,082% so với độ dài vector gốc. Bảng bên dưới cho cả đường cong: giữ 1 chiều thì sai số 6,234981957 hay 49,447%, giữ 2 chiều thì còn 1,598510078, giữ 3 chiều thì 1,145214405, và giữ đủ 8 chiều thì 0,000000000.

Hàng cuối cùng đó là ca vàng của bài. Khi chiều tiềm ẩn bằng chiều đầy đủ và ma trận chiếu khả nghịch thì phép chiếu lên là nghịch đảo đúng của phép chiếu xuống, nên vector dựng lại trùng khít vector gốc và sai số bằng 0. Nó không chứng minh rằng nén không mất mát, vì ở đó chẳng có nén nào cả. Nó chứng minh một điều khác, quan trọng hơn với người đọc: phép chiếu trong sim này được cài đúng. Cổng kiểm số của bài khẳng định điều đó với sai số 1e-12, trên cả hai loại ma trận, cho mọi chiều từ 2 tới 12.

Hai điều nữa nhìn thấy được bằng số ở bảng đó.

Sai số không bao giờ giảm khi ta cắt bớt chiều tiềm ẩn. Kéo thanh r xuống thì cột sai số chỉ đi lên hoặc đứng yên, chưa bao giờ đi xuống. Điều này không phải may mắn của một vector cụ thể: hai họ ma trận trong sim đều lồng nhau, ma trận r hàng đúng là ma trận r + 1 hàng bỏ đi hàng cuối, nên không gian giữ lại chỉ có thể nhỏ đi khi r nhỏ đi.

Cái quyết định là không gian giữ lại, không phải mức nén. Đổi ô Ma trận chiếu sang loại hàng không trực giao mà giữ nguyên mọi thứ khác: sai số ở r bằng 3 nhảy từ 1,145214405 lên 4,628632878. Cùng một vector, cùng một mức nén, sai số khác gấp bốn lần, chỉ vì hai ma trận giữ lại hai không gian con khác nhau. Đây chính là chỗ để hiểu vì sao mô hình thật phải học ma trận chiếu chứ không lấy một ma trận cho sẵn.

Sim này không nói được gì

Phải nói thẳng chỗ này, vì nó là ranh giới giữa cái bài đo được và cái bài không đo được.

Ma trận chiếu ở phần 3 do tôi đặt cố định. Nó là một công thức đóng, tất định, không có số ngẫu nhiên nào, và đó là lý do cổng kiểm số kiểm được nó. Nhưng mô hình thật học cả hai ma trận xuống và lên từ dữ liệu. Ma trận học được biết những hướng nào thật sự xuất hiện trong K và V của mô hình đó, nên nó dồn ngân sách r chiều vào đúng những hướng ấy. Sai số dựng lại của mô hình thật vì thế nhỏ hơn nhiều con số bạn thấy ở đây, và không có cách nào suy nó ra từ sim này. Con số 9,082% ở trên là sai số của phép chiếu cosin trên một vector tôi chọn, nó không phải sai số của DeepSeek-V2 và cũng không xấp xỉ sai số đó.

Phép chiếu lên trong sim là giả nghịch đảo của phép chiếu xuống, tức cách dựng lại tốt nhất nếu ta coi mọi hướng trong không gian quan trọng như nhau. Mô hình thật học hai ma trận riêng biệt, và chúng không phải giả nghịch đảo của nhau.

Nói gọn: phần kế toán bộ nhớ và tham số ở mục 1 và mục 2 là chính xác, đúng theo cấu hình bạn đặt, và mọi con số ở đó bị cổng kiểm số canh. Phần chất lượng ở mục 3 chỉ minh hoạ cơ chế, nó cho bạn thấy nén thì mất mát và r bằng d thì không mất mát, chứ nó không định lượng được chất lượng mô hình. Muốn biết MLA làm điểm số của một mô hình thay đổi bao nhiêu thì phải huấn luyện, và đó là việc không tab trình duyệt nào làm được.

Vài chỗ dễ vấp

  • Số lớp và kiểu số không đổi hệ số nén. Chúng nhân vào cả hai cột nên triệt tiêu trong tỉ lệ. Chúng chỉ đổi con số byte tuyệt đối, mà con số tuyệt đối mới là thứ quyết định bạn có nạp nổi mô hình hay không.
  • Byte bộ đệm là hàm bậc nhất của d_c, không phải tỉ lệ thuận. Vì có số hạng d_rope cộng thêm, nhân đôi d_c không nhân đôi bộ đệm. Đặt d_rope về 0 thì mới tỉ lệ thuận đúng.
  • Kéo d_c về 0 là vô nghĩa và sim chặn. Không còn gì để dựng lại K và V thì không còn mô hình. Kéo d_c lên trên 2 × n_h × d_h cũng vô nghĩa theo hướng khác: bộ đệm "nén" to hơn bộ đệm đầy đủ, và sim nói thẳng ra như vậy chứ không im lặng cho ra một con số nhỏ hơn 1.
  • Bảng tham số dùng một đơn giản hoá. Bản gốc chia khoá thành phần nén và phần vị trí với hai bề rộng khác nhau, và có thể gộp W_UK vào ma trận truy vấn để không bao giờ dựng K ra. Ở đây tôi dùng một bề rộng d_h cho cả K và V và không gộp. Kế toán bộ đệm không đổi vì thế, nhưng số tham số và số phép tính thì lệch so với bản cài thật.
  • Preset mang tên mô hình là ảnh chụp một thời điểm. Số kiến trúc của hai preset DeepSeek ghi lại ngày 27/07/2026, đọc từ báo cáo kỹ thuật và tệp cấu hình phát hành kèm mô hình. Lĩnh vực này đổi rất nhanh nên hãy tự đối chiếu lại. Cấu hình thứ ba cố tình không mang tên mô hình nào: nó là số tròn để tính tay.
Điều rút ra

MLA đổi một bộ đệm to thành một bộ đệm nhỏ cộng thêm hai phép nhân ma trận. Hệ số nén tính được chính xác, xấp xỉ 2 × n_h × d_h / d_c và chính xác là 2 × n_h × d_h / (d_c + d_rope), nên kéo chiều tiềm ẩn xuống là bộ đệm co lại rất nhanh. Chiều ngược lại thì không tính được từ đây: sai số dựng lại chắc chắn tăng khi r giảm, nhưng tăng bao nhiêu thì phụ thuộc vào ma trận chiếu mà mô hình học được, và đó là thứ chỉ huấn luyện mới biết.

Câu hỏi tự kiểm0/3 đúngchưa trả lời
  1. 1Cấu hình mặc định có 128 đầu, mỗi đầu 128 chiều, chiều tiềm ẩn 512, phần RoPE tách riêng 64. Hệ số nén bộ đệm bằng bao nhiêu?
  2. 2Ở phần 3 của sim, đặt chiều tiềm ẩn r bằng chiều đầy đủ d thì sai số dựng lại về 0. Điều này chứng minh gì?
  3. 3Ở cấu hình mặc định, bốn ma trận của MLA cộng lại là 1.183.580.160 tham số, còn W_K và W_V của MHA đầy đủ là 10.066.329.600. Kết luận nào đúng?

Đi tiếp

Bài sau trong chương này đi theo hướng khác hẳn: thay vì nén cái phải nhớ, ta bỏ bớt cái phải nhớ, tức là chỉ cho mỗi token nhìn về sau một cửa sổ hữu hạn. Còn nếu bạn muốn xem lại vì sao một ma trận r hàng d cột nén được vector, và vì sao phép nén đó mất mát khi r nhỏ hơn d, thì quay lại ma trận biến hình ở môn Toán ML.