AI-Hypercomputer/maxtext

A simple, performant, and scalable Jax LLM!

What it solves

MaxTextは、大規模言語モデル(LLM)のトレーニングのための、高性能でスケーラブルなリファレンス実装を提供します。JAXとXLAコンパイラを活用することで、単一ホストから大規模なクラスター(最大数万個のチップ)にわたって、高いモデルFLOPs利用率(MFU)とスループットを実現し、複雑な手動の最適化の必要性を排除します。

How it works

純粋なPythonとJAXで記述されたMaxTextは、Google CloudのTPUおよびGPUをターゲットとしています。Flax、Tunix、Orbax、Optax、GrainといったJAX AIライブラリのスイートを統合し、完全なトレーニングスタックを提供します:Flaxはニューラルネットワーク用、Tunixはポストトレーニング用、Orbaxはチェックポインティング用、Optaxは最適化用、Grainはデータローディング用です。スクラッチからの事前学習(pre-training)と、Supervised Fine-Tuning (SFT) や Reinforcement Learning (GRPO and GSPO) といった手法を用いたスケーラブルな事後学習(post-training)の両方をサポートしています。

Who it’s for

実験、アイデア出し、および大規模なモデルトレーニングのための高性能な出発点が必要な、野心的なLLMプロジェクトを構築しているAI研究者やプロダクションエンジニア向けに設計されています。

Highlights

  • Broad Model Support: Gemma, Llama, DeepSeek, Qwen, Mistral, and GPT-OSSのリファレンス実装を含みます。
  • Scalable Training: 大規模なクラスターでの事前学習と、Tunixを介したスケーラブルな事後学習をサポートします。
  • Multi-modal Capability: Gemma 3, Gemma 4, and Llama 4 VLMsのようなマルチモーダルモデルのトレーニングをサポートします。
  • Optimization-Free: JAXとXLAを通じて、自動的に高い効率とスループットを実現します。
  • Decoupled Mode: Google Cloud Platform (GCP) への依存なしに実行可能です。