AI-Hypercomputer/maxtext · 上手攻略

  • 仓库:AI-Hypercomputer/maxtext
  • 链接:https://github.com/AI-Hypercomputer/maxtext
  • 分类:LLM 训练框架 / JAX
  • 作者:Tom
  • 更新:2026-08-15

是什么

MaxText 是 Google AI-Hypercomputer 团队开源的高性能、高可扩展 LLM 训练与后训练库,全部用纯 Python/JAX 编写,目标是 TPU 和 NVIDIA GPU 上的高效分布式训练与微调。官方于 2025-09-15 发布 PyPI 包,当前稳定版本 0.2.3(2026 年初)。

MaxText 的定位是"LLM 项目发射台":既是从零预训练的参考实现,也支持 SFT、GRPO(类 RLHF)、GSPO 等后训练技术,并内置 Gemma、Llama、DeepSeek、Qwen、Mistral 等主流模型库。截至 2026-08,它已支持 DeepSeek V4 Flash (284B)、Qwen3.5 MoE (397B)、Gemma 4 (31B dense / 26B MoE)、Kimi K2 等最新模型。


解决什么问题

LLM 训练框架通常在易用性和规模扩展性之间二选一:简单框架难以 scale 到上万 TPU/GPU,定制框架则门槛极高。MaxText 通过 JAX + XLA 编译器实现了"优化极少但自动极优"——开发者写纯 Python/JAX,XLA 自动完成算子融合、内存布局、分布式并行,最大化 MFU(Model FLOPs Utilization)。

痛点 MaxText 对策
分布式训练配置复杂 YAML 配置 + 自动多主机扩展,无需手写分割逻辑
不同后端(TPU/GPU)代码分化 统一 Python/JAX API,TPU 和 GPU 共用同一套代码
后训练(RL/SFT)框架分散 内置 Tunix 后训练框架,统一管理 SFT/GRPO/GSPO
换模型要从头配 60+ 模型配置(configs/models/),checkpoint 转换工具齐全
评测体系混乱 内置 eval 框架,支持 lm-eval / evalchemy / 自定义 benchmark

快速安装

环境要求

  • Python 3.12(主推版本,其他版本可能出现兼容问题)
  • Linux(仅在 Linux 上做正式测试)
  • TPU VM 或 NVIDIA GPU VM
  • Cloud Storage Bucket(存 checkpoint 和日志)

方式一:PyPI 安装(推荐)

# 创建虚拟环境(使用 uv)
uv venv --python 3.12 --seed maxtext-env
source maxtext-env/bin/activate

# TPU 预训练/推理
uv pip install maxtext[tpu]==0.2.3 --resolution=lowest
install_tpu_pre_train_extra_deps

# GPU 预训练/推理(CUDA 12)
uv pip install maxtext[cuda12]==0.2.3 --resolution=lowest
install_cuda12_pre_train_extra_deps

# TPU 后训练(SFT / RL / vLLM 推理)
UV_TORCH_BACKEND=cpu uv pip install maxtext[tpu-post-train]==0.2.3 --resolution=lowest
install_tpu_post_train_extra_deps

⚠️ --resolution=lowest 是官方强制的:它安装 MaxText 验证过的特定版本依赖,而非最新版,以保证结果可复现。切换安装目标(如从 tpu 切到 tpu-post-train)建议重建虚拟环境以避免依赖冲突。

方式二:源码安装

git clone https://github.com/AI-Hypercomputer/maxtext.git
cd maxtext
uv venv --python 3.12 --seed maxtext-env
source maxtext-env/bin/activate
uv pip install -e .[tpu] --resolution=lowest
install_tpu_pre_train_extra_deps

验证安装

uv pip check
python3 -c "import maxtext"
python3 -m maxtext.trainers.pre_train.train --help

核心用法

1. 预训练(单 host TPU/GPU)

# 必须先创建 Cloud Storage bucket 并配置 base_output_directory
python3 -m maxtext.trainers.pre_train.train \
  config=configs/base.yml \
  base_output_directory=gs://YOUR_BUCKET/maxtext/ \
  steps=100 \
  log_period=10

configs/base.yml 包含 ~1B 参数的解码器-only 模型配置,所有参数可通过命令行覆盖:

python3 -m maxtext.trainers.pre_train.train \
  config=configs/base.yml \
  base_output_directory=gs://YOUR_BUCKET/maxtext/ \
  steps=1000 \
  log_period=100 \
  per_device_batch_size=8

2. HuggingFace Checkpoint 转换

MaxText 使用 Orbax 格式,不直接读 HF 格式,需要先转换:

# 参见 https://maxtext.readthedocs.io/en/latest/guides/checkpointing_solutions/convert_checkpoint.html

3. 运行推理(vLLM 加速)

