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를 대상으로 합니다. 완전한 학습 스택을 위해 JAX AI 라이브러리 제품군을 통합합니다: 신경망을 위한 Flax, 사후 학습을 위한 Tunix, 체크포인팅을 위한 Orbax, 최적화를 위한 Optax, 그리고 데이터 로딩을 위한 Grain입니다. 처음부터 시작하는 사전 학습(pre-training)과 Supervised Fine-Tuning(SFT) 및 Reinforcement Learning(GRPO 및 GSPO)과 같은 기술을 사용한 확장 가능한 사후 학습(post-training)을 모두 지원합니다.
Who it’s for
실험, 아이디어 구상 및 대규모 모델 학습을 위한 고성능 시작점을 필요로 하는 야심 찬 LLM 프로젝트를 구축하는 AI 연구원 및 프로덕션 엔지니어를 위해 설계되었습니다.
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) 의존성 없이 실행할 수 있습니다.