基于MaxText在TPU上复现OLMo 3 7B预训练:大规模训练案例研究
MaxText团队利用JAX/XLA在谷歌云TPU上从零成功复现了AI2的OLMo 3 7B大模型,在预训练及中期训练阶段的保留评估指标均精准对齐原PyTorch-GPU基准。该方案实现了高达57.4%的模型算力利用率(MFU),并在中途集群扩缩容及跨代TPU迁移中展现出极高鲁棒性。此外,团队通过保留验证发现并修复了导致训练损失虚低的数据加载器隐性记忆漏洞,验证了完备评估机制的重要性。
MaxText团队利用JAX/XLA在谷歌云TPU上从零成功复现了AI2的OLMo 3 7B大模型,在预训练及中期训练阶段的保留评估指标均精准对齐原PyTorch-GPU基准。该方案实现了高达57.4%的模型算力利用率(MFU),并在中途集群扩缩容及跨代TPU迁移中展现出极高鲁棒性。此外,团队通过保留验证发现并修复了导致训练损失虚低的数据加载器隐性记忆漏洞,验证了完备评估机制的重要性。
Google Cloud 已将 TPU 原生支持集成至开源推理框架 vLLM,开发者可通过 GKE 弹性扩展嵌入模型流水线。针对 Qwen3-Embedding-8B 等模型的 15K+ 超长上下文推理,工程团队引入了硬件安全张量对齐、JAX/XLA 编译预热及分块预填充的混合 StepPool 架构等专项优化,在实现与 GPU 基线近乎一致的数值精度的同时提升吞吐,相关配置方案已在 GitHub 开源。
autofinetune 项目推出了一套全自动大模型后训练研究闭环,涵盖监督微调(SFT)和基于 GRPO 的强化学习。开发者只需在单个 Markdown 文档中设定边界条件与评估指标,AI Agent 便能自主修改训练脚本、发起实验,并将验证有效的超参数优化自动提交至 Git。该项目基于 Google 的 Tunix、Gemma 和 Cloud TPU 技术栈构建,消除了繁琐的人工微调过程,已在函数调用与数学推理任务中验证了无人工干预的性能提升。