Bộ nhớ khi huấn luyện đi đâu hết
Bộ nhớ khi huấn luyện đi đâu hết
Trọng số của mô hình 8 tỉ tham số ở 2 byte là 14,9 GiB. Bạn có một thẻ 80 GiB, dư gấp năm lần. Vậy vì sao huấn luyện nó vẫn không nổi?
Ở bộ đệm KV bạn đã đếm bộ nhớ lúc suy luận: trọng số là hằng số, bộ đệm lớn lên theo từng token. Bài đó kết thúc bằng một câu để dành: huấn luyện là hoá đơn khác hẳn. Đây là bài trả nợ câu đó.
Chỗ khiến trực giác sai không phải là con số cuối cùng, mà là số lần cùng một tham số bị trả tiền. Khi suy luận, mỗi tham số nằm trong bộ nhớ đúng một lần. Khi huấn luyện, nó nằm ở đó dưới nhiều cái tên: bản trọng số mà lượt xuôi đọc, gradient của nó, bản fp32 để bước cập nhật làm việc cho chính xác, và một hoặc hai bộ đệm nữa mà bộ tối ưu giữ riêng. Mỗi cái tên là một khoản byte riêng, và mỗi khoản có bề rộng riêng.
Nên câu hỏi đúng không phải "một tham số tốn bao nhiêu byte khi huấn luyện". Câu hỏi đúng là phép cộng đó gồm những hạng nào. Bảng dưới đây để từng hạng thành một ô số bạn sửa được, rồi cộng lại trước mắt bạn.
1 · Hình dạng mô hình ✎ sửa được
2 · Tổng byte mỗi tham số, cộng theo từng vai ✎ sửa được
| vai | byte mỗi tham số | tổng byte | chia từ bậc | mỗi thiết bị |
|---|---|---|---|---|
| trọng số làm việc | 2 | 14,9 GiB | 3 | 14,9 GiB |
| gradient | 2 | 14,9 GiB | 2 | 14,9 GiB |
| bản trọng số chủ fp32 | 4 | 29,8 GiB | 1 | 29,8 GiB |
| moment thứ nhất | 4 | 29,8 GiB | 1 | 29,8 GiB |
| moment thứ hai | 4 | 29,8 GiB | 1 | 29,8 GiB |
| tổng trạng thái mô hình | 16 | 119,2 GiB | · | 119,2 GiB |
2 + 2 + 4 + 4 + 4 = 16 byte mỗi tham số
3 · Bộ nhớ kích hoạt ✎ sửa được
17 × 4.096 × 2 × 32 × 4.096 × 1 = 18.253.611.008 byte
34·s·b·h trong bài của Korthikanti và cộng sự. Bản gốc còn một hạng thức bậc hai theo độ dài chuỗi cho ma trận điểm chú ý, sim này không tính hạng thức đó.4 · Chia trạng thái cho nhiều thiết bị ✎ sửa được
5 · So với chỉ suy luận
| khoản | byte | so với thẻ 80 GiB |
|---|---|---|
| nạp trọng số để suy luận (2 byte mỗi tham số, chưa tính bộ đệm KV) | 14,9 GiB | vừa |
| trạng thái mô hình khi huấn luyện, mỗi thiết bị | 119,2 GiB | · |
| kích hoạt, mỗi thiết bị | 17,0 GiB | · |
| tổng một bước huấn luyện, mỗi thiết bị | 136,2 GiB | thiếu 56,2 GiB |
- Nó cộng byte. Nó không đo tốc độ, không đo chất lượng, không biết mô hình có học được gì không. Đổi Adam sang SGD trơn ở đây cắt ba phần tư hoá đơn trạng thái, và sim không hề nói rằng mô hình vẫn huấn luyện tốt như cũ. Cái đó phải chạy thật mới biết.
- Phần kích hoạt là xấp xỉ, tuyến tính theo lô, độ dài chuỗi, số lớp và chiều ẩn. Bản dẫn giải trong y văn còn một hạng thức bậc hai theo độ dài chuỗi cho ma trận điểm chú ý, mà sim này không tính, vì nhân chú ý hiện nay thường không dựng hẳn ma trận đó ra. Nên coi con số kích hoạt là sàn có một hạng thức đã biết là thiếu.
- Con số “+0,0% lượt tính” khi bật tái tính là đếm lượt nhân ma trận theo quy ước một lượt xuôi và hai lượt ngược, không phải đo đồng hồ. Thời gian thật phụ thuộc nhân, băng thông và cách xếp lớp.
- Bậc 3 chia cả tham số, nhưng lúc tính vẫn phải gom trọng số của từng lớp về rồi bỏ đi. Con số ở đây là mức thường trú, nên nó nhỏ hơn đỉnh thật. Sim không mô hình hoá phần gom đó.
- Preset mang tên mô hình chỉ chép lại ba con số công bố kèm nguồn và ngày. Cổng kiểm số canh phép tính, không canh chuyện mô hình ngoài kia có đúng hình dạng đó hay không.
Không có con số 16 thần thánh
Bạn sẽ gặp rất nhiều chỗ nói "huấn luyện tốn 16 byte cho mỗi tham số". Bạn cũng sẽ gặp chỗ nói 18, và chỗ nói 20. Cả ba đều có thật, và tranh nhau xem số nào mới đúng là việc vô ích, vì chúng chỉ khác nhau ở chỗ đặt fp32 vào đâu.
Ở trạng thái mở bài, bảng cộng như sau: 2 + 2 + 4 + 4 + 4 = 16. Trọng số làm việc 2 byte, gradient 2 byte, bản trọng số chủ fp32 4 byte, moment thứ nhất 4 byte, moment thứ hai 4 byte. Với 8 tỉ tham số thì thành 119,2 GiB trạng thái mô hình.
Bây giờ thử ba lần sửa, mỗi lần một ô:
- Đặt gradient thành 4 byte, tức giữ gradient ở fp32: tổng thành 18.
- Đặt cả trọng số làm việc thành 4 byte nữa: tổng thành 20.
- Bấm tắt độ chính xác trộn: bản trọng số chủ biến thành 0 byte và tổng tụt về 12.
Ba con số 16, 18, 20 mà bạn đọc thấy ngoài kia chính là ba dòng đầu. Chúng không phải ba phe, chúng là ba cách kế toán. Cái đáng nhớ là năm cái tên, không phải cái tổng.
Có một chỗ trong danh sách hay bị bỏ quên, và nó là chỗ đắt: bản trọng số chủ fp32 là một khoản riêng, không nằm chung với trọng số làm việc. Khi huấn luyện độ chính xác trộn, lượt xuôi và lượt ngược đọc bản 2 byte cho nhanh, còn bước cập nhật cộng dồn những lượng rất nhỏ nên phải làm trên bản 4 byte, nếu không thì cộng mãi mà không nhích. Hai bản cùng tồn tại. Ai cộng "trọng số cộng gradient cộng hai moment" rồi dừng lại sẽ ra 12 và thiếu mất 4 byte mỗi tham số.
Ba phần tư hoá đơn là thứ chỉ bộ tối ưu dùng
Nhìn dòng ngay dưới bảng vai. Trong 119,2 GiB trạng thái mô hình đó, 89,4 GiB tức 75,0% là bản trọng số chủ cộng hai moment, tức phần mà chỉ bước cập nhật đọc tới. Lượt xuôi không cần chúng. Suy luận thì không cần một byte nào trong số đó.
Đây là lý do thanh chọn bộ tối ưu ở góc trên đổi con số mạnh tới vậy. Adam giữ hai moment cho mỗi tham số, trung bình động của gradient và trung bình động của bình phương gradient, nên nó cõng thêm đúng hai lần số tham số nhân bề rộng moment. Bấm sang SGD có momentum: chỉ còn một moment, tổng byte mỗi tham số từ 16 xuống 12. Bấm sang SGD trơn: không moment nào, còn 8. Ở cấu hình mở bài, riêng cú bấm đó đưa tổng mỗi thiết bị từ 136,2 GiB xuống 76,6 GiB, tức vừa một thẻ 80 GiB mà không cần thêm bất cứ mẹo nào.
Và ngay đây phải nói thẳng một câu, vì sim rất dễ bị đọc quá: nó không nói rằng đổi sang SGD trơn thì mô hình vẫn học tốt như cũ. Nó chỉ nói cú bấm đó tiết kiệm bao nhiêu byte. Vì sao Adam được dùng gần như mặc định trong huấn luyện mô hình ngôn ngữ là chuyện thực nghiệm, và bạn xem lại các bộ tối ưu để biết hai moment kia dùng làm gì. Ở trang này chúng chỉ là hai khoản byte.
Kích hoạt: khoản không dính gì tới số tham số
Phần thứ hai của hoá đơn không nhân với số tham số. Lượt truyền ngược cần lại các giá trị trung gian mà lượt xuôi đã tính, nên chúng phải được giữ, và số lượng của chúng phụ thuộc lô, độ dài chuỗi, số lớp, chiều ẩn, chứ không phụ thuộc mô hình có bao nhiêu tham số. Nếu bạn còn mơ hồ vì sao lượt ngược lại cần chúng thì lan truyền ngược là chỗ để xem lại.
Ở trạng thái mở bài, sim ước lượng 17 tensor cỡ (lô × chuỗi × chiều ẩn) cho mỗi lớp, mỗi phần tử 2 byte:
17 × 4.096 × 2 = 139.264 byte cho mỗi token ở mỗi lớp
tức 136,0 KiB mỗi token mỗi lớp, nhân 32 lớp ra 4,3 MiB mỗi token, nhân 4.096 token ra 17,0 GiB. Cộng với 119,2 GiB trạng thái được 136,2 GiB cho mỗi thiết bị. Thẻ 80 GiB thiếu 56,2 GiB.
Con số 17 tensor ấy là xấp xỉ, và phải nói rõ nó xấp xỉ tới đâu. Ở mức 17 tensor nhân 2 byte, nó bằng 34 byte cho mỗi token mỗi lớp mỗi chiều ẩn, tức đúng con số 34·s·b·h trong bài của Korthikanti và cộng sự về giảm tái tính kích hoạt. Nhưng bản dẫn giải đó còn một hạng thức bậc hai theo độ dài chuỗi cho ma trận điểm chú ý, mà sim này không tính, vì nhân chú ý hiện nay thường không dựng hẳn ma trận đó ra. Nên hãy đọc con số kích hoạt ở đây là sàn có một hạng thức đã biết là thiếu, không phải dự báo. Ô số đó sửa được chính là để bạn thay bằng con số của mình.
Tái tính kích hoạt cắt được nhiều, nhưng không cắt đúng chỗ
Bấm bật tái tính kích hoạt. Cách làm là mỗi lớp chỉ giữ lại đầu vào của nó, còn phần bên trong thì lúc truyền ngược tính lại từ đầu vào đó. Sim mô hình hoá thành: giữ 1 tensor mỗi lớp thay vì 17.
Kết quả: kích hoạt từ 17,0 GiB xuống 1,0 GiB, tiết kiệm 16,0 GiB. Cái giá hiện ngay ô bên cạnh: +33,3% lượt tính, vì theo cách đếm một lượt xuôi và hai lượt ngược thì thêm một lượt xuôi nữa là từ 3 lên 4.
Nhưng nhìn tổng: 120,2 GiB. Vẫn không vừa thẻ 80 GiB. Vì mẹo này chỉ chạm vào phần kích hoạt, còn 119,2 GiB trạng thái mô hình thì không đổi một byte. Đây là điều mà một bảng chỉ hiện tổng sẽ không dạy được cho bạn: ba cần điều khiển của bài này tác động vào ba phần khác nhau của phép cộng, và biết cái nào chạm cái gì thì mới chọn đúng.
Nói thêm cho đủ: con số +33,3% là đếm lượt nhân ma trận, không phải đo đồng hồ. Thời gian thật còn phụ thuộc nhân, băng thông bộ nhớ và cách xếp lớp, và sim không đo được thứ nào trong đó.
Chia trạng thái cho nhiều thiết bị, và cái không chia được
Cần thứ ba là chia trạng thái theo bậc, mỗi bậc chia thêm một vai. Bậc 1 chia trạng thái bộ tối ưu, tức bản trọng số chủ và các moment. Bậc 2 chia thêm gradient. Bậc 3 chia thêm cả trọng số làm việc. Cột "chia từ bậc" trong bảng vai ghi rõ mỗi vai bắt đầu bị chia từ bậc nào.
Với 8 thiết bị, thang bậc ở trạng thái mở bài đọc như sau:
| bậc | mỗi thiết bị | vừa thẻ 80 GiB |
|---|---|---|
| 0 | 136,2 GiB | không |
| 1 | 58,0 GiB | vừa |
| 2 | 44,9 GiB | vừa |
| 3 | 31,9 GiB | vừa |
Bậc 1 một mình đã đủ để lọt, và cũng dễ hiểu: nó chia đúng cái phần chiếm ba phần tư. Kéo số thiết bị xuống thì các dòng này đổi theo, và sim cho biết ít nhất bao nhiêu thiết bị mới vừa: ở bậc 3 là 2 thiết bị, ở bậc 1 là 3 thiết bị, vì bậc 1 chỉ chia được 96 trong 128 tỉ byte.
Chỗ quan trọng nhất của mục này lại là chỗ không đổi: kích hoạt không được chia. Cách chia này là song song dữ liệu, mỗi thiết bị vẫn nạp lô riêng của nó, nên ô "cỡ lô mỗi thiết bị" đúng là lô của một thiết bị và 17,0 GiB kia đứng nguyên dù bạn có 8 hay 1.024 thiết bị. Muốn thấy chỗ đó cắn thật thì bấm preset Llama 3.1 70B: kích hoạt của nó là 85,0 GiB, và ở bậc 3 với 1.024 thiết bị mỗi thiết bị vẫn phải chứa 86,0 GiB, tức vẫn quá thẻ 80 GiB. Bật thêm tái tính thì tụt xuống 6,0 GiB. Chia bao nhiêu cũng không cứu được phần không chia.
Còn một chuyện sim không mô hình hoá và cần biết: bậc 3 chia cả tham số, nhưng lúc tính vẫn phải gom trọng số của từng lớp về rồi bỏ đi. Con số ở đây là mức thường trú, nên đỉnh thật cao hơn.
Vậy 8 tỉ tham số huấn luyện được ở đâu
Gộp lại thành câu trả lời cho tiêu đề. Cùng một mô hình 8 tỉ tham số, cùng một thẻ 80 GiB:
- Nạp trọng số để suy luận, 2 byte mỗi tham số: 14,9 GiB. Vừa thoải mái, vừa cả thẻ 24 GiB. Nhớ rằng đó là chưa tính bộ đệm KV, thứ mà bài trước đã đếm và ở ngữ cảnh dài còn nặng hơn cả trọng số.
- Huấn luyện đầy đủ ở cấu hình mở bài: 136,2 GiB. Không vừa, thiếu 56,2 GiB, tức tốn 9,1 lần so với chỉ nạp trọng số.
- Tái tính kích hoạt một mình: 120,2 GiB. Vẫn không vừa.
- Bậc 3 với 8 thiết bị một mình: 31,9 GiB. Vừa.
- Bậc 3 với 8 thiết bị cộng tái tính: 15,9 GiB. Vừa cả thẻ 24 GiB.
Nên câu "mô hình này chạy được trên thẻ X" và câu "mô hình này huấn luyện được trên thẻ X" là hai câu khác nhau hoàn toàn, và khoảng cách giữa chúng ở đây là gần một bậc độ lớn. Preset thứ tư trong sim ghi rõ là số minh hoạ, không phải mô hình nào cả, và nó ở đó để bạn thấy chỗ nào thì khoảng cách ấy hẹp lại: một mô hình 1,5 tỉ tham số với 24 lớp và chiều ẩn 2.048 tốn 28,7 GiB một bước huấn luyện đầy đủ, tức vừa một thẻ 80 GiB mà không cần mẹo nào.
Nếu bạn tự hỏi có thể cắt số byte của các vai xuống nữa không thì câu trả lời là có, và đó là chuyện của định dạng số: mỗi ô byte trong bảng vai chính là chỗ một định dạng ít bit hơn đi vào. Nhưng cắt bit của gradient và của moment không giống cắt bit của trọng số lúc suy luận, vì bước cập nhật cộng dồn những lượng rất nhỏ, và cái giá của việc cắt thì trang này không đo được.
Còn một đường khác hẳn, và nó không cắt bề rộng của các ô: cắt số tham số có gradient và có moment ngay từ đầu. Bảng ở đây giả định mọi tham số đều được cập nhật, nên cả năm vai đều nhân với trọn số tham số. Nếu bạn đóng băng phần lớn mô hình và chỉ huấn luyện một phần nhỏ thì bốn vai sau co lại theo phần nhỏ đó, còn vai trọng số làm việc thì không. Đó là chủ đề của tinh chỉnh hạng thấp, và nó là cách phổ biến nhất để đưa hoá đơn này về mức một thẻ.
Huấn luyện trả tiền cho cùng một tham số dưới năm cái tên: trọng số làm việc, gradient, bản trọng số chủ fp32, và một hoặc hai moment của bộ tối ưu. Đừng nhớ tổng, hãy nhớ năm cái tên đó, vì tổng ra 16, 18 hay 20 chỉ tuỳ chỗ bạn đặt fp32. Ở cấu hình mở bài, ba phần tư hoá đơn trạng thái là thứ chỉ bước cập nhật đọc tới, nên đổi bộ tối ưu là cú cắt mạnh nhất trên phần đó. Rồi nhớ tiếp rằng hoá đơn có hai nửa và ba mẹo cắt vào ba chỗ khác nhau: đổi bộ tối ưu cắt trạng thái, tái tính cắt kích hoạt, chia bậc cắt trạng thái mỗi thiết bị mà không cắt kích hoạt. Mô hình 8 tỉ tham số suy luận hết 14,9 GiB và huấn luyện đầy đủ hết 136,2 GiB: chỉ khi ghép chia bậc với tái tính thì nó mới xuống 15,9 GiB và vừa một thẻ. Cộng byte không phải là đo tốc độ hay đo chất lượng, và phần kích hoạt ở đây còn thiếu hẳn một hạng thức.
- 1Bạn đọc thấy ba chỗ khác nhau nói huấn luyện tốn 16, 18 và 20 byte cho mỗi tham số. Ai đúng?
- 2Mô hình 8 tỉ tham số ở cấu hình mở bài cần 136,2 GiB mỗi thiết bị, trong đó 119,2 GiB là trạng thái mô hình và 17,0 GiB là kích hoạt. Bạn bật tái tính kích hoạt. Tổng thành bao nhiêu, và nó vừa thẻ 80 GiB chưa?
- 3Preset Llama 3.1 70B có 85,0 GiB kích hoạt. Bạn đặt bậc chia trạng thái là 3 và số thiết bị là 1.024, tức chia mọi vai cho 1.024. Mỗi thiết bị còn phải chứa bao nhiêu, và vì sao?