AI-Hypercomputer/maxtext

A simple, performant, and scalable Jax LLM!

何を解決するか

MaxTextは、大規模言語モデル(LLM)の学習に向けた高パフォーマンスでスケーラブルな参照実装を提供することを目的としています。単一ホストからGoogle Cloud TPUsおよびGPUの巨大クラスタに至る広範なハードウェア環境において、高いモデルFLOPs利用率(MFU)とスループット(トークン/秒)を達成する課題に対処しつつ、コードベースをシンプルかつ「最適化フリー」に保っています。

仕組み

純粋なPythonとJAXで記述されており、XLAコンパイラを活用してパフォーマンスを最適化しています。Flax(ニューラルネットワーク)、Tunix(事後学習)、Orbax(チェックポイント)、Optax(最適化)、Grain(データロード)といったJAX AIライブラリのセットと統合されています。事前学習からスケーラブルな事後学習(Supervised Fine-Tuning(SFT)や強化学習(GRPOおよびGSPO)を用いた手法)までをサポートしています。

対象ユーザー

研究および生産環境で野心的なLLMプロジェクトを構築する研究者や開発者向けです。特に、人気のあるオープンソースモデルの学習およびファインチューニングにスケーラブルで高パフォーマンスなフレームワークが必要な方々に適しています。

特徴

  • 広範なモデル対応: Gemma、Llama、DeepSeek、Qwen、Mistralの参照実装を含み、MoE(エキスパートの混合)およびマルチモーダルモデルもサポートしています。
  • スケーラビリティ: 最大数万チップでの事前学習をサポートしています。
  • 事後学習フレームワーク: Tunixを介してSFTおよびRL(vLLMによるサンプリングを使用)のスケーラブルなフレームワークを提供しています。
  • ハードウェア最適化: 最大効率を発揮するよう、Google Cloud TPUsおよびGPUに特化して最適化されています。
  • マルチモーダル対応: Gemma 3、Gemma 4、Llama 4 VLMのマルチモーダル学習をサポートしています。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト