Đọc config.json rồi đếm ra số tham số
Đọc config.json rồi đếm ra số tham số
Bảy con số trong một tệp văn bản là đủ để tự tính ra con số mà nhà sản xuất in trên thẻ mô hình. Bài này bắt bạn tự cộng, từng ma trận một, rồi đối chiếu.
Cả học phần này đã đi qua từng bộ phận: lan truyền xuôi qua một lớp tuyến tính, chú ý theo nhóm làm hẹp chiếu khoá và chiếu giá trị, khối tiến có cổng, chuẩn hoá, bảng nhúng. Giờ là bài chốt, và nó không thêm khái niệm mới nào cả. Nó chỉ yêu cầu một việc: mở tệp config.json của một mô hình đã phát hành, lấy ra bảy con số, rồi cộng cho tới lúc bạn đứng ngang con số nhà sản xuất công bố.
Việc này quan trọng hơn vẻ ngoài của nó. Khi bạn tự cộng ra được 8.030.261.248 cho một mô hình mà thiên hạ gọi gọn là "8B", bạn thôi coi số tham số là một nhãn dán và bắt đầu coi nó là một phép tính có thể kiểm. Bạn cũng lập tức trả lời được những câu hỏi mà trước đó phải tra: tăng từ vựng lên gấp bốn thì mô hình dày thêm bao nhiêu, bật tie_word_embeddings thì tiết kiệm được gì, và tham số của mô hình thật ra nằm ở đâu. Câu cuối cùng là chỗ hầu hết người học đoán ngược.
1 · Các trường của config.json ✎ sửa được
{
"vocab_size": 128256,
"hidden_size": 4096,
"num_hidden_layers": 32,
"num_attention_heads": 32,
"num_key_value_heads": 8,
"intermediate_size": 14336,
"tie_word_embeddings": false
}hidden_size 4.096 / num_attention_heads 32 = 128. Trường bạn gõ (128) đang bị bỏ qua. Từ đó ra hai bề rộng của khối chú ý: num_attention_heads × head_dim = 4.096 cho chiếu truy vấn, và num_key_value_heads × head_dim = 1.024 cho chiếu khoá và chiếu giá trị.2 · Cộng từng tensor một
| tensor trong tệp trọng số | hình dáng | mỗi bản | số bản | tham số | phần của mô hình |
|---|---|---|---|---|---|
model.embed_tokens.weightbảng nhúng token | 128.256 × 4.096 | 525.336.576 | 1 | 525.336.576 | 6.54% |
layers.*.self_attn.q_proj.weightchiếu truy vấn | 4.096 × 4.096 | 16.777.216 | 32 | 536.870.912 | 6.69% |
layers.*.self_attn.k_proj.weightchiếu khoá | 1.024 × 4.096 | 4.194.304 | 32 | 134.217.728 | 1.67% |
layers.*.self_attn.v_proj.weightchiếu giá trị | 1.024 × 4.096 | 4.194.304 | 32 | 134.217.728 | 1.67% |
layers.*.self_attn.o_proj.weightchiếu đầu ra của khối chú ý | 4.096 × 4.096 | 16.777.216 | 32 | 536.870.912 | 6.69% |
layers.*.mlp.gate_proj.weightcổng của khối tiến | 14.336 × 4.096 | 58.720.256 | 32 | 1.879.048.192 | 23.40% |
layers.*.mlp.up_proj.weightnhánh lên của khối tiến | 14.336 × 4.096 | 58.720.256 | 32 | 1.879.048.192 | 23.40% |
layers.*.mlp.down_proj.weightnhánh xuống của khối tiến | 4.096 × 14.336 | 58.720.256 | 32 | 1.879.048.192 | 23.40% |
layers.*.input_layernorm.weightchuẩn hoá trước khối chú ý | 4.096 (vectơ) | 4.096 | 32 | 131.072 | 0.0016% |
layers.*.post_attention_layernorm.weightchuẩn hoá trước khối tiến | 4.096 (vectơ) | 4.096 | 32 | 131.072 | 0.0016% |
model.norm.weightchuẩn hoá cuối | 4.096 (vectơ) | 4.096 | 1 | 4.096 | 0.0001% |
lm_head.weightđầu ra ngôn ngữ | 4.096 × 128.256 | 525.336.576 | 1 | 525.336.576 | 6.54% |
| tổng cộng | 8.030.261.248 | 100% | |||
2 × hidden × (dQ + dKV)3 × hidden × intermediate3 · Tham số nằm ở đâu
Khối nào lớn hơn thì do một bất đẳng thức quyết định, và nó đọc được thẳng từ tệp cấu hình: khối tiến vượt khối chú ý khi 3 × intermediate_size lớn hơn 2 × (dQ + dKV). Ở cấu hình này hai bên là 43.008 và 10.240, nên phần thắng thuộc về khối tiến. Preset cuối cùng đặt riêng để lật ngược chiều này, nên đừng nhớ nó như một luật.
4 · So với con số nhà sản xuất công bố
| tôi đếm được | 8.030.261.248 |
|---|---|
| nhà sản xuất ghi | 8,03 tỉ, tức 8.030.000.000 |
| lệch | +261.248 tham số, tức 0.0033% |
| độ chính xác của con số công bố | viết tới chữ số cuối là 10.000.000, nên nó không phân biệt được khoảng cách dưới 5.000.000 |
- Nó đếm tham số theo một quy ước kiến trúc cố định: khối tiến ba ma trận, hai lớp RMSNorm mỗi khối, không độ chệch ở đâu cả. Mô hình nào lệch khỏi quy ước đó thì tổng ở đây lệch theo, và phần chú thích của preset nói rõ lệch bao nhiêu.
- Nó không đọc tệp trọng số. Nó đọc bảy con số, còn tệp thật có thể chứa tensor mà tệp cấu hình không nhắc tới. Muốn chắc thì phải liệt kê tensor trong checkpoint.
- Số tham số không phải bộ nhớ chạy máy, không phải tốc độ, không phải chất lượng. Bao nhiêu byte thì còn tuỳ định dạng số, và phần bộ nhớ khi huấn luyện nằm ở bài bộ nhớ khi huấn luyện.
- Với mô hình trộn chuyên gia thì phép đếm này sai hẳn về mặt khái niệm: tổng tham số và tham số hoạt động là hai con số khác nhau, và đó là chuyện của bài định tuyến và tham số hoạt động.
- Phần “khớp với con số công bố” chỉ khẳng định hai con số không chống nhau tới độ chính xác đã công bố. Cổng kiểm số của bài canh phép trừ và phép so, chứ không canh chuyện nhà sản xuất ghi đúng.
Bảy trường, và cái bẫy nằm ở trường thứ sáu
Sim mở bài ở tệp cấu hình của Llama 3 8B, và bạn đọc được ngay năm trường đầu: vocab_size 128.256, hidden_size 4.096, num_hidden_layers 32, num_attention_heads 32, num_key_value_heads 8, intermediate_size 14.336. Trường tie_word_embeddings để false.
Trường thứ sáu, head_dim, thì không có trong tệp đó. Đây không phải chuyện nhỏ: mọi bề rộng của khối chú ý đều dựng từ nó. Khi tệp không khai báo, bộ nạp suy ra bằng hidden_size / num_attention_heads, ở đây là 4096 / 32 = 128. Sim hiện cả hai con số cạnh nhau và nói rõ nó đang dùng cái nào, vì một giá trị suy ra mà trông như giá trị đọc được là đúng kiểu sai lặng lẽ. Có mô hình khai báo hẳn head_dim và giá trị đó không bằng thương kia: bấm preset Gemma 2 2B thì head_dim là 256 trong khi 2304 / 8 = 288, và tắt công tắc head_dim sẽ thấy tổng lệch hẳn đi. Ngược lại, ở trạng thái mở bài giá trị gõ và giá trị suy ra trùng nhau đúng ở 128, nên bật tắt công tắc không đổi con số nào, và sim nói thẳng ra như vậy chứ không để bạn tưởng mình vừa làm gì.
Từ head_dim ra hai bề rộng, và giữ chúng riêng biệt là toàn bộ khác nhau giữa đếm đúng và đếm sai:
num_attention_heads × head_dim = 32 × 128 = 4096, bề rộng của chiếu truy vấn và chiếu đầu ra.num_key_value_heads × head_dim = 8 × 128 = 1024, bề rộng của chiếu khoá và chiếu giá trị.
Cộng từng ma trận một
Bảng phân tách trong sim liệt kê mười hai tensor, đúng tên như chúng nằm trong tệp trọng số, nên bạn kiểm được từng dòng bằng máy tính cầm tay. Ở trạng thái mở bài:
Bảng nhúng token. vocab_size × hidden_size = 128.256 × 4.096 = 525.336.576. Một bản duy nhất, không thuộc lớp nào.
Khối chú ý, cho mỗi lớp. Bốn ma trận: chiếu truy vấn 4096 × 4096 = 16.777.216, chiếu khoá 1024 × 4096 = 4.194.304, chiếu giá trị cùng hình dáng nên cũng 4.194.304, chiếu đầu ra 4096 × 4096 = 16.777.216. Cộng lại 41.943.040, viết gọn thành 2 × hidden × (dQ + dKV). Nếu bạn dùng num_attention_heads cho cả chiếu khoá thì con số này phồng lên 67.108.864, và đó chính là chỗ chú ý theo nhóm tiết kiệm được.
Khối tiến kiểu SwiGLU, cho mỗi lớp. Ba ma trận, không phải hai: cổng, nhánh lên, nhánh xuống. Mỗi cái 4096 × 14336 = 58.720.256, tổng 3 × 4096 × 14336 = 176.160.768. Đếm hai ma trận là sai ngay một phần ba khối lớn nhất của mô hình.
Chuẩn hoá. Hai lớp RMSNorm mỗi khối, mỗi lớp là một vectơ dài bằng chiều ẩn và không có độ chệch, nên 2 × 4096 = 8.192 cho mỗi lớp. Thêm một lớp chuẩn hoá cuối trước đầu ra, 4.096. Cả mô hình chỉ có 32 × 8192 + 4096 = 266.240 tham số chuẩn hoá, tức 0,0033%. Sim để bốn chữ số thập phân cho riêng dòng này, vì với hai chữ số nó ra 0,00% và người đọc kết luận là bằng không.
Đầu ra ngôn ngữ. hidden_size × vocab_size = 525.336.576, bằng đúng bảng nhúng, vì tie_word_embeddings đang tắt.
Cộng lại: một lớp là 41.943.040 + 176.160.768 + 8.192 = 218.112.000, ba mươi hai lớp là 6.979.584.000, thêm bảng nhúng 525.336.576, thêm chuẩn hoá cuối 4.096, thêm đầu ra 525.336.576, ra 8.030.261.248.
Chỗ hầu hết người tự đếm sai
Công tắc tie_word_embeddings. Khi nó bật, tệp trọng số không có tensor lm_head.weight nữa: đầu ra dùng lại đúng bảng nhúng. Trong bảng phân tách, dòng đó vẫn còn để bạn thấy nó đáng lẽ tốn bao nhiêu, nhưng cột số bản ghi 0 và cột tham số ghi 0.
Bật nó lên ở trạng thái mở bài thì tổng tụt từ 8.030.261.248 xuống 7.504.924.672, tức bay đúng 525.336.576 tham số, hơn nửa tỉ. Với từ vựng 32.000 của Mistral 7B thì cùng công tắc ấy chỉ cắt 131.072.000, vì phần bị xoá luôn đúng bằng vocab_size × hidden_size. Đây là lý do một người tự cộng theo trí nhớ rất dễ lệch nửa tỉ: họ nhân bảng nhúng lên hai lần cho một mô hình có nhúng dùng chung, hoặc chỉ đếm một lần cho một mô hình không dùng chung. Trong bảy preset thì Llama 3.2 1B, Gemma 2 2B và preset minh hoạ bật công tắc này, bốn preset còn lại tắt.
Tham số nằm ở đâu, và vì sao ai cũng đoán ngược
Bảng phần trăm là phần đáng giá nhất của bài. Ở trạng thái mở bài:
| khối | tham số | phần của mô hình |
|---|---|---|
| khối tiến, cả 32 lớp | 5.637.144.576 | 70,20% |
| khối chú ý, cả 32 lớp | 1.342.177.280 | 16,71% |
| bảng nhúng token | 525.336.576 | 6,54% |
| đầu ra ngôn ngữ | 525.336.576 | 6,54% |
| chuẩn hoá, tất cả | 266.240 | 0,0033% |
Chú ý là thứ mọi bài báo nói tới, mọi hình vẽ tô đậm, mọi bài giảng dành nhiều thời gian nhất. Nó chiếm 16,71%. Khối tiến, thứ thường bị gọi qua là "một MLP hai lớp", chiếm 70,20%, tức gấp 4,2 lần khối chú ý ngay trong cùng một lớp: 176.160.768 so với 41.943.040.
Thứ tự đó không phải luật, và ở đây có thể chứng minh cả hai chiều. Khối tiến vượt khối chú ý đúng khi 3 × intermediate_size lớn hơn 2 × (dQ + dKV). Ở trạng thái mở bài hai bên là 43.008 và 10.240, nên khối tiến thắng rất xa. Preset cuối cùng của sim đặt riêng để lật ngược: chú ý nhiều đầu thuần cộng intermediate_size bằng đúng chiều ẩn cho 12.288 so với 16.384, và lúc đó khối chú ý mới là khối lớn hơn, 67.108.864 so với 50.331.648. Dòng nhận xét dưới sim đổi lời theo, vì nó tính chứ không đọc thuộc. Đo trên cả sáu preset mang tên mô hình thật thì khối tiến luôn thắng, luôn trên 60% mô hình, còn khối chú ý luôn dưới 32%. Nhưng đó là kết quả đo trên sáu cấu hình, không phải một định lý.
Bảng nhúng thì tuỳ hai chuyện khác nhau, và hai preset tách được chúng ra. So Llama 3 8B với Mistral 7B v0.1: khối của hai mô hình giống hệt, cùng 218.112.000 tham số mỗi lớp, chỉ khác từ vựng 128.256 so với 32.000. Chênh lệch giữa hai tổng là 788.529.152, đúng bằng 2 × (128.256 - 32.000) × 4.096, và phần nhúng tụt từ 6,54% xuống 1,81%. Còn Llama 3.2 1B giữ nguyên từ vựng 128.256 nhưng chỉ có 16 lớp hẹp, nên bảng nhúng nhảy lên 21,25% mô hình dù kích thước tuyệt đối của nó còn nhỏ hơn. Nói cách khác, tỉ lệ của bảng nhúng do từ vựng đặt cạnh chồng lớp quyết định, không do mô hình to hay nhỏ.
Hai kết luận này đi tiếp vào các bài khác. Khối tiến chiếm phần lớn tham số là lý do tinh chỉnh LoRA lắp bộ chuyển vào những ma trận đó thì đắt hơn hẳn so với chỉ lắp vào chiếu truy vấn và chiếu giá trị. Nó cũng là lý do kiến trúc trộn chuyên gia chọn đúng khối tiến để nhân bản: đó là chỗ có nhiều tham số nhất để tách ra, và với mô hình như vậy thì phép đếm của bài này sai hẳn về mặt khái niệm, vì tổng tham số và tham số hoạt động là hai con số khác nhau.
"Khớp với con số công bố" nghĩa là gì
Nhà sản xuất Llama 3 8B ghi 8,03 tỉ tham số. Tôi đếm được 8.030.261.248. Hai con số đó lệch +261.248, tức 0,0033%.
Đây là chỗ dễ tự lừa mình, nên sim làm phép so một cách rõ ràng. Con số công bố viết tới ba chữ số có nghĩa, tức chữ số cuối của nó đứng ở hàng 10.000.000. Vậy mọi tổng thật nằm giữa 8,025 và 8,035 tỉ đều làm tròn về đúng 8,03 tỉ, và khoảng cách tối đa mà con số công bố không thể phân biệt là 5.000.000. Khoảng cách 261.248 của tôi nhỏ hơn hẳn, nên hai con số không hề chống nhau. Nhưng chữ "khớp" ở đây chỉ có nghĩa là khớp tới độ chính xác đã công bố, không phải khớp từng đơn vị: nếu nhà sản xuất viết 8,0290 tỉ thì đúng cùng phép đếm ấy sẽ lệch, vì lúc này mức không phân biệt được chỉ còn 500.000.
Sáu preset mang tên mô hình đều nằm trong mức đó, và khoảng cách lớn nhất là 0,3375% của Llama 3.2 1B:
| preset | tôi đếm được | công bố | lệch |
|---|---|---|---|
| Llama 3 8B | 8.030.261.248 | 8,03 tỉ | +261.248 |
| Llama 2 7B | 6.738.415.616 | 6,74 tỉ | -1.584.384 |
| Mistral 7B v0.1 | 7.241.732.096 | 7,24 tỉ | +1.732.096 |
| Qwen2.5 7B | 7.615.487.488 | 7,62 tỉ | -4.512.512 |
| Gemma 2 2B | 2.614.222.080 | 2,61 tỉ | +4.222.080 |
| Llama 3.2 1B | 1.235.814.400 | 1,24 tỉ | -4.185.600 |
Còn khi lệch thật thì phải nói là lệch. Sim liệt kê các lý do khả dĩ theo thứ tự hay gặp: quy ước đếm lớp nhúng khác đi, vì có nơi báo số tham số không tính lớp nhúng; mô hình có thêm độ chệch ở một số phép chiếu; số lớp chuẩn hoá mỗi khối khác 2; khối tiến không phải kiểu ba ma trận; hoặc có lớp phụ mà tệp cấu hình không nói ra. Hai preset ở đây rơi đúng vào hai lý do đó, và phần chú thích nguồn của chúng nói ra bằng số: Qwen2.5 7B có độ chệch ở ba chiếu Q, K, V mà sim bỏ qua, tức thiếu 4.608 tham số mỗi lớp và 129.024 trên cả mô hình; Gemma 2 dùng bốn lớp chuẩn hoá mỗi khối chứ không phải hai, tức thiếu 119.808. Cả hai khoản đó nhỏ hơn mức làm tròn của con số công bố, nên chúng không lộ ra trong bảng trên, nhưng chúng có thật và bài phải ghi. Cái không được làm là bẻ một con số cho khớp.
Đọc preset cho đúng
Preset mang tên mô hình chỉ chép lại các trường đã công bố, kèm tên kho, danh sách từng trường lấy về, con số công bố với độ chính xác của nó, và ngày chép. Ngày đó hiện ngay dưới hàng preset. Lĩnh vực này đổi rất nhanh, một bản mới có thể đổi cả từ vựng lẫn số lớp, nên hãy coi preset là ví dụ đã đóng băng thời điểm. Preset thứ bảy ghi rõ là số minh hoạ, không phải mô hình nào cả, và nó không có con số công bố: bịa ra một con số nhà sản xuất cho một cấu hình mình tự đặt sẽ là kiểu giả mạo tệ nhất, nên chỗ đó để trống.
Vài điểm nhỏ nhưng hay bị bỏ qua:
- Sim đếm theo một quy ước kiến trúc cố định: khối tiến ba ma trận, hai lớp RMSNorm mỗi khối, không độ chệch ở đâu cả. Đó là quy ước của họ Llama, không phải chân lý. Mô hình lệch khỏi quy ước thì tổng lệch theo, và phần chú thích của preset nói lệch bao nhiêu.
- Số tham số không phải bộ nhớ. Bao nhiêu byte còn tuỳ định dạng số, và bộ nhớ lúc huấn luyện thì lớn hơn nhiều lần vì còn trạng thái bộ tối ưu, xem bộ nhớ khi huấn luyện. Con số bạn vừa cộng ra ở đây chính là đầu vào của bài mô hình này chạy được trên máy nào: ở đó nó được nhân với số byte mỗi tham số rồi cộng thêm bộ đệm KV.
- Bảy con số không phải cả tệp. Một
config.jsonthật cònrope_theta,max_position_embeddings,rms_norm_epsvà nhiều trường khác. Chúng đổi hành vi của mô hình nhưng không thêm một tham số nào, nên bài này không nhắc tới. Cũng có thể tệp trọng số chứa tensor mà tệp cấu hình không nói ra, và cách duy nhất để chắc là mở checkpoint ra liệt kê. - Giá trị ngoài khoảng bị kẹp và được báo. Gõ
hidden_sizebằng một triệu thì sim kẹp về16.384và hiện một dòng cảnh báo nói rõ bạn đặt bao nhiêu, hệ dùng bao nhiêu, khoảng cho phép là gì. Nó không âm thầm sửa.
Số tham số của một mô hình không phải nhãn dán, nó là một phép nhân bạn tự làm được từ bảy trường trong config.json. Ba chỗ quyết định kết quả: dùng num_key_value_heads chứ không phải num_attention_heads cho chiếu khoá và chiếu giá trị; khối tiến có cổng là ba ma trận; và tie_word_embeddings xoá hẳn đầu ra ngôn ngữ, đúng vocab_size × hidden_size tham số, hơn nửa tỉ với từ vựng 128.256. Cộng xong thì phần thưởng là một bức tranh khác hẳn trực giác: khối tiến chiếm 70,20% mô hình, khối chú ý chỉ 16,71%, và toàn bộ chuẩn hoá chỉ 0,0033%. Cuối cùng, "khớp với con số công bố" chỉ có nghĩa tới độ chính xác mà nhà sản xuất viết ra, không phải tới từng đơn vị, và khi lệch nhiều hơn thế thì việc phải làm là đi tìm quy ước nào khác nhau, không phải bẻ số cho khớp.
- 1Một tệp config.json ghi hidden_size 4096, intermediate_size 14336, num_hidden_layers 32. Khối tiến kiểu SwiGLU của MỘT lớp có bao nhiêu tham số?
- 2Vẫn mô hình đó với vocab_size 128256, đang có tie_word_embeddings false và tổng 8.030.261.248 tham số. Đổi trường đó thành true thì chuyện gì xảy ra?
- 3Bạn muốn đoán xem khối chú ý hay khối tiến chiếm nhiều tham số hơn trong một lớp, mà chỉ được đọc config.json. Kiểm bằng cách nào?