nanochat을 TPU로 포팅하며: PyTorch에서 그대로 가져온 것과 깨진 것
Porting nanochat to a TPU: what carries over from PyTorch, and what breaks

Karpathy의 nanochat은 원래 8×H100 GPU 노드에서 실행되며, JAX 포팅도 이미 여러 개 있습니다. 이 글은 nanochat-jax 프로젝트의 TPU v6e-8 포팅 과정을 기록한 것으로, config과 아키텍처를 nanochat과 최대한 동일하게 유지하면서 모델 품질과 학습 성능을 따라잡는 것을 목표로 했습니다. 품질은 CORE 점수로 검증했는데, nanochat R4(d24)의 점수 분포(0.2512–0.2677)를 상회하는 0.2695를 달성했습니다. 하지만 성능은 아직 갭이 있어 MFU가 약 24%로 H100 대비 절반 수준입니다. 또한 토크나이저, base 모델, SFT 각 단계의 실행 결과와 TPU 관련 인사이트(메모리 대비 컴퓨트 비율, MXU 패딩 문제 등)를 공유합니다.
v6e의 MXU는 256×256으로 커졌는데, 텐서 차원이 256의 배수가 아니면 XLA가 0으로 패딩해서 일부 유닛이 낭비됩니다.