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

MHA, MQA và GQA

Sửa cả sáu tham sốByte đúng tới từng byteKhông đo được chất lượng

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ì KV 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ố.

Chia sẻ đầu K và V · từ MHA qua MQA tới GQA
chế độ GQAtiết kiệm byte mỗi token 128 KiB
Llama 3 8Bsố theo 04/2024
Nguồn: thẻ mô hình và tệp config.json bản 8B, kèm báo cáo Llama 3
32 đầu Q chia thành 8 nhóm, mỗi nhóm 4 đầu. Đây là cấu hình GQA hay được lấy làm mẫu nhất, và là trạng thái mặc định của sim.
Mỗi preset ghi số theo bản công bố ở thời điểm đứng trong ngoặc, không phải số được đọc lại hôm nay. Lĩnh vực này đổi rất nhanh nên hãy coi chúng là ví dụ đã đóng băng thời điểm, và đối chiếu lại tệp config.json của mô hình khi cần. Mọi tham số trong sim đều sửa được, kể cả sau khi bấm preset.
Cấu hình ✎ sửa được
Đầu Q nào dùng chung cặp khoá và giá trị nào
GQA · 8 nhóm × 4 đầu Qđầu Qcặp K, Vnhóm số1234567891011121314151617181920212223242526272829303132KV1KV2KV3KV4KV5KV6KV7KV8
32 đầu Q chia thành 8 nhóm, mỗi nhóm 4 đầu Q dùng chung đúng một cặp KV. 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.
hệ số tiết kiệm so với MHA
32 / 8 = 4 đúng bằng số đầu Q chia số nhóm KV
byte KV cache mỗi token
128 KiB
2 × 8 × 128 × 2 × 32 tức K và V, số nhóm, chiều mỗi đầu, byte mỗi số, số lớp
bộ đệm cho 8 192 token
1 GiB
một chuỗi, chưa tính chi phí phân trang. MHA cùng cấu hình cần 4 GiB, MQA cần 128 MiB
Ba chế độ trên đúng cấu hình này
Chế độNhóm KVĐầu Q mỗi nhómByte KV mỗi tokenBộ đệm cho 8 192 tokenTiết kiệm so với MHATham số chiếu K và VTham số cả 4 phép chiếu
MHA
MHA · mỗi đầu Q một cặp KV riêng
321512 KiB
524 288 byte
4 GiB
4 294 967 296 byte
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
84128 KiB
131 072 byte
1 GiB
1 073 741 824 byte
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
13216 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óm1.07 tỉ tham số, giống nhau ở cả ba dòng
Sim này không đo được chất lượng. Mọi con số ở trên là kế toán: byte và tham số, đúng tới từng byte với cấu hình bạn đặt. Nhưng chia sẻ đầu KV làm mô hình mất bao nhiêu độ chính xác thì chỉ huấn luyện mới biết, và ở đây không có huấn luyện nào cả. Nếu chỉ nhìn bảng này thì MQA luôn thắng, mà thực tế thì không phải vậy. Ngoài ra bảng bỏ qua hệ số chệch của các phép chiếu, và tính bộ đệm cho một chuỗi duy nhất, chưa cộng chi phí phân trang khi phục vụ nhiều yêu cầu.

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 , 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ómmỗ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_KW_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 MHA hoặc trùng dòng MQA khi 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ống nghe như bỏ mất nhiều, nhưng đã cắt 4 GiB xuống 1 GiB, còn đi tiếp tới 32× chỉ cắt thêm từ 1 GiB xuống 128 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 int4 thì 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ều rút ra

Đổ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 KV. 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.

Câu hỏi tự kiểm0/3 đúngchưa trả lời
  1. 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. 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. 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.