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

Ngân sách tính toán để huấn luyện

Kế toán chính xácSố công bố có ngàyNói rõ chỗ công thức hụt

Ngân sách tính toán để huấn luyện

Toàn bộ chi phí huấn luyện một mô hình lớn gói được trong một phép nhân ba số. Bài này bắt bạn tự làm phép nhân đó, rồi chỉ đúng ba chỗ nó nói dối.

bộ đệm KVđịnh dạng số, thứ phải đếm là bộ nhớ lúc chạy. Bài này đếm một thứ khác, đắt hơn nhiều và chỉ trả một lần: số phép tính để huấn luyện xong mô hình.

Điều đáng ngạc nhiên là phép đếm ấy gần như không cần biết gì về kiến trúc. Với N tham số và D token dữ liệu, chi phí xấp xỉ 6 × N × D phép tính dấu phẩy động. Không có số lớp, không có số đầu chú ý, không có kiểu vị trí. Lý do thì thẳng thắn: một phép nhân ma trận dùng mỗi trọng số đúng một lần, và dùng nó thì tốn một phép nhân cộng một phép cộng, nên lượt xuôi qua một token chạm vào mỗi tham số hai lần, tức 2ND cho cả tập dữ liệu. Lượt ngược phải tính hai gradient, một theo đầu vào của lớp và một theo trọng số của lớp, mỗi cái lại là một phép nhân ma trận cùng cỡ, nên nó tốn khoảng 4ND. Cộng lại thành 6ND.

Sim dưới đây làm đúng chuỗi đó rồi đi tiếp: chia FLOPs cho công suất thật của thiết bị ra giờ máy, nhân giá thuê ra tiền. Mọi tham số đều sửa được, và nó cũng in ra thẳng cái mà 6ND bỏ qua.

Ngân sách tính toán để huấn luyện · từ 6ND tới hoá đơn
Tổng 722,7 ZFLOPGiờ GPU 507.201Tiền 1,27 triệu đô
Mô hình nhỏ nhưng nuốt 15 nghìn tỉ token, tức khoảng 1.868 token cho mỗi tham số, xa hẳn tỉ lệ 20 của Chinchilla. Đây là trạng thái mở bài.
Nguồn số: Thẻ mô hình Meta Llama 3 (công bố 18/04/2024) cho số giờ GPU, loại thiết bị và mốc "hơn 15 nghìn tỉ token"; tệp config.json công bố kèm mô hình cho 32 lớp, chiều rộng 4096, độ dài chuỗi 8192. Thẻ mô hình chỉ ghi "8B", còn 8,03 tỉ là số tham số đếm được từ bộ trọng số đã công bố. MFU 0,4 và giá 2,5 đô mỗi giờ là GIẢ ĐỊNH của sim, không phải số Meta công bố.
Số chép vào bài ngày 28/07/2026 và có thể đã lạc hậu: thẻ mô hình bị sửa, bảng giá thuê đổi hàng tháng. Mọi tham số bên dưới đều sửa được, nên có số mới hơn thì gõ vào là tính lại ngay.

1 · Mô hình và dữ liệu ✎ sửa được

Ô số token đang cầm lái, nên ô tỉ lệ tạm ẩn. Tỉ lệ bên dưới là số suy ra từ N và D, không phải ô nhập.
8,03
tỉ tham số, tức N = 8.030.000.000
15,00 nghìn tỉ token
D, tức 15.000.000.000.000 token
1.868,0
token cho mỗi tham số, tức 93,4 lần tỉ lệ 20
1.831.054.688
chuỗi dài 8.192 token xếp kín được từ D

2 · Phép nhân ra FLOPs

