Tinh chỉnh hạng thấp (LoRA)
Tinh chỉnh hạng thấp
Huấn luyện lại 8 tỉ tham số thì bộ tối ưu một mình đã ngốn gần 90 GiB. Nếu bản cập nhật chỉ cần nằm trong một không gian hẹp thì cái giá đó tụt xuống bao nhiêu, và tới hạng nào thì mẹo này hết rẻ?
Tinh chỉnh đầy đủ một mô hình nghĩa là mỗi trọng số đều được cập nhật, nên bộ tối ưu phải giữ trạng thái cho từng tham số một. Với Adam và một bản trọng số chính ở fp32 thì đó là 12 byte cho mỗi tham số, và trên mô hình 8.030.261.248 tham số con số ấy thành 89,75 GiB, chưa tính trọng số gốc và chưa tính kích hoạt trung gian. Đây là lý do rất nhiều người không thể tinh chỉnh mô hình mở dù đã tải được nó về.
Ý tưởng của LoRA là đóng băng toàn bộ trọng số gốc rồi giả định thêm một điều: bản cập nhật cần thiết không cần phủ hết không gian. Với một ma trận trọng số d_out × d_in, thay vì học một bản cập nhật đầy đủ ΔW có d_in · d_out số, ta học hai ma trận gầy, B cỡ d_out × r và A cỡ r × d_in, rồi dùng ΔW = B · A. Số tham số huấn luyện được đi từ d_in · d_out xuống r · (d_in + d_out). Đó là cả cơ chế, phần còn lại của bài là kế toán, và kế toán thì đếm được chính xác.
1 · Lắp bộ chuyển vào ma trận nào ✎ sửa được
| ma trận | hình dáng d_out × d_in | cập nhật đầy đủ, mỗi lớp | bộ chuyển hạng 16, mỗi lớp | tỉ lệ | r* hoà vốn | hạng rẻ hơn tối đa | ở hạng 16 |
|---|---|---|---|---|---|---|---|
W_qchiếu truy vấn | 4.096 × 4.096 | 16.777.216 | 131.072 | 0.78% | 2048 | 2.047 | còn rẻ hơn |
W_kchiếu khoá | 1.024 × 4.096 | 4.194.304 | 81.920 | 1.95% | 819.2 | 819 | còn rẻ hơn |
W_vchiếu giá trị | 1.024 × 4.096 | 4.194.304 | 81.920 | 1.95% | 819.2 | 819 | còn rẻ hơn |
W_ochiếu đầu ra | 4.096 × 4.096 | 16.777.216 | 131.072 | 0.78% | 2048 | 2.047 | còn rẻ hơn |
W_gatecổng của khối tiến | 14.336 × 4.096 | 58.720.256 | 294.912 | 0.50% | 3185.78 | 3.185 | còn rẻ hơn |
W_upnhánh lên của khối tiến | 14.336 × 4.096 | 58.720.256 | 294.912 | 0.50% | 3185.78 | 3.185 | còn rẻ hơn |
W_downnhánh xuống của khối tiến | 4.096 × 14.336 | 58.720.256 | 294.912 | 0.50% | 3185.78 | 3.185 | còn rẻ hơn |
2 · Hạng r, hệ số alpha, và điểm hoà vốn ✎ sửa được
Bản cập nhật là ΔW = (alpha / r) · B · A, nên tỉ lệ nhân đang là 32 chia 16 bằng 2. Hệ số này không đổi một tham số nào: kéo alpha mà xem, hai ô số ở thẻ trên đứng im. Nó chỉ đổi độ lớn của bản cập nhật được cộng vào trọng số gốc. Hạng r thì đổi cả hai: nó đổi số tham số và đổi tỉ lệ nhân, vì r nằm ở mẫu.
3 · Chiều của mô hình ✎ sửa được
Chiều truy vấn là 32 × 128 = 4.096, chiều khoá giá trị là 8 × 128 = 1.024. Một lớp có 41.943.040 tham số ở khối chú ý, 176.160.768 ở khối tiến và 8.192 ở hai lớp chuẩn hoá, tổng 218.112.000. Nhân 32 lớp, cộng 1.050.673.152 tham số nhúng và đầu ra cùng 4.096 tham số chuẩn hoá cuối, ra 8.030.261.248 tham số.
4 · Bộ nhớ khi huấn luyện ✎ sửa được
| khoản | LoRA hạng 16 | tinh chỉnh đầy đủ |
|---|---|---|
| tham số huấn luyện được | 6.815.744 | 8.030.261.248 |
| trọng số gốc đóng băng | 14.96 GiB | 14.96 GiB |
| trọng số bộ chuyển (2 byte mỗi tham số) | 13 MiB | 0 B |
| trạng thái bộ tối ưu (12 byte mỗi tham số huấn luyện được) | 78 MiB | 89.75 GiB |
| tổng ba khoản trên | 15.05 GiB | 104.7 GiB |
- Nó đếm tham số và byte. Nó không đo chất lượng. Hạng 4 và hạng 64 hiện ra ở đây gọn gàng như nhau, còn chuyện hạng nào đủ cho việc của bạn thì chỉ chạy thực nghiệm mới biết, và câu trả lời đổi theo từng tác vụ.
- Bộ nhớ ở đây là trọng số cộng trạng thái bộ tối ưu. Nó bỏ qua kích hoạt trung gian, bỏ qua chuyện tính lại kích hoạt để đổi tính toán lấy bộ nhớ, bỏ qua phân mảnh và bộ đệm của thư viện. Số thật khi chạy luôn lớn hơn.
- Số byte của bộ tối ưu là một quy ước bạn chọn trong menu, không phải chân lý. Adam giữ hai moment cộng một bản trọng số chính là 12 byte mỗi tham số, bộ tối ưu 8 bit thì khác, SGD thì khác nữa.
- Phép đếm bỏ qua bias và giả định khối tiến kiểu có cổng với ba ma trận. Mô hình có bias, hoặc khối tiến hai ma trận, sẽ lệch một chút, và preset nào có bias thì dòng nguồn nói rõ lệch bao nhiêu.
- Điểm hoà vốn là chuyện đếm tham số, không phải chuyện biểu diễn. Ở đúng mốc r* thì hai bên tốn tham số bằng nhau, nhưng bên LoRA vẫn bị ràng buộc hạng còn bên kia thì không, nên hai thứ bằng nhau về giá chứ không bằng nhau về khả năng.
Ba con số ở trạng thái mở bài
Phần huấn luyện được nhỏ tới mức khó tin, và nó là phép cộng chứ không phải phép nhân. Sim mở ra ở chiều của một mô hình 8 tỉ tham số với hạng r = 16, lắp bộ chuyển vào hai ma trận: chiếu truy vấn và chiếu giá trị. Chiếu truy vấn là 4.096 × 4.096, nên bộ chuyển tốn 16 × (4.096 + 4.096) = 131.072 tham số mỗi lớp. Chiếu giá trị là 1.024 × 4.096, nên nó tốn 16 × (1.024 + 4.096) = 81.920. Cộng lại 212.992 tham số mỗi lớp, nhân 32 lớp ra 6.815.744 tham số huấn luyện được, đúng 0,0849% mô hình. Cập nhật đầy đủ đúng hai ma trận đó thôi đã là 671.088.640 tham số, nên bộ chuyển chỉ bằng 1,0156% của nó.
Chiều của bảy ma trận không giống nhau, và đây là chỗ hay bị đếm sai. Bảng trong sim để cả bảy ma trận cạnh nhau chính vì vậy. Với chú ý theo nhóm, chiếu khoá và chiếu giá trị chỉ cao 8 × 128 = 1.024 thay vì 32 × 128 = 4.096, tức bằng một phần tư chiếu truy vấn về số tham số. Ba ma trận của khối tiến thì đi hướng ngược lại: chúng dùng chiều trung gian 14.336, nên W_gate và W_up là 14.336 × 4.096 còn W_down là 4.096 × 14.336, mỗi cái 58.720.256 tham số, tức gấp ba lần rưỡi chiếu truy vấn. Bấm chip để lắp bộ chuyển vào chúng là thấy tổng nhảy hẳn một bậc. Nếu bạn muốn hiểu vì sao khoá và giá trị lại hẹp đi như vậy thì bài MHA, MQA và GQA đếm riêng chuyện đó.
Chỗ tiết kiệm lớn nhất là trạng thái bộ tối ưu, không phải trọng số. Bảng bộ nhớ nói thẳng: trạng thái bộ tối ưu đi từ 89,75 GiB xuống 78 MiB, tiết kiệm 89,67 GiB. Trọng số gốc thì không nhỏ đi một byte nào, vẫn 14,96 GiB ở bf16, vì mô hình vẫn phải nằm trong bộ nhớ để chạy xuôi và chạy ngược. Cộng ba khoản lại, một vòng huấn luyện đầy đủ cần 104,7 GiB còn LoRA cần 15,05 GiB. Và phần thắng ấy không đến từ việc chọn ít ma trận: nếu tinh chỉnh đầy đủ nhưng chỉ đúng hai ma trận đã chọn thì trạng thái bộ tối ưu vẫn là 7,5 GiB, gấp gần trăm lần. Nó đến từ hạng thấp.
Alpha không phải là hạng, và nó không đổi số tham số
Đây là chỗ bị hiểu lẫn nhiều nhất, nên nói cho rõ. Bản cập nhật thật là ΔW = (alpha / r) · B · A, tức có một hệ số nhân đứng trước. Ở trạng thái mở bài, alpha = 32 và r = 16, nên tỉ lệ nhân bằng 2.
Kéo núm alpha lên 64 mà xem: tỉ lệ nhân đi lên 4, còn hai thẻ số tham số đứng im. Alpha không tạo thêm một trọng số nào, nó chỉ nhân bản cập nhật lên. Kéo núm hạng r lên 32 thì cả hai đổi: số tham số gấp đôi thành 13.631.488, và tỉ lệ nhân tụt từ 2 xuống 1, vì r nằm dưới mẫu. Chia cho r là một quy ước, và lý do người ta nêu ra khi chọn nó là để đổi hạng thì không phải dò lại tốc độ học từ đầu. Đó là lời giải thích về ý định, không phải thứ sim này kiểm được.
Hai điều sim này không chứng minh được, và bạn nên nghi ngờ ai nói ngược lại. Thứ nhất, alpha nên đặt bằng bao nhiêu là câu hỏi thực nghiệm; bài LoRA gốc dùng cách khác với bài QLoRA, và preset thứ hai trong sim để r = 64 với alpha = 16, tức tỉ lệ nhân chỉ 0,25. Thứ hai, ở lúc khởi tạo thì B bằng 0 nên ΔW bằng 0 bất kể alpha là bao nhiêu; hệ số này chỉ có ý nghĩa sau khi B đã học được gì đó. Sim đếm tham số và byte, nó không mô phỏng quá trình học.
Điểm hoà vốn, và hai chi tiết mà cách nói gọn làm mất
Bộ chuyển rẻ hơn không phải là chuyện đương nhiên, nó chỉ đúng khi r còn nhỏ. Đặt hai bên bằng nhau: r · (d_in + d_out) = d_in · d_out, suy ra
r* = (d_in · d_out) / (d_in + d_out)
Với ma trận vuông d × d thì r* = d² / 2d = d / 2, đúng như người ta thường nói. Chiếu truy vấn 4.096 × 4.096 có r* = 2.048, và sim ghi đúng con số đó.
Nhưng cách nói gọn "hoà vốn ở d / 2" làm mất hai chi tiết, và cả hai đều đo được:
- Ở đúng mốc
r*thì hai bên bằng khít nhau, chứ không phải LoRA vẫn còn rẻ hơn. Hạng 2.047 cho chiếu truy vấn tốn 16.769.024 tham số, dưới 16.777.216 của bản cập nhật đầy đủ. Hạng 2.048 tốn đúng 16.777.216, tức bằng nhau. Hạng 2.049 tốn 16.785.408, đã đắt hơn. Vậy hạng cuối cùng còn rẻ hơn thật là 2.047, không phải 2.048. - Với ma trận chữ nhật thì
r*không phải một nửa cạnh nào cả. Chiếu giá trị1.024 × 4.096cór* = 4.194.304 / 5.120 = 819,2. Một nửa cạnh ngắn là 512, mà ở hạng 512 bộ chuyển vẫn còn rẻ hơn. Một nửa trung bình hai cạnh là 1.280, mà ở hạng 1.280 nó đã đắt hơn rồi. Cả hai cách nhớ tắt đều sai, vàr*thật là một nửa trung bình điều hoà của hai chiều.
Ba ma trận khối tiến hoà vốn ở 3.185,78, còn cả cụm hai ma trận đang chọn thì hoà vốn ở 1.575,38 nên hạng 1.575 vẫn rẻ, hạng 1.576 thì hết. Thang hạng trong sim luôn chèn sẵn hai hàng ở hai bên mốc để bạn khỏi phải đi tìm. Muốn tự tay kéo qua mốc thì bấm preset khối nhỏ: ở đó chiếu truy vấn là 512 × 512 nên hoà vốn đúng hạng 256, và cụm hai ma trận trở nên đắt hơn ngay từ hạng 197.
Còn một mốc nữa, khác loại: tích B · A không thể có hạng cao hơn min(d_out, d_in). Đặt r vượt mốc đó thì bạn trả thêm tham số mà không mua thêm được hướng biểu diễn nào. Với chiếu giá trị 1.024 × 4.096 thì mốc ấy là 1.024, còn r* là 819,2, nên mốc chi phí tới trước.
Và nó luôn tới trước, không phụ thuộc hình dáng ma trận. Chứng minh một dòng: đặt a ≤ b là hai chiều, thì r* = a·b / (a + b) nhỏ hơn a·b / b = a = min(d_out, d_in), vì mẫu số a + b lớn hơn b khi a dương. Nói cách khác, LoRA hết rẻ hơn cập nhật đầy đủ trước khi nó chạm giới hạn biểu diễn, ở mọi hình dáng. Nên mốc biểu diễn là chuyện lý thuyết đáng biết, còn mốc phải nhìn trong thực tế là mốc chi phí.
QLoRA: đóng băng ở 4 bit
LoRA cắt trạng thái bộ tối ưu nhưng không cắt trọng số gốc, mà trọng số gốc ở đây là 14,96 GiB. QLoRA vá đúng chỗ đó: mô hình gốc đóng băng ở 4 bit, chỉ bộ chuyển giữ độ chính xác cao. Vì phần gốc không được huấn luyện, nén nó không làm mất gradient của ai.
Bật núm QLoRA trong sim: trọng số gốc tụt từ 14,96 GiB xuống 5,21 GiB, và cả vòng huấn luyện còn 5,3 GiB. Chú ý con số 5,21 GiB không phải một phần tư của 14,96 GiB, và sim nói rõ vì sao: nó chỉ hạ 6.979.321.856 tham số của các lớp tuyến tính trong khối xuống nửa byte, còn 1.050.939.392 tham số của bảng nhúng, đầu ra và các lớp chuẩn hoá vẫn để 2 byte, đúng như cách người ta cài thật. Với từ vựng 128.256, riêng phần đó đã là 1,96 GiB không nén được. Chuyện 4 bit mua được gì và mất gì thì hai bài một bit dùng vào việc gì và lượng tử hoá theo nhóm đo riêng. Ở đây cần nhớ một điều: sim không đếm các hệ số tỉ lệ theo nhóm mà lượng tử hoá 4 bit phải giữ kèm, nên con số 5,21 GiB là chặn dưới chứ không phải số cuối.
Vài điểm nhỏ nhưng hay bị bỏ qua:
- Lắp bộ chuyển vào nhiều ma trận hơn không đồng nghĩa tốn nhiều bộ nhớ hơn. Bấm preset thứ hai, kiểu QLoRA: cả bảy ma trận, hạng 64, phần huấn luyện được nhảy lên 167.772.160 tham số tức 2,0892% mô hình, gấp hơn hai mươi bốn lần preset đầu, mà tổng bộ nhớ lại thấp hơn, 7,4 GiB so với 15,05 GiB, vì phần gốc đã xuống 4 bit. Hai đại lượng đó không cùng chiều, và trộn chúng vào một câu là cách dễ nhất để nói sai.
- Số byte của bộ tối ưu là quy ước bạn chọn. Menu trong sim có 12 byte cho Adam kèm bản chính fp32, 6 byte cho bộ tối ưu 8 bit, 4 byte cho SGD có động lượng. Đổi menu là mọi con số bộ nhớ đổi theo, và không có lựa chọn nào là chân lý.
- Phần trăm phụ thuộc mẫu số. Cùng một bộ chuyển 6.815.744 tham số, trên mô hình 8 tỉ là 0,0849%, còn bấm sang preset Mistral 7B với từ vựng 32.000 thì thành 0,0941%, chỉ vì mô hình nhỏ hơn hai bảng nhúng. Con số phần trăm nghe ấn tượng nhưng nó nói về mẫu số nhiều hơn về bộ chuyển.
- Hạng thấp là một ràng buộc, không phải quà tặng. Ý này giống hệt chuyện nén dữ liệu về ít chiều ở bài phân tích thành phần chính: bạn giả định thứ mình cần nằm trong một không gian con hẹp. Nếu giả định đó sai với tác vụ của bạn thì không có hạng nào cứu được, và sim này không hề kiểm tra giả định đó.
LoRA đổi d_in · d_out tham số của một bản cập nhật đầy đủ thành r · (d_in + d_out), nên với hạng 16 trên hai ma trận chú ý thì phần huấn luyện được chỉ còn 6.815.744 tham số, tức 0,0849% mô hình, và trạng thái bộ tối ưu tụt từ 89,75 GiB về 78 MiB. Phép hoán đổi ấy chỉ có lãi khi r còn dưới r* = d_in · d_out / (d_in + d_out), mà r* bằng d / 2 chỉ với ma trận vuông; ở đúng mốc thì hai bên bằng khít nhau chứ không phải LoRA còn thắng. Alpha là một núm khác hẳn: nó nhân bản cập nhật lên alpha / r mà không tạo thêm một tham số nào. Và tất cả những con số này là chi phí, không phải chất lượng: hạng nào đủ dùng thì chỉ thực nghiệm mới trả lời được.
- 1Một chiếu khoá có hình dáng 1.024 × 4.096. Lắp bộ chuyển hạng 8 vào nó thì có bao nhiêu tham số huấn luyện được, tính cho một lớp?
- 2Đang ở hạng r = 16 với alpha = 32, phần huấn luyện được là 6.815.744 tham số và tỉ lệ nhân là 2. Bạn đổi alpha thành 64 và không đổi gì khác. Cái gì đổi?
- 3Vẫn ma trận 1.024 × 4.096. Hạng r bằng bao nhiêu thì bộ chuyển hết rẻ hơn một bản cập nhật đầy đủ của chính ma trận đó?