# 需要 tpu-post-train 安装或 maxtext[runner]
# 参考:https://maxtext.readthedocs.io/en/latest/tutorials/inference.html

4. 后训练(SFT / RL)

# SFT 单 host TPU
# 参考:https://maxtext.readthedocs.io/en/latest/tutorials/posttraining/sft.html

# RL(GRPO/GSPO)单 host
# 参考:https://maxtext.readthedocs.io/en/latest/tutorials/posttraining/rl.html

# 多 host RL
# 参考:https://maxtext.readthedocs.io/en/latest/tutorials/posttraining/rl_on_multi_host.html

5. 评测

# 内置 eval 框架,支持 lm-eval / evalchemy / 自定义 benchmark
# 参考:https://maxtext.readthedocs.io/en/latest/guides/eval_framework.html

6. 多主机扩展(Kubernetes / GKE)

# 推荐通过 GKE 运行多主机
# 参考:https://maxtext.readthedocs.io/en/latest/run_maxtext/run_maxtext_via_cluster_toolkit.html
# 或 XPK:https://maxtext.readthedocs.io/en/latest/run_maxtext/run_maxtext_via_xpk.html

支持的主要模型(2026 年 8 月)

系列 模型
Google Gemma Gemma 4 (26B MoE, 31B dense), Gemma 3 (4B/12B/27B), Gemma 2 (2B/9B/27B)
DeepSeek AI DeepSeek V4 Flash (284B), DeepSeek V3.2 (671B), DeepSeek R1-0528
Qwen(阿里) Qwen3.5 MoE (35B/397B), Qwen3 Next (80B), Qwen3 (30B~480B)
Moonshot AI Kimi K2 / K2-Thinking / K2.5 / K2.6
Meta Llama Llama 4 Scout (109B) / Maverick (400B), Llama 3.3 (70B)
Mistral AI Mixtral (8x7B / 8x22B), Mistral (7B)
GPT-OSS GPT-OSS (20B, 120B)

典型适用场景

  1. 学术研究:用 MaxText 作为从零训练新架构的参考实现(改动配置而非重写框架)
  2. 后训练微调:用 Tunix 框架对开源模型做 SFT/GRPO/GSPO,适配私有数据或特殊任务
  3. 大规模预训练:利用 JAX + XLA 的 scale 特性,在 TPU pod 或多 GPU 集群上预训练数百 B 参数模型
  4. 评测基准:用统一 eval 框架对比不同模型的性能,支持 lm-eval / evalchemy
  5. Checkpoint 实验:用 MaxText 的 Orbax 格式管理实验,切换不同模型配置做 ablations

坑与注意

  1. Python 3.12 强制:其他版本可能不兼容,建议严格使用 3.12
  2. 依赖冲突高风险:MaxText 依赖链深,--resolution=lowest + 全新虚拟环境是官方唯一保证可跑的方式;切换安装选项必须重建环境
  3. HF Checkpoint 必须转换:直接喂 HuggingFace 格式会报错,需走 checkpoint 转换流程
  4. TPU 和 GPU 安装选项不同[tpu] vs [cuda12],选错会导致运行时缺少 XLA 编译路径
  5. 多主机门槛高:虽然文档覆盖 GKE / XPK,但实际多主机部署需要 Google Cloud 环境和 Cluster Toolkit 配置,门槛不低
  6. Linux only:Windows / macOS 无法正常运行(只能在 Linux VM 或 Colab TPU 上跑)
  7. PyPI 版本落后 main 分支:PyPI 0.2.3 是稳定版,main 分支可能包含更新的模型支持(如有需要从源码装)

与同类对比

维度 MaxText Megatron-LM transformers DeepSpeed
语言 Python/JAX Python/PyTorch Python/PyTorch Python/PyTorch
硬件 TPU + GPU GPU 为主 任意 GPU 为主
自动并行 ✅ XLA ❌ 手动 部分
后训练 ✅ SFT/GRPO/GSPO
模型库 60+ 内置 有限 HuggingFace 有限
多模态 Gemma 3/4 VL
上手难度 中(YAML 配置) 高(分布式配置) 低(HuggingFace)

一句话推荐结论

MaxText 是目前最成熟的 JAX 原生 LLM 训练库——如果你在 TPU 生态或需要极致 scale-out 性能,它是首选;如果你在纯 GPU 生态且追求快速出活,PyTorch 生态(HuggingFace + DeepSpeed)更友好。核心价值在于用纯 Python 写出规模无关的代码,XLA 自动处理所有并行复杂度

📖 文档:https://maxtext.readthedocs.io/ 🧵 Discord:https://discord.com/invite/2H9PhvTcDU ⚠️ 官方建议生产环境用 PyPI latest stable 版本,main 分支不保证 production-ready