khoảncông thứcFLOPsdạng khoa học
lượt xuôi2 × N × D240,9 ZFLOP2,41 × 1023
lượt ngược4 × N × D481,8 ZFLOP4,82 × 1023
cộng lại, tức 6ND6 × N × D722,7 ZFLOP7,23 × 1023
phần chú ý mà 6ND bỏ qua6 × L × d × D × (s + 1)96,65 ZFLOP9,66 × 1022
TỔNG dùng để tính giờ6ND (bỏ phần chú ý)722,7 ZFLOP7,23 × 1023
Phần chú ý bằng 13,4% của con số 6ND, tức 11,8% của tổng thật. Bạn đang bỏ nó ngoài tổng, đúng như cách công thức 6ND làm, nên hoá đơn bên dưới thiếu đúng chỗ đó. Tỉ lệ này không phụ thuộc D: cả hai khoản đều tuyến tính theo số token nên D triệt tiêu, chỉ còn L × d × (s + 1) / N.

3 · Hình dạng mô hình, chỉ phần chú ý dùng tới ✎ sửa được

Một chuỗi dài 8.192 token mở s(s + 1)/2 ô mặt nạ nhân quả, tức 33.558.528 ô. Mỗi ô tốn 4 × d FLOPs ở lượt xuôi (16.384 FLOPs với chiều rộng 4.096), nhân 32 lớp, nhân 3 cho cả lượt ngược, nhân 1.831.054.688 chuỗi. Ba núm ở trên không làm đổi con số 6ND, vì N là số bạn gõ thẳng chứ không phải số dựng lại từ hình dạng. Chúng chỉ đổi phần chú ý, và đó là điều đáng nhìn: gấp đôi s thì phần chú ý gấp đôi theo mỗi token, còn 6ND đứng yên.

4 · Từ FLOPs ra giờ và ra tiền ✎ sửa được

989,5 TFLOPS đỉnh × MFU 0,40 = 395,8 TFLOP mỗi giây, nhân 3600 giây ra 1,42 EFLOP mỗi giờ máy. Chia tổng 722,7 ZFLOP cho con số đó là ra số giờ.

khoảncách tínhgiá trị
giờ GPUFLOPs / (đỉnh × MFU × 3600)507.201 giờ
đổi ra ngày máygiờ GPU / 2421.133,4 ngày
thời gian thực tế trên 1.024 thiết bịgiờ GPU / (số thiết bị × 24)20,6 ngày
tiền thuêgiờ GPU × giá mỗi giờ1.268.002 đô = 1,27 triệu đô

5 · Nếu huấn luyện đúng tỉ lệ 20

cùng N = 8,03 tỉ tham sốcấu hình đang xemtỉ lệ 20hơn bao nhiêu lần
số token15,00 nghìn tỉ token160,6 tỉ token93,4
FLOPs722,7 ZFLOP7,74 ZFLOP93,4
giờ GPU507.201 giờ5.430 giờ93,4
tiền thuê1.268.002 đô13.576 đô93,4
Cột bên phải không nói rằng tỉ lệ 20 tốt hơn. Nó chỉ nói cùng số tham số đó, huấn luyện tới điểm tối ưu theo compute thì rẻ hơn 93,4 lần. Người ta cố tình trả cái giá đắt hơn ấy vì mô hình nhỏ mà học nhiều token thì suy luận rẻ hơn suốt đời phục vụ, và sim này không đo được phần suy luận đó.

6 · Đối chiếu với số đã công bố

