MHA, MQA và GQA
MHA, MQA và GQA
Ba cái tên nghe như ba kiến trúc khác nhau. Thật ra chúng là một cơ chế với đúng một con số bị đổi: có bao nhiêu nhóm đầu khoá và giá trị cho các đầu truy vấn dùng chung.
Ở self-attention bạn đã thấy mỗi đầu chú ý có ba phép chiếu riêng: một cho truy vấn Q, một cho khoá K, một cho giá trị V. Bài KV cache rồi cho thấy lúc sinh chữ thì K và V của mọi token trước phải được giữ lại trong bộ nhớ, còn Q thì không. Ghép hai điều đó lại là ra ngay một câu hỏi tiền bạc: nếu mỗi đầu truy vấn phải kéo theo một cặp khoá và giá trị riêng, thì bộ đệm phình theo số đầu.
Câu trả lời của mọi mô hình mở gần đây là cho các đầu truy vấn dùng chung đầu khoá và giá trị. Chia sẻ tới mức nào thì có một thang: mỗi đầu một cặp riêng là MHA, tất cả chung đúng một cặp là MQA, còn chia nhóm là GQA. Ba cái tên, một con số.
K và V. Bề rộng mỗi ô KV bằng bề rộng nhóm đầu Q đang dùng nó, nên nhìn hình là đọc được cách chia.32 / 8 = 4 đúng bằng số đầu Q chia số nhóm KV2 × 8 × 128 × 2 × 32 tức K và V, số nhóm, chiều mỗi đầu, byte mỗi số, số lớp| Chế độ | Nhóm KV | Đầu Q mỗi nhóm | Byte KV mỗi token | Bộ đệm cho 8 192 token | Tiết kiệm so với MHA | Tham số chiếu K và V | Tham số cả 4 phép chiếu |
|---|---|---|---|---|---|---|---|
| MHA MHA · mỗi đầu Q một cặp KV riêng | 32 | 1 | 512 KiB 524 288 byte | 4 GiB 4 294 967 296 byte | 1× | 1.07 tỉ 1 073 741 824 | 2.15 tỉ 2 147 483 648 |
| GQAđang đặt GQA · các đầu Q chia nhóm, mỗi nhóm một cặp KV | 8 | 4 | 128 KiB 131 072 byte | 1 GiB 1 073 741 824 byte | 4× | 268.4 triệu 268 435 456 | 1.34 tỉ 1 342 177 280 |
| MQA MQA · mọi đầu Q dùng chung đúng một cặp KV | 1 | 32 | 16 KiB 16 384 byte | 128 MiB 134 217 728 byte | 32× | 33.6 triệu 33 554 432 | 1.11 tỉ 1 107 296 256 |
Chiếu Q và chiếu ra O: không phụ thuộc cách chia nhóm | 1.07 tỉ tham số, giống nhau ở cả ba dòng | ||||||
Hãy làm đúng ba việc trong sim, theo thứ tự.
Một: kéo thanh trượt Số nhóm KV từ 8 lên 32. Hình vẽ nối lại ngay: mỗi ô K, V thu về đúng bề rộng một đầu Q, không còn đầu nào dùng chung với đầu nào, và nhãn chế độ đổi thành MHA. Bảng đọc như sau. Byte KV mỗi token nhảy từ 128 KiB lên 512 KiB, bộ đệm cho 8 192 token nhảy từ 1 GiB lên 4 GiB, và hệ số tiết kiệm rơi xuống 1×, tức là không tiết kiệm gì so với chính nó. Kéo tiếp về 1 thì đúng chiều ngược lại: mọi đầu Q đổ về một ô K, V duy nhất, byte mỗi token còn 16 KiB, bộ đệm còn 128 MiB, hệ số tiết kiệm lên 32×, và nhãn đổi thành MQA.
Ba con số đó cùng nằm trên một đường thẳng, và đó là điều đáng nhớ nhất của bài này:
Hệ số tiết kiệm bộ đệm đúng bằng số đầu Q chia số nhóm KV.
Ở mặc định là 32 / 8 = 4, nên 512 KiB chia 4 ra 128 KiB. Đặt 32 nhóm thì 32 / 32 = 1. Đặt 1 nhóm thì 32 / 1 = 32. Không có phép nhân bí ẩn nào ở đây: bộ đệm giữ hai vector, khoá và giá trị, cho mỗi nhóm ở mỗi lớp, nên số byte tỉ lệ thuận với số nhóm, và tỉ lệ thuận thì chia đôi số nhóm là chia đôi số byte.
Hai: đọc cột Tham số chiếu K và V. Chỗ này nhiều người bỏ qua. Chia sẻ đầu KV không chỉ làm nhỏ bộ đệm lúc chạy, nó còn làm nhỏ chính ma trận trọng số, vì W_K và W_V giờ chỉ cần sinh ra số nhóm × chiều mỗi đầu số thay vì số đầu Q × chiều mỗi đầu số. Ở mặc định, hai phép chiếu đó đi từ 1.07 tỉ tham số của MHA xuống 268.4 triệu, cũng đúng 4 lần, cùng một hệ số.
Nhưng cột kế bên bắt bạn đứng lại. Phép chiếu Q và phép chiếu ra O không đổi gì cả, vì chúng vẫn phải làm việc trên đủ 32 đầu. Nên tổng tham số của bốn phép chiếu chỉ đi từ 2.15 tỉ xuống 1.34 tỉ, tức khoảng 1.6 lần, chứ không phải 4 lần. Và bốn phép chiếu này còn chưa phải cả mô hình: phần lớn tham số của một mô hình 8 tỉ nằm ở các lớp truyền thẳng và ở bảng nhúng, những chỗ mà cách chia nhóm KV không chạm tới. Bộ đệm nhỏ đi 4 lần, tham số chú ý nhỏ đi 1,6 lần, tổng mô hình gần như không đổi. Ba con số khác nhau cho cùng một thay đổi, và nói lẫn chúng là chỗ hay sai nhất khi đọc báo cáo kỹ thuật.
Ba: kéo Số nhóm KV sang 3. Cả bảng biến mất và thay bằng một khung báo lỗi: 32 không chia hết cho 3, còn dư 2. Sim không làm tròn hộ bạn, vì một nhóm 10 đầu rưỡi thì không tồn tại. Mỗi nhóm phải có đúng bằng nhau số đầu Q, nên số nhóm buộc phải là một ước của số đầu Q. Với 32 đầu Q thì các số nhóm hợp lệ là 1, 2, 4, 8, 16, 32, và khung lỗi bày sẵn đúng dãy đó thành các nút bấm được. Đây là lý do các mô hình thật hay chọn số đầu và số nhóm đều là lũy thừa của 2: mọi ước đều có sẵn.
Vài chỗ dễ hiểu lệch, nói rõ luôn:
- MHA và MQA không phải hai cơ chế đứng cạnh GQA. Chúng là hai đầu của cùng một thang. Đặt số nhóm bằng số đầu Q thì công thức ra đúng con số của MHA, đặt bằng 1 thì ra đúng con số của MQA. Sim thể hiện chuyện đó bằng cách in ra dòng
trùng dòng MHAhoặctrùng dòng MQAkhi bạn chạm vào hai đầu thang. - GQA thắng vì nó ở giữa, không vì nó tiết kiệm nhất. Nhìn riêng bảng byte thì MQA luôn thắng, mà thực tế thì các mô hình lớn không chọn MQA. Lý do nằm ngoài bảng này: nén 32 cặp khoá và giá trị về 1 cặp làm mô hình mất chất lượng rõ rệt, còn nén về 8 nhóm thì mất rất ít mà đã lấy được phần lớn khoản tiết kiệm. Từ
32×xuống4×nghe như bỏ mất nhiều, nhưng4×đã cắt4 GiBxuống1 GiB, còn đi tiếp tới32×chỉ cắt thêm từ1 GiBxuống128 MiB. Phần dễ ăn nằm ở những bước đầu. - Ô ngữ cảnh nhân vào, không đổi hệ số. Đổi độ dài ngữ cảnh thì cột bộ đệm đổi theo, nhưng hệ số tiết kiệm đứng yên, vì nó là tỉ số và ngữ cảnh có mặt ở cả tử số lẫn mẫu số. Kiểu số cũng vậy: chuyển bf16 sang fp8 chia đôi mọi con số byte mà hệ số vẫn nguyên.
- Kiểu số dưới một byte cho ra byte thập phân. Chọn
int4thì mỗi số chỉ chiếm nửa byte, nên con số byte mỗi token có thể không còn nguyên. Đó là kế toán đúng, nhưng lượng tử hoá bộ đệm là một quyết định khác hẳn với chia sẻ đầu KV, và bài này không nói gì về việc nó làm mô hình sai thêm bao nhiêu.
Sim này không nói được gì
Đây là chỗ phải nói thẳng, và nó quan trọng hơn mọi con số ở trên.
Sim không đo được chất lượng. Nó tính byte và tham số, chính xác tới từng byte trên cấu hình bạn đặt. Nhưng câu hỏi thật sự khi chọn số nhóm KV là "chia sẻ tới mức này thì mô hình kém đi bao nhiêu", và câu đó chỉ huấn luyện mới trả lời được. Ở đây không có huấn luyện, không có dữ liệu, không có perplexity. Nếu bạn chỉ đọc bảng này thì kết luận sẽ là MQA luôn tốt nhất, và kết luận đó sai.
Nói cách khác: sim cho bạn một nửa của phép đánh đổi, đúng tuyệt đối ở nửa đó, và im lặng hoàn toàn ở nửa kia. Cổng kiểm số của bài canh phần kế toán, không canh phần chất lượng, và cũng không canh chuyện mô hình nào ngoài kia thật sự có cấu hình nào.
Vì lẽ đó, các preset mang tên mô hình đều ghi kèm thời điểm công bố và nguồn ngay trên giao diện. Chúng là ví dụ đã đóng băng thời điểm, không phải sự thật vĩnh viễn, nên nếu bạn có số mới hơn thì gõ thẳng vào sáu ô cấu hình là kiểm lại được ngay. Hai preset ghi rõ chỉ là minh hoạ thì không phải cấu hình của mô hình nào cả.
Bảng số còn bỏ qua hai thứ, và bỏ qua có chủ ý: hệ số chệch của các phép chiếu, vì chúng nhỏ tới mức không đáng kể trước hàng triệu tham số trọng số, và chi phí phân trang bộ đệm khi phục vụ nhiều yêu cầu cùng lúc, vì cái đó phụ thuộc bộ máy phục vụ chứ không phụ thuộc kiến trúc.
Đổi một con số duy nhất, số nhóm KV, là đi hết từ MHA sang MQA. Hệ số tiết kiệm bộ đệm bằng đúng số đầu Q chia số nhóm KV, và cùng hệ số đó cũng thu nhỏ hai ma trận chiếu K và V. Nhưng nó không chạm vào chiếu Q với chiếu ra, nên phần tham số tiết kiệm được nhỏ hơn phần bộ đệm tiết kiệm được. GQA được chọn không phải vì nó tiết kiệm nhất mà vì nó nằm ở chỗ đã lấy gần hết khoản tiết kiệm trong khi còn mất ít chất lượng, và phần "mất ít chất lượng" ấy là phần duy nhất mà cái sim này không đo được.
- 1Một mô hình có 64 đầu Q và 8 nhóm KV. So với MHA cùng hình dáng, bộ đệm KV mỗi token nhỏ đi bao nhiêu lần?
- 2Đặt số nhóm KV bằng 6 cho một mô hình 32 đầu Q. Sim phải làm gì?
- 3Ở cấu hình mặc định, bộ đệm KV nhỏ đi 4 lần khi chuyển từ MHA sang GQA 8 nhóm. Tổng tham số của bốn phép chiếu chú ý thì thế nào?
Bài kế tiếp trong chương đi xa hơn một bước: thay vì chia nhóm các đầu khoá và giá trị, MLA nén bộ đệm bằng chiếu tiềm ẩn rồi bung ra khi cần. Còn nếu bạn muốn xem lại vì sao bộ đệm tồn tại ngay từ đầu thì quay lại KV cache.