TL;DR

  • MaxText의 JAX/TPU 실행은 PyTorch 기준 Olmo 3 7B의 사전학습을 약 5.93조 토큰, 141만 스텝 규모로 재현하고, 1단계 종료 시 Ai2의 손실 곡선과 일치함.
  • PyTorch에서 JAX로 변환한 체크포인트는 KL 약 1.5e-3의 로짓 차이를 보이고, 8192토큰 문맥에서 최상위 토큰이 98.75% 일치함.
  • 홀드아웃 평가에서 데이터 로더 버그로 인한 가짜 성능 향상을 찾아냈으며, 체크포인트 재개 후 손실과 퍼플렉서티 차이가 0.000으로 유지됨.
  • 학습 도중 용량을 4분의 1로 줄이거나 TPU 세대를 바꿔도 학습 레시피를 바꾸지 않고 실행했으며, 디바이스당 처리량은 1% 이내로 유지되고 2단계 MFU는 57.4%임.
  • Ironwood에서 44.5% MFU를 달성했고, 어텐션 헤드 구성을 바꾸는 부가 실험은 품질을 유지하면서 12.4% 더 빠른 실행을 보임.

주요 결과

  • PyTorch에서 JAX로 모델을 변환함. Olmo 3의 재정렬 정규화 블록(reordered-norm block), QK 정규화(QK-norm), 3:1 슬라이딩/전역 어텐션 구조를 MaxText로 옮기고 로짓 동등성 검사(logit-parity check)로 검증함.
  • 변환된 0스텝 체크포인트와 HuggingFace 기준 모델의 쿨백-라이블러 발산(KL divergence)은 약 1.5e-3로, 같은 모델을 서로 다른 프레임워크에서 실행할 때의 잡음 수준임.
  • 전체 8192토큰 문맥을 bfloat16으로 처리할 때 두 모델의 최상위 토큰이 98.75% 일치함.
  • 실제 버그를 잡아내는 검증을 수행함. 홀드아웃 평가에서 MaxText가 기준 모델보다 앞서는 것처럼 보이게 한 데이터 로더 버그를 발견했으며, 그 성능 향상은 암기에서 비롯됨.
  • 여러 주에 걸친 학습의 신뢰성을 검증함. 체크포인트 저장 및 재개는 실행을 정확히 재현함.
  • 통제된 A/B 비교에서 재개 후 모든 스텝의 차이(Δ)는 0.000임.
  • 호스트 장애로 2단계 실행이 중단된 뒤 재개했을 때 127스텝을 다시 학습했으며, 기록된 손실과 퍼플렉서티 차이는 0.000임.
  • 실행 중 학습 작업의 크기를 조정함. 약 105만 스텝에서 처리 용량의 4분의 3을 잃은 뒤, 레시피를 바꾸지 않고 기존 크기의 4분의 1인 장비 구성에서 실행을 재개함.
  • 동일한 run_olmo3_7b_stage1.sh 스크립트가 디바이스당 배치 크기를 조정해 전역 배치 크기(GBS)를 유지함.
  • 양방향으로 측정한 디바이스당 처리량은 1% 이내로 유지됐으며, 강한 스케일링 효율은 약 100%임.
  • 레시피 실행 중 TPU 세대를 변경함. 2단계에서 동일한 실행기를 Ironwood 대신 v5p에 지정해 디바이스 유형만 변경했으며, 57.4% MFU를 유지함.
  • 성능 최적화로 연산 비용을 절감함. Ironwood의 7B 모델에서 SparseCore 집단 통신 오프로딩, 재계산(rematerialization) 조정, 최적 샤딩으로 44.5% MFU를 달성했으며, 전체 연산 예산의 약 3분의 1을 절감함.
  • TPU에 맞춘 공동 설계로 동일한 품질에서 속도를 높임. 파라미터 수와 부동소수점 연산량(FLOPs)을 그대로 유지하면서 어텐션 구성을 헤드 32개 × 헤드 차원 128에서 헤드 16개 × 헤드 차원 256으로 바꿈.
  • 헤드 차원 256이 Ironwood의 256×256 MXU를 완전히 활용해 실행 속도가 12.4% 빨라짐.
  • 손실 곡선은 1200억 토큰(3만 스텝)까지 원래 구성과 일치함. 이는 부가 실험이며, 재현 작업에는 원래 아키텍처를 유지함.
  • Ai2의 0스텝 PyTorch 가중치와 동일한 핵심 레시피에서 시작해 MaxText 실행 결과가 약 5.93조 토큰 / 141만 스텝의 전체 학습 예산에 걸쳐 Ai2가 공개한 손실 곡선을 따라가며, 1단계가 끝날 때 곡선이 일치함.
  • Ai2가 두 개를 이어 붙인 학습률 코사인 스케줄 대신 단일 코사인 학습률 스케줄을 사용하고, 공개된 데이터 혼합을 사용하는 등 레시피 세부사항 두 가지를 단순화했지만 일치는 유지됨.
  • 이어지는 내용에서는 각 요소의 구현과 측정 과정, 그리고 성능 향상을 거의 거짓으로 보이게 할 뻔한 한 사례를 다룸.

Olmo 3를 재현하는 이유

  • Olmo 3는 공개 가중치, 공개 데이터, 완전히 명시된 학습 레시피, Weights & Biases의 공개 기준 실행을 제공하는 드문 개방형 프런티어급 언어 모델임.
  • 독립적으로 학습한 실행 결과를 손실 곡선만이 아니라 홀드아웃 지표에서도 맞추는 일은 최적화기, 손실, 데이터 파이프라인, 수치 연산을 포함한 MaxText 스택이 충실하며 단지 학습되는 것처럼 보이는 데 그치지 않는다는 강한 증거임.
  • MaxText는 TPU용으로 구축된 JAX/XLA 대형 언어 모델(LLM) 학습 프레임워크임.
  • 검증 목표는 GPU 기반 PyTorch 레시피를 TPU 기반 JAX에서 비트 단위로 똑같이 맞추는 대신 중요한 지표에서 충실하게 재현할 수 있는지, 그리고 이를 어떻게 입증하는지임.
  • Olmo 3의 레시피는 일반 사전학습, 중간 학습(어닐링), 장문맥 적응의 3단계 커리큘럼임.
  • 이 글에서 다루는 범위는 약 5.9조 토큰 규모의 1단계 사전학습과 2단계 중간 학습이며, 둘 다 처음부터 끝까지 학습해 Ai2의 기준 결과와 비교함.
  • 3단계와 사후 학습(SFT/RL)은 Tunix를 이용한 레시피를 작성했지만 아직 실행하지 않은 단계임.