AI-Hypercomputer/maxtext
A simple, performant, and scalable Jax LLM!
What it solves
MaxText 提供了一个高性能、可扩展的参考实现,用于训练大语言模型 (LLM)。它通过利用 JAX 和 XLA 编译器来实现高模型 FLOPs 利用率 (MFU) 和吞吐量,从而消除了对复杂手动优化的需求,涵盖了从单机到大规模集群(多达数万个芯片)的场景。
How it works
MaxText 使用纯 Python 和 JAX 编写,针对 Google Cloud TPUs 和 GPUs 进行优化。它集成了一套 JAX AI 库以提供完整的训练栈:用于神经网络的 Flax,用于训练后处理的 Tunix,用于检查点的 Orbax,用于优化的 Optax,以及用于数据加载的 Grain。它支持从零开始的预训练以及使用监督微调 (SFT) 和强化学习 (GRPO 和 GSPO) 等技术进行的可扩展训练后处理。
Who it’s for
它专为 AI 研究人员和生产工程师设计,旨在为那些正在构建宏大 LLM 项目、需要高性能实验、构思和大规模模型训练起点的用户提供支持。
Highlights
- Broad Model Support: 包括 Gemma, Llama, DeepSeek, Qwen, Mistral, 和 GPT-OSS 的参考实现。
- Scalable Training: 支持在大规模集群上进行预训练,并通过 Tunix 进行可扩展的训练后处理。
- Multi-modal Capability: 支持 Gemma 3, Gemma 4, 和 Llama 4 VLMs 等多模态模型的训练。
- Optimization-Free: 通过 JAX 和 XLA 自动实现高效率和高吞吐量。
- Decoupled Mode: 可以在没有 Google Cloud Platform (GCP) 依赖的情况下运行。