Preset này công bố 1.300.000 giờ máy. Thẻ mô hình ghi 1,3 triệu giờ GPU H100-80GB cho bản 8B (tổng cả họ là 7,7 triệu giờ).
MFU hàm ý = tổng FLOPs / (đỉnh × 3600 × giờ công bố) = 15,6%
Còn ước lượng của sim ở MFU 0,40 507.201 giờ, tức 2,56 lần ít hơn số công bố. Hai con số này lệch nhau và bài không bẻ số cho khớp: xem phần văn xuôi để biết vì sao con số hàm ý đáng nghi.
Sim này không làm được gì
  • đếm phép nhân. Nó không biết mô hình huấn luyện với ngân sách này có tốt hay không. Hai lần chạy cùng 6ND có thể ra hai mô hình chênh nhau rất xa, và không con số nào ở đây thấy được điều đó.
  • MFU là đầu vào, không phải phép đo. Mặc định 0,4 là con số hay được nhắc cho huấn luyện dày trên thiết bị hiện đại, nhưng sim không đo được nó. Chỗ duy nhất sim nói được điều gì thật về MFU là mục 6: khi có số giờ đã công bố thì nó tính ra MFU mà số giờ đó hàm ý.
  • 6ND tự nó là quy ước. Nó bỏ softmax, chuẩn hoá lớp, hàm kích hoạt, bước cập nhật của bộ tối ưu, phần nhúng, và toàn bộ chi phí truyền dữ liệu giữa các máy. Phần chú ý ở mục 2 cũng chỉ đếm nửa tam giác của mặt nạ, tức một cài đặt có mặt nạ tốt. Nên tổng ở đây là chặn dưới, không phải hoá đơn thật.
  • Giá thuê là con số bịa cho dễ nhìn. Không ai huấn luyện mô hình lớn bằng cách thuê lẻ theo giờ, và giá công khai đổi liên tục. Coi cột tiền là thang đo độ lớn, đừng coi là báo giá.
  • Cổng kiểm số của bài canh phép tính: cấu hình này thì FLOPs, giờ và tiền phải bằng đúng những con số này. Nó không canh, và không thể canh, chuyện mô hình ngoài kia có đúng tham số đó hay không.

Bốn phép nhân ở trạng thái mở bài

Preset mở bài là Llama 3 8B: N bằng 8,03 tỉ tham số, D bằng 15 nghìn tỉ token.

Phép nhân thứ nhất, ra FLOPs. Lượt xuôi 2 × 8,03 × 10^9 × 15 × 10^12 bằng 240,9 ZFLOP, lượt ngược gấp đôi con số đó thành 481,8 ZFLOP, tổng 6ND722,7 ZFLOP, hay viết cách khác 7,23 × 10^23. Một số ít người có cảm giác trực quan về đại lượng này, nên hãy neo nó lại như sau: đó là hơn bảy trăm tỉ tỉ phép tính, và toàn bộ ước lượng gói trong ba con số.

Phép nhân thứ hai, ra giờ máy. Thiết bị mặc định là H100 SXM với 989,5 TFLOPS ở bf16. Con số đó lấy từ thông số nhà sản xuất công bố là 1.979 TFLOPS rồi chia hai, vì mốc 1.979 là mốc có dùng sparsity, mà huấn luyện dày thì không được nửa đó. Đem nhân với MFU 0,4 ra 395,8 TFLOP mỗi giây, nhân 3.600 giây ra 1,42 EFLOP mỗi giờ máy. Chia 722,7 ZFLOP cho con số đó được 507.201 giờ GPU, tức 21.133,4 ngày máy, tức 20,6 ngày thực tế nếu 1.024 thiết bị chạy song song trọn thời gian.

Phép nhân thứ ba, ra tiền. 507.201 giờ nhân 2,5 đô mỗi giờ bằng 1.268.002 đô, tức 1,27 triệu đô. Cả giá thuê lẫn MFU đều là con số tôi đặt ra, không phải số Meta công bố, và bên dưới sẽ thấy vì sao chỗ này đáng ngờ nhất trong cả bài.

Phép nhân thứ tư, cái mà 6ND không thấy. Mọi phép nhân ma trận có trọng số đều nằm trong 6ND. Nhưng hai phép nhân ma trận trong lớp chú ý thì không có trọng số nào: QK^T lấy truy vấn nhân khoá, rồi trọng số chú ý nhân giá trị. Chi phí của chúng lớn lên theo độ dài chuỗi, không theo N, nên 6ND đơn giản là không nhìn thấy chúng.

Chỗ công thức hụt, đo bằng số

