El desarrollador tucan9389 ha portado nanochat, el proyecto de entrenamiento de modelos de chat de Andrej Karpathy pensado para nodos de 8 GPU H100, a TPU usando JAX, Flax y Pallas. El objetivo del puerto, bautizado nanochat-jax, es conservar al máximo la configuración y arquitectura originales y verificar que los resultados de calidad se reproducen en hardware distinto.
La verificación se ejecutó en un slice TPU v6e-8 (ocho chips Trillium) en Google Cloud. El script speedrun.sh reproduce las fases de tokenizador (5,3 minutos), modelo base (6,02 horas más 44,5 minutos de evaluación) y ajuste supervisado o SFT (68,7 minutos y unas 3,5 horas de evaluación), pero detiene el proceso antes del aprendizaje por refuerzo. En la métrica CORE de 22 tareas, el modelo base alcanza 0,2695, por encima de GPT-2 (0,2565) y dentro de la banda R4 de Karpathy (0,2512–0,2677). El val bpb se queda en 0,7343 frente a 0,7185 de la referencia, diferencia atribuible al tokenizador.
El rendimiento aún muestra una brecha notable: el MFU ronda el 24 % en d24, aproximadamente la mitad del 47–48 % que Karpathy obtuvo con H100 en d20. El entrenamiento costó 60,8 dólares en tarifa spot durante 12,19 horas (unos 263 dólares a tarifa on-demand) e incluyó una preempción con recuperación. El artículo detalla particularidades del hardware TPU v6e, como su MXU de 256×256, sus 32 GB de HBM por chip y su carácter compute-heavy, y explica los ajustes necesarios —atención Splash, gradientes fused, lm-head en precisión alta— para que el código encaje en esta arquitectura. También alerta sobre el coste de olvidar un slice encendido: unos 120 dólares al día.
