# Tái hiện quá trình huấn luyện OLMo 3 7B trên Google Cloud TPU bằng MaxText

- Nguồn: Google Developers Blog
- Thời gian phát hành: 2026-09-24 22:13 (giờ Việt Nam)
- Điểm AI: 48/100
- Link AIHOT.vn: https://aihot.vn/items/6bf64cbc9f7b10dd
- Nguồn dữ liệu AI HOT: https://aihot.news/items/cmufoas7707oyro8w46swq4rx
- Link gốc: https://developers.googleblog.com/reproducing-olmo-3-7b-pre-training-in-maxtext-case-study-of-large-scale-training-on-tpus

## Tóm tắt AI

Đội ngũ Google đã sử dụng MaxText để tái hiện thành công quá trình huấn luyện từ đầu của mô hình OLMo 3 7B trên Google Cloud TPU, bao gồm cả giai đoạn tiền huấn luyện và tinh chỉnh, với kết quả đạt chuẩn trên các bộ chỉ số đánh giá.

## Thân bài

![Reproducing OLMo 3 7B Pre-training in MaxText: case study of large scale training on TPUs](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/header.2e16d0ba.fill-1200x600_UqwAgBm.jpg)

24 THÁNG 9, 2026

[Ran Ran](https://developers.googleblog.com/search/?author=Ran+Ran) Kỹ sư phần mềm

[OLMo 3](https://docs.allenai.org/models/olmo), được phát triển bởi Viện Trí tuệ Nhân tạo Allen (AI2), là một mô hình ngôn ngữ hoàn toàn mở, hiện đại, được huấn luyện với kiến trúc tân tiến và quy trình huấn luyện đa giai đoạn. Để đánh giá khả năng của MaxText trên Google Cloud TPU, nhóm chúng tôi đã bắt tay vào việc tái tạo OLMo 3 7B của AI2 từ đầu. Chúng tôi chọn OLMo 3 vì nó kết hợp ba đặc tính hiếm khi xuất hiện cùng nhau. Đây là một mô hình 7B mạnh mẽ, hiện đại, được huấn luyện ở quy mô sản xuất thực tế. AI2 công khai gần như toàn bộ luồng mô hình, bao gồm dữ liệu, mã nguồn, cấu hình, điểm kiểm tra (checkpoints), nhật ký và các đánh giá. Cuối cùng, nó cung cấp cho chúng tôi một tham chiếu độc lập trên PyTorch và GPU để kiểm thử MaxText và TPU.

Chúng tôi đã tái tạo OLMo 3 7B của AI2 trong [MaxText](https://github.com/AI-Hypercomputer/maxtext) trên Google Cloud TPU, bao gồm cả giai đoạn 1 tiền huấn luyện và giai đoạn 2 tinh chỉnh (anneal), đồng thời chứng minh sự tương đồng trên các chỉ số đánh giá độc lập (held-out metrics), chứ không chỉ dựa trên đường cong mất mát (loss curve):

![table 0](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/table_0.original.png)

Các điểm nổi bật chính, mỗi điểm sẽ được trình bày chi tiết trong phần sau của bài viết:

- Chuyển đổi mô hình từ PyTorch sang JAX. Kiến trúc của OLMo 3 (khối reordered-norm, QK-norm, cơ chế chú ý 3:1 sliding/global) được chuyển sang MaxText và xác minh bằng kiểm tra logit-parity: điểm kiểm tra bước 0 sau khi chuyển đổi khớp với tham chiếu HuggingFace ở mức KL ≈ 1.5e-3, đây là ngưỡng nhiễu "cùng mô hình, khác framework", và ở ngữ cảnh đầy đủ 8192-token trong định dạng bfloat16, cả hai khớp với nhau ở token top-1 với tỷ lệ 98,75%.

- Xác minh giúp phát hiện lỗi thực tế. Các đánh giá độc lập đã phát hiện một lỗi trong bộ nạp dữ liệu (data-loader) khiến MaxText trông như đang vượt qua tham chiếu; thực tế mức tăng đó là do mô hình ghi nhớ dữ liệu.

- Độ tin cậy trong quá trình chạy kéo dài nhiều tuần. Tính năng checkpoint-and-resume tái lập quá trình chạy một cách chính xác: một thử nghiệm A/B có kiểm soát cho thấy Δ = 0.000 tại mọi bước sau khi khôi phục, và khi một lỗi máy chủ làm gián đoạn quá trình chạy giai đoạn 2, quá trình khôi phục đã huấn luyện lại 127 bước với Δ = 0.000 trong nhật ký loss và perplexity.

- Thay đổi quy mô công việc huấn luyện khi đang chạy. Tại bước ~1.05M, chúng tôi mất ba phần tư công suất; quá trình chạy được khôi phục trên một phân đoạn (slice) chỉ bằng một phần tư kích thước mà không cần thay đổi quy trình (cùng tập lệnh [run_olmo3_7b_stage1.sh](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/trainers/pre_train/scripts/olmo/run_olmo3_7b_stage1.sh) tự động điều chỉnh kích thước batch trên mỗi thiết bị để giữ GBS không đổi), thông lượng trên mỗi thiết bị được duy trì trong phạm vi 1% (≈100% strong scaling, đo lường theo cả hai hướng).

- Thay đổi thế hệ TPU giữa quy trình. Giai đoạn 2 trỏ cùng một trình khởi chạy vào v5p thay vì Ironwood, chỉ thay đổi loại thiết bị, và duy trì mức 57,4% MFU.

- Công việc tối ưu hiệu năng mang lại hiệu quả xứng đáng. Đạt 44,5% MFU trên Ironwood ở quy mô 7B thông qua việc giảm tải tập thể SparseCore, tinh chỉnh remat và phân đoạn (sharding) tối ưu, giúp tiết kiệm khoảng một phần ba ngân sách tính toán.

- Đồng thiết kế cho TPU: nhanh hơn với cùng chất lượng. Thay đổi hình dạng cơ chế chú ý từ 32 đầu × head-dim 128 thành 16 × 256, với cùng số lượng tham số và FLOPs, chạy nhanh hơn +12,4% vì head-dim 256 tận dụng tối đa MXU 256×256 của Ironwood, và đường cong loss của nó khớp với bản gốc qua 120B token (30k bước). Đây là một thử nghiệm phụ; bản tái tạo vẫn giữ nguyên kiến trúc gốc.

Bắt đầu từ trọng số PyTorch bước 0 của AI2 và cùng một quy trình cốt lõi, quá trình chạy MaxText bám sát đường cong loss đã công bố của AI2 trong suốt ngân sách ~5,93T-token / 1,41M-bước và đạt kết quả tương đương ở cuối giai đoạn 1. Chúng tôi thậm chí đã đơn giản hóa hai chi tiết quy trình (một lịch trình cosine LR duy nhất thay vì hai lịch trình ghép lại như AI2, và tập dữ liệu phát hành công khai; xem [quy trình](https://docs.google.com/document/d/1fHvtl172UgdWaEu4wwlVyGkaJIfr2VCCsFPKtABznHY/edit?tab=t.0#bookmark=id.bxrs32gtdeub) bên dưới), và sự tương đồng vẫn được duy trì. Phần còn lại của bài viết này là cách mỗi thành phần được xây dựng, đo lường và, trong một trường hợp mang tính giáo huấn, suýt chút nữa đã bị làm giả.

### Tại sao phải tái tạo OLMo 3?

[OLMo 3](https://allenai.org/olmo) là một trong số ít các mô hình ngôn ngữ tiên phong thực sự mở: trọng số mở, dữ liệu mở và quy trình huấn luyện được chỉ định đầy đủ với một lần chạy tham chiếu công khai trên Weights & Biases. Việc khớp với lần chạy được huấn luyện độc lập đó, dựa trên các chỉ số đánh giá độc lập thay vì chỉ đường cong loss, là bằng chứng mạnh mẽ cho thấy hệ thống MaxText (bộ tối ưu hóa, loss, đường ống dữ liệu, tính toán số học) là trung thực, chứ không chỉ là "trông có vẻ như đang huấn luyện".

MaxText là một framework huấn luyện LLM dựa trên JAX/XLA được xây dựng cho TPU. Câu hỏi chúng tôi đặt ra để trả lời: liệu một quy trình PyTorch-on-GPU có thể được tái tạo trung thực trong JAX-on-TPU, khớp với các chỉ số quan trọng thay vì khớp từng bit, và làm thế nào để chứng minh điều đó?

Quy trình của OLMo 3 là một chương trình 3 giai đoạn: tiền huấn luyện tổng quát, huấn luyện giữa kỳ (annealing) và thích ứng ngữ cảnh dài. Bài viết này bao gồm giai đoạn 1 (quá trình tiền huấn luyện ~5,9T-token) và giai đoạn 2 (huấn luyện giữa kỳ), cả hai đều được huấn luyện từ đầu đến cuối và đối chiếu với các tham chiếu của AI2. Giai đoạn 3 và hậu huấn luyện (SFT/RL qua Tunix) là các quy trình chúng tôi đã viết nhưng chưa chạy.

![image11](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/image11_ArbtJ6A.original.png)

Chương trình tiền huấn luyện OLMo-3: giai đoạn 1 (Ironwood) và giai đoạn 2 (v5p) được tái tạo trong bài viết này; giai đoạn 3 ngữ cảnh dài (seq 65k, YaRN) và hậu huấn luyện (SFT sau đó là GRPO qua Tunix) sẽ thực hiện tiếp theo.

### Quy trình thực hiện

OLMo 3 7B là mô hình transformer dày đặc 32 lớp, 4096-dim với một vài lựa chọn phi tiêu chuẩn: khối "reordered norm", QK-norm và hỗn hợp 3:1 giữa sliding-window và global attention. Cấu hình MaxText (olmo3-7b-pt.yml, được sử dụng cho giai đoạn 1 và 2) khớp chính xác với cấu hình này:

![table 1 (1)](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/table_1_1.original.png)

Quy trình huấn luyện phản ánh pretrain-1.py của OLMo-core; các thông số cần khớp để các đường cong trùng nhau:

![table 2 (1)](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/table_2_1.original.png)

Chúng tôi bắt đầu huấn luyện từ checkpoint PyTorch bước 0 của AI2, được chuyển đổi sang Orbax, vì vậy MaxText bắt đầu từ chính xác các trọng số giống như tham chiếu. Bản thân việc chuyển đổi là điểm kiểm tra đầu tiên: một lượt forward pass trên các trọng số đã chuyển đổi khớp với tham chiếu HuggingFace ở mức KL ≈ 1.5e-3 với độ trùng lặp 9/10 token top-10, ngưỡng nhiễu "cùng mô hình, khác framework".

Liệu sự tương đồng có phụ thuộc vào việc kế thừa khởi tạo của AI2 không? Rõ ràng là không. Như một kiểm tra độc lập, chúng tôi cũng đã chạy huấn luyện từ khởi tạo ngẫu nhiên của riêng MaxText trong khoảng 50k bước (3,5% lộ trình); loss huấn luyện của nó bám sát đường cong đã công bố của AI2, chạy thấp hơn một chút. Đó là một kiểm tra nhanh về loss huấn luyện, không phải là bản tái tạo đầy đủ, nhưng nó cho thấy sự tương đồng không phụ thuộc vào việc bắt đầu từ trọng số của AI2.

Đường ống dữ liệu phản ánh chính xác OLMo-core: token hóa và nối tất cả các tài liệu (có EOS ở giữa), cắt thành các thực thể 8192-token không chồng lấp, xáo trộn chỉ mục toàn cục với một hạt giống cố định và áp dụng bộ lọc lặp lại n-gram để loại bỏ các thực thể có >32 n-gram lặp lại. dataset_type=olmo_grain của MaxText (được xây dựng trên [Grain](https://github.com/google/grain)) thực hiện điều này.

Hai sự khác biệt có chủ ý so với lần chạy của AI2. (1) Lịch trình LR: AI2 ban đầu dự kiến ~5T token và mở rộng quá trình chạy giữa chừng lên tổng cộng ~5,93T, vì vậy dấu vết LR của nó ghép hai đường cong cosine (có thể thấy trong lần chạy WandB công khai); chúng tôi đã chạy một đường cong cosine duy nhất trên toàn bộ lộ trình. (2) Dữ liệu: chúng tôi huấn luyện trên tập dữ liệu OLMo-3 được phát hành công khai, lược bỏ một phần nhỏ (<0,5% ngân sách token, chủ yếu là các shard s2pdf không có trong danh sách tệp được phát hành) mà lần chạy nội bộ của AI2 đã thấy. Cả hai đều là những đơn giản hóa mà chúng tôi chọn, không phải ngẫu nhiên, và MaxText vẫn khớp trên mọi chỉ số đánh giá độc lập. Đây cũng là lý do tại sao chúng tôi nói "tái tạo trong phạm vi nhiễu giữa các lần chạy", chứ không phải khớp từng bit (xem phân tích KL trong [§G](https://docs.google.com/document/d/10PwEU4OfJWksBoloAQgZr-4HyJ7HyVsfF-R98zYBkK4/edit#bookmark=id.pf43by2wull6)).

### Những gì chúng tôi phải xây dựng

OLMo 3 chưa có trong MaxText khi chúng tôi bắt đầu; quá trình tái tạo đã bổ sung và đưa mọi thứ bên dưới vào upstream. "Tự tái tạo" ở cuối bài viết này là cấu hình, không phải mã nguồn.

- Bản thân mô hình: khối reordered-norm, QK-norm và mô hình chú ý sliding/global 3:1 ([#3004](https://github.com/AI-Hypercomputer/maxtext/pull/3004), [#3112](https://github.com/AI-Hypercomputer/maxtext/pull/3112)).
- Bộ tối ưu hóa skip-step khớp với ngữ nghĩa của OLMo-core đến tận độ lệch chuẩn chạy (running std) có hiệu chỉnh Bessel (bỏ qua ở mức 6σ trên cửa sổ 128 bước) ([#3490](https://github.com/AI-Hypercomputer/maxtext/pull/3490)).
- z-loss ([#3211](https://github.com/AI-Hypercomputer/maxtext/pull/3211)) và tính năng che weight-decay theo từng tham số để có thể loại trừ các embedding ([#3280](https://github.com/AI-Hypercomputer/maxtext/pull/3280)).
- Pipeline dữ liệu olmo_grain ([#3749](https://github.com/AI-Hypercomputer/maxtext/pull/3749)): đọc truy cập ngẫu nhiên các shard đã được token hóa, xáo trộn chỉ mục toàn cục có hạt giống (seeded) với cơ chế bảo vệ bằng dấu vân tay (fingerprint) chống lại việc hoán đổi dữ liệu âm thầm khi khởi động lại, bộ lọc lặp n-gram và (từ giai đoạn 2) checkpoint trạng thái trình lặp Grain.
- Chuyển đổi checkpoint HF↔Orbax với kiểm tra logit-parity, công cụ đứng sau mọi con số "nền nhiễu khung làm việc" (framework noise floor) trong bài viết này ([#3112](https://github.com/AI-Hypercomputer/maxtext/pull/3112), [#3832](https://github.com/AI-Hypercomputer/maxtext/pull/3832)).
- Trình khởi chạy giai đoạn 1 và 2 (script chạy dựa trên env + wrapper XPK với các lệnh submit / monitor / resume_until_done) ([#3886](https://github.com/AI-Hypercomputer/maxtext/pull/3886)).
- Các thẻ TensorBoard parity (optim/step_skipped, perf/total_tokens) để mọi chỉ số trên bảng điều khiển W&B của AI2 đều có đối tác tương ứng trong MaxText để so sánh.

### Nó có hội tụ không?

Tiêu đề là một lớp phủ duy nhất: lm_loss giai đoạn 1 của MaxText so với đường cong WandB đã công bố của AI2, được căn chỉnh theo bước và nhóm thành các trung bình 2k bước. Qua khoảng 800k bước, hai đường này bám sát nhau trong phạm vi ±0.012; từ khoảng 0.9M, MaxText thấp hơn AI2 và không bao giờ cắt ngang trở lại, dấu hiệu đầu tiên của lỗi dữ liệu được phân tích trong phần tiếp theo.

![image5 (1)](https://storage.googleapis.com/gweb-developer-goog-blog-assets/images/image5_1.original.png)

So sánh loss giai đoạn 1 của MaxText và AI2 qua 1.41M bước, với khoảng cách được làm mịn 18k bước ở bảng dưới. Trên: các đường cong không thể phân biệt được ở quy mô này cho đến phần đuôi. Dưới: khoảng cách duy trì trong phạm vi ±0.012 đến khoảng 800k, nghiêng về giá trị âm từ khoảng 0.9M khi các lỗi dữ liệu lặp lại tích tụ, và giảm sâu qua 1.25M, đạt −0.22 trong các bin 2k bước thô trước khi áp dụng làm mịn 18k bước cho cả hai bảng. Không có điều nào trong số này ảnh hưởng đến loss trên tập held-out hoặc độ chính xác hạ nguồn, như các phần tiếp theo sẽ cho thấy. (Bảng theo cột mốc: Phụ lục A; các đường cong đầy đủ được commit tại olmo_stage1_loss_curve.tsv.)

Nhưng chỉ riêng đường cong loss là bằng chứng yếu: hai lần chạy có thể khớp nhau về loss huấn luyện nhưng lại khác biệt ở mọi thứ bạn thực sự quan tâm. Vì vậy, chúng tôi đã xác minh sự hội tụ trên bốn bề mặt độc lập tại sáu cột mốc bước trải dài 915k bước:

1. Held-out C4 lm_loss: đánh giá chỉ chuyển tiếp (forward-only) trên 16M token của C4-en held-out, các batch giống hệt nhau cho cả hai lần chạy.
2. Bộ 8 tác vụ lm-eval-harness: MMLU, HellaSwag, ARC-easy/challenge, OpenBookQA, PIQA, BoolQ, WinoGrande.
3. Độ phức tạp (perplexity) held-out đa miền: một đợt quét theo phong cách Paloma trên các miền web, tin tức, bách khoa toàn thư và hỗn hợp.
4. Token-level KL: khoảng cách phân phối token tiếp theo trên các đầu vào giống hệt nhau.

Cách chúng tôi đo lường. Tất cả các đánh giá chạy trên các cặp checkpoint được căn chỉnh theo bước: lần chạy MaxText trực tiếp so với checkpoint AI2 tại cùng một bước, tức là bản sửa đổi HuggingFace công khai allenai/Olmo-3-1025-7B@stage1-step{N} được chuyển đổi sang Orbax (quá trình chuyển đổi tái tạo tham chiếu PyTorch với KL ≤ 1.8e-3, nền nhiễu khung làm việc). lm-eval sử dụng lm-eval-harness tiêu chuẩn (5-shot MMLU, các phần khác để mặc định); σ là stderr của harness theo từng tác vụ, được kết hợp theo phương pháp cầu phương (quadrature) cho các delta.

_Bài gốc còn tiếp._ Xem tiếp tại: <https://developers.googleblog.com/reproducing-olmo-3-7b-pre-training-in-maxtext-case-study-of-large-scale-training-on-tpus>