Đếm phần chú ý thì phải quay lại đúng cái tam giác của cửa sổ trượt và của self-attention. Một chuỗi dài s với mặt nạ nhân quả mở s(s + 1)/2 ô. Ở độ dài 8.192, đó là 33.558.528 ô cho mỗi chuỗi. Mỗi ô tốn 4 × d FLOPs ở lượt xuôi, tức 16.384 FLOPs với chiều rộng 4.096, vì QK^T góp 2 × d và phép nhân với giá trị góp thêm 2 × d khi cộng dồn qua các đầu. Nhân 32 lớp, nhân 3 cho cả lượt ngược, nhân 1.831.054.688 chuỗi mà 15 nghìn tỉ token xếp kín được, ra 96,65 ZFLOP.

96,65 so với 722,7 là 13,4%. Nói cách khác 6ND hụt khoảng một phần tám ở cấu hình này, và bấm nút cộng vào tổng thì hoá đơn đi từ 722,7 lên 819,3 ZFLOP, từ 507.201 lên 575.030 giờ, từ 20,6 lên 23,4 ngày, từ 1,27 lên 1,44 triệu đô.

Chuyện đáng nhìn hơn là tỉ lệ ấy đổi theo cái gì. Nó bằng L × d × (s + 1) / N, tức không phụ thuộc D chút nào: cả hai khoản đều tuyến tính theo số token nên D triệt tiêu. Cái nó phụ thuộc là độ dài chuỗi. Giữ nguyên mô hình 8B rồi kéo s từ 8.192 lên 131.072 mà xem: con số 6ND đứng yên, còn phần chú ý gấp 16 lần, và tỉ lệ đi từ 13,4% lên 213,9% của 6ND, tức chiếm 68,1% tổng thật. Ở ngưỡng đó công thức quen thuộc không còn là ước lượng gần đúng nữa, nó sai hơn hai lần. Preset "minh hoạ" cuối hàng bật sẵn phần chú ý ở độ dài 131.072 để bạn thấy chuyện này với một mô hình to hơn: 43,0% của 6ND, tức 30,0% tổng.

Cần nói rõ ranh giới của chính phép đếm phần chú ý. Nó chỉ đếm nửa tam giác, tức giả định một cài đặt biết bỏ qua nửa ô bị mặt nạ chặn; đếm cả hình vuông thì con số gấp đôi. Nó cũng vẫn bỏ softmax, chuẩn hoá lớp, hàm kích hoạt, bước cập nhật của bộ tối ưu, phần nhúng và toàn bộ chi phí truyền dữ liệu giữa các máy. Vậy nên tổng trong sim là chặn dưới, không phải hoá đơn. Sim này biết nhân, nó không biết lập lịch nhân.

Tỉ lệ 20, và vì sao gần như không ai dùng nữa

6ND có hai biến chứ không phải một, nên cùng một ngân sách chia được nhiều cách: mô hình to học ít token, hay mô hình nhỏ học nhiều token. Bài Chinchilla trả lời câu đó bằng thực nghiệm và ra một tỉ lệ dễ nhớ: với ngân sách cố định, chỗ tốt nhất nằm quanh D ≈ 20 × N. Preset Chinchilla 70B là đúng điểm ấy, 70 tỉ tham số nhân 20 ra 1,4 nghìn tỉ token, tốn 588,0 ZFLOP.

Giờ soi lại preset mở bài bằng thước đó. Llama 3 8B dùng khoảng 1.868 token cho mỗi tham số, tức 93,4 lần tỉ lệ 20. Bảng số 5 tính hộ phần đối chiếu: cùng 8,03 tỉ tham số mà huấn luyện tới đúng điểm Chinchilla thì chỉ cần 160,6 tỉ token, 7,74 ZFLOP, 5.430 giờ máy, 13.576 đô. Rẻ hơn 93,4 lần.

Vậy tại sao có người tự nguyện trả cái giá đắt gấp chín mươi lần? Vì Chinchilla tối ưu chi phí huấn luyện, mà chi phí huấn luyện trả một lần còn chi phí suy luận trả mãi mãi. Một mô hình 8 tỉ tham số học rất nhiều token có thể đạt chất lượng của một mô hình lớn hơn nhiều, rồi phục vụ hàng tỉ lượt gọi với bộ đệm KV nhỏ hơn và độ trễ thấp hơn. Đổi mấy triệu đô một lần để mỗi lượt gọi rẻ đi mãi là một phép tính hoàn toàn khác, và sim này không tính được phép tính đó: nó không biết gì về suy luận, cũng không biết gì về chất lượng. Cột bên phải của bảng 5 nói "rẻ hơn 93,4 lần", nó không nói "tốt hơn".

Còn GPT-3 175B thì lệch về phía ngược lại, chỉ 1,71 token cho mỗi tham số, tức dưới tỉ lệ 20 hơn mười lần. Đó chính là kết luận mà bài Chinchilla đưa ra khi nhìn lại lứa mô hình đó: chúng quá lớn so với lượng dữ liệu đã học.

Số giờ GPU công bố nói gì, và chỗ tôi không bẻ số

Preset GPT-3 có một thứ quý: bài báo công bố thẳng tổng compute, 3.640 PF-ngày, tức 3.640 × 10^15 × 86.400 bằng 314,5 ZFLOP. Còn 6ND trên đúng ND đã công bố cho ra 315,0 ZFLOP. Lệch 0,16%. Với một công thức bỏ qua toàn bộ phần chú ý và mọi thứ ngoài phép nhân ma trận, khớp tới bốn chữ số là chuyện đáng chú ý, và nó cũng gợi rằng con số trong bài báo bản thân nó được tính theo lối kế toán tương tự. Bằng chứng nhỏ cho ý đó: cộng thêm phần chú ý vào thì khoảng lệch rộng ra, không hẹp lại.

Hai preset Llama 3 thì cho một thứ khác: số giờ GPU. Thẻ mô hình ghi 1,3 triệu giờ H100 cho bản 8B và 6,4 triệu giờ cho bản 70B. Từ đó lật ngược phép chia được MFU mà con số ấy hàm ý:

  • bản 8B: 722,7 ZFLOP trên 1,3 triệu giờ H100 hàm ý MFU 15,6%, hay 17,7% nếu tính cả phần chú ý;
  • bản 70B: 6,35 YFLOP trên 6,4 triệu giờ hàm ý MFU 27,9%, hay 30,0% nếu tính cả phần chú ý.

Con số của bản 70B nghe hợp lý. Con số của bản 8B thì không: 15,6% là thấp bất thường cho huấn luyện dày trên H100, và điều làm nó đáng nghi không phải bản thân nó mà là khoảng cách giữa hai bản, vì cả hai chạy trên cùng một hạ tầng của cùng một đội, ở cùng thời điểm.

Có nhiều cách giải thích khả dĩ. "Giờ GPU" trên thẻ mô hình có thể là giờ đã cấp phát chứ không phải giờ tính toán, tức gồm cả khởi động lại sau sự cố, gồm cả đánh giá giữa chừng, gồm cả lúc máy chờ. Mô hình nhỏ hơn khó chia việc cho cụm lớn nên phần chờ chiếm tỉ lệ cao hơn. Mốc "hơn 15 nghìn tỉ token" có thể là số làm tròn xuống. Có thể phần hậu huấn luyện cũng bị cộng vào.

Điều tôi không làm là chọn một MFU nào cho mỗi mô hình để hai bên khớp nhau. Làm thế thì con số nào cũng khớp và bài học biến mất. Sim in ra MFU hàm ý, gọi nó đúng tên là số hàm ý, rồi để nó trông kỳ như nó vốn kỳ. Với MFU 0,4 mặc định, ước lượng của sim cho bản 8B là 507.201 giờ, tức số công bố cao gấp 2,56 lần. Chênh lệch đó là dữ kiện, không phải lỗi cần vá.

Đọc preset cho đúng

Preset mang tên mô hình chỉ chép lại tham số đã công bố, kèm nguồn và ngày chép, và ngày đó hiện ngay dưới hàng preset. Vài điểm nhỏ nhưng hay bị bỏ qua:

  • Khối thiết bị gần như luôn là giả định của tôi. Chinchilla huấn luyện trên TPU, GPT-3 trên V100, nên gán H100 cho chúng là câu hỏi ngược "nếu chạy lại hôm nay tốn bao nhiêu", không phải mô tả chuyện đã xảy ra. Chỗ duy nhất phần cứng là số công bố thật là hai preset Llama 3, và đó cũng là lý do chỉ ở đó mới tính MFU hàm ý.
  • N là số bạn gõ vào, không phải số dựng lại từ hình dạng. Ba núm số lớp, chiều rộng và độ dài chuỗi không làm đổi con số 6ND, chúng chỉ đổi phần chú ý. Nếu muốn N đổi thì phải sửa ô N. Sim không suy N từ kiến trúc vì nó sẽ phải đoán số chiều lớp truyền thẳng, cách gắn bó trọng số nhúng và mấy lựa chọn khác.
  • MFU bị kẹp trong khoảng 0,01 tới 1. Bằng 1 nghĩa là thiết bị chạy đúng mốc đỉnh nhà sản xuất công bố, và không lần chạy nào vượt được mốc đó. Gõ 1,01 thì hệ kẹp về 1 và nói cho bạn biết là đã kẹp, thay vì âm thầm sửa. Mọi ô số đều xử sự như vậy.
  • Preset cuối ghi rõ là số minh hoạ, không phải cấu hình của mô hình nào cả, vì thà nói là minh hoạ còn hơn dán tên một mô hình lên những con số không kiểm được.
Điều rút ra

Ngân sách huấn luyện gói trong 6 × N × D, trong đó 2 phần là lượt xuôi và 4 phần là lượt ngược, và trên GPT-3 công thức đó khớp số công bố tới 0,16%. Chia tiếp cho đỉnh × MFU × 3600 là ra giờ máy, nhân giá là ra tiền: 722,7 ZFLOP thành 507.201 giờ H100 và 1,27 triệu đô ở mức MFU 0,4. Hai chỗ công thức gãy đều đáng nhớ. Thứ nhất, nó không thấy phần chú ý bậc hai, thứ hụt đó là 13,4% ở chuỗi 8.192 nhưng thành 213,9% ở chuỗi 131.072. Thứ hai, MFU là thứ bạn giả định chứ không đo được, và khi có số giờ đã công bố để lật ngược thì con số hàm ý có thể ra 15,6%, thấp tới mức phải nghi cách người ta đếm giờ. Đếm được phép nhân không có nghĩa là biết mô hình sẽ tốt: tỉ lệ 20 của Chinchilla rẻ hơn 93,4 lần mà gần như không ai chọn nữa, vì chi phí suy luận trả mãi mãi còn chi phí huấn luyện trả một lần.

Câu hỏi tự kiểm0/3 đúngchưa trả lời
  1. 1Một mô hình 7 tỉ tham số huấn luyện trên 2 nghìn tỉ token. Ngân sách huấn luyện xấp xỉ bao nhiêu FLOPs?
  2. 2Vẫn mô hình 8,03 tỉ tham số với 32 lớp và chiều rộng 4.096, vẫn 15 nghìn tỉ token, nhưng huấn luyện ở độ dài chuỗi 131.072 thay vì 8.192. Con số 6ND và phần chú ý đổi thế nào?
  3. 3Thẻ mô hình ghi 1,3 triệu giờ H100 cho Llama 3 8B. Sim tính 6ND ra 722,7 ZFLOP, nên số giờ đó hàm ý MFU 15,6%, trong khi bản 70B cùng cụm máy hàm ý 27,9%. Nên kết luận gì?