google/paxml · 上手攻略
- 仓库:google/paxml
- 链接:https://github.com/google/paxml
- 分类:ML Framework · JAX · TPU Training
- 作者:Tom
- 更新:2026-09-17
是什么
Pax(Parallel AXis)是 Google 基于 JAX 构建的大规模神经网络训练框架,专门针对需要跨多个 TPU 加速器切片(slice)或 Pod 扩展的超大规模模型设计。Pax 的核心理念是最大化 ML 研究与部署的效率,尤其在超大规模分布式训练场景下表现突出,官方称其达到了业界领先的 Model FLOPs Utilization(MFU) 利用率。
Pax 并非从零造轮子,而是站在 JAX 生态的肩膀上:
| 依赖层 | 组件 | 作用 |
|---|---|---|
| 底层引擎 | JAX + XLA | 自动微分 + 硬件加速编译 |
| 网络结构 | Flax | 高性能神经网络层 API |
| 数据输入 | SeqIO | 序列数据预处理 |
| 优化器 | Optax | 梯度优化工具集 |
| 检查点 | Orbax + TensorStore | 大型多维数组的读写 |
| 配置系统 | Fiddle | ML 友好的配置管理 |
| AutoML | PyGlove | 超参搜索与自动调优 |
| 层实现 | Praxis | Pax 专属的层库(内部实现) |
简单说:Flax = 层 API,Pax = 训练框架 + 分布式策略 + 配套工具。Pax 用 Flax 定义模型结构,但在其上封装了完整的分布式训练、配置管理、检查点、实验管理流水线。
解决什么问题
训练超大模型(十亿到万亿参数)是工程难题:数据并行、模型并行、流水线并行的配置复杂;跨多 TPU Pod 的通信和调度需要专门代码;实验可复现性难以保证。Pax 解决了这些问题:
- 开箱即用的分布式策略:支持
pmap(单主机多设备)和pjit/SPMD(多主机跨 Pod 横向扩展),无需手写collective原语 - 配置驱动的实验管理:通过 Fiddle 配置类(而非硬编码超参)管理实验变体,配合 XManager 做实验编排
- MFU 驱动的性能基准:提供 TPU v4 Pod 弱扩展(weak scaling)基准,展示 1B → 16B → GPT3-XL 在不同规模下的训练效率
- 多切片(Multi-slice)支持:通过 Google Cloud Queued Resources API 调度跨 Pod 切片训练(v4-128 及以上规模)
快速安装
在 Cloud TPU VM 上安装(推荐)
# 1. 创建 TPU VM(单主机 v4-8 示例)
export ZONE=us-central2-b
export VERSION=tpu-vm-v4-base
export PROJECT=<your-project>
export ACCELERATOR=v4-8
export TPU_NAME=paxml
gcloud compute tpus tpu-vm create $TPU_NAME \
--zone=$ZONE --version=$VERSION \
--project=$PROJECT --accelerator-type=$ACCELERATOR
# 2. SSH 登录 TPU VM
gcloud compute tpus tpu-vm ssh $TPU_NAME --zone=$ZONE
# 3. 安装 Stable 版本(PyPI)
python3 -m pip install -U pip
python3 -m pip install paxml jax[tpu] \
-f https://storage.googleapis.com/jax-releases/libtpu_releases.html
# 4. 安装 Dev 版本(需要先装 praxis)
git clone https://github.com/google/praxis
pip install -e praxis
git clone https://github.com/google/paxml
pip install -e paxml
pip install "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
⚠️ 依赖问题注意:如果遇到 transitive dependency 冲突,请使用 release branch 的 requirements.txt 而非直接 pip install,例如:
git clone -b r0.4.0 https://github.com/google/paxml
pip install --no-deps -r paxml/paxml/pip_package/requirements.txt
⚠️ Python 版本:官方示例路径中写的是 python3.8,建议使用 Python 3.8–3.11,具体以对应 release 的 requirements.txt 为准。
核心用法
运行一个预置实验(单主机 pmap)
python3 .local/lib/python3.8/site-packages/paxml/main.py \
--exp=tasks.lm.params.lm_cloud.LmCloudTransformerAdamLimitSteps \
--job_log_dir=gs://<your-bucket> \
--pmap_use_tensorstore=True
运行 SPMD 多主机实验
python3 .local/lib/python3.8/site-packages/paxml/main.py \
--exp=tasks.lm.params.lm_cloud.LmCloudSpmd2BLimitSteps \
--job_log_dir=gs://<your-bucket>
C4 数据集上的基准模型
| 模型规模 | TPU 配置 | 运行命令 |
|---|---|---|
| 1B 参数 | v4-8(4 Replicas) | --exp=tasks.lm.params.c4.C4Spmd1BAdam4Replicas |
| 16B 参数 | v4-64(32 Replicas) | --exp=tasks.lm.params.c4.C4Spmd16BAdam32Replicas |
| GPT3-XL | v4-128(Pipeline) | --exp=tasks.lm.params.c4.C4SpmdPipelineGpt3SmallAdam64Replicas |
Multi-slice 跨 Pod 训练(2 slices × v4-128)
# Step 1: 创建 Queued Resource(多切片)
export TPU_PREFIX=<your-prefix>
export QR_ID=$TPU_PREFIX
export ACCELERATOR=v4-128
export NODE_COUNT=2
gcloud alpha compute tpus queued-resources create $QR_ID \
--accelerator-type=$ACCELERATOR \
--runtime-version=tpu-vm-v4-base \
--node-count=$NODE_COUNT \
--node-prefix=$TPU_PREFIX
# Step 2: 在所有 worker 上安装依赖
for ((i=0; i<$NODE_COUNT; i++)); do
gcloud compute tpus tpu-vm ssh $TPU_PREFIX-$i \
--zone=us-central2-b --worker=all \
--command="pip install paxml && pip install orbax==0.1.1 && pip install 'jax[tpu]' -f https://storage.googleapis.com/jax-releases/libtpu_releases.html"
done
# Step 3: 终端 0 — 运行 slice 0
export TPU_PREFIX=<your-prefix>
export EXP_NAME=C4Spmd22BAdam2xv4_128
export LIBTPU_INIT_ARGS="--xla_jf_spmd_threshold_for_windowed_einsum_mib=0 --xla_tpu_spmd_threshold_for_allgather_cse=10000 --xla_enable_async_all_gather=true --xla_tpu_enable_latency_hiding_scheduler=true TPU_MEGACORE=MEGACORE_DENSE"
gcloud compute tpus tpu-vm ssh $TPU_PREFIX-0 --zone=us-central2-b --worker=all \
--command="LIBTPU_INIT_ARGS=$LIBTPU_INIT_ARGS python3 /home/yooh/.local/lib/python3.8/site-packages/paxml/main.py \
--exp=tasks.lm.params.c4_multislice.${EXP_NAME} --job_log_dir=gs://<your-bucket>"
# Step 4: 终端 1 — 并发运行 slice 1(命令同上,$TPU_PREFIX-1)
Jupyter Notebook 教程
# 在 TPU VM 上启用 port forwarding
gcloud compute tpus tpu-vm ssh $TPU_NAME --project=$PROJECT_NAME \
--zone=$ZONE --ssh-flag="-4 -L 8080:localhost:8080"
# TPU VM 内安装 jupyter
pip install notebook
pip install markupsafe==2.0.1
export PATH=/home/$USER/.local/bin:$PATH
# 启动 notebook(记录生成的 token)
jupyter notebook --no-browser --port=8080
# 本地浏览器访问 http://localhost:8080/ 输入 token
⚠️ 注意:每次运行新 notebook 前记得 pkill -9 python3 释放 TPU 资源。
典型适用场景
- 超大规模语言模型预训练:540B PaLM 模型背后训练基础设施的同类架构,适合 1B–540B 参数量级的 Transformer 训练研究
- 多模态大模型研究:官方 repo 包含 vision、speech、text 多种模态的配置示例
- TPU Pod 级分布式训练:需要横向扩展到数百乃至数千 TPU 核的场景(multi-slice),而非单机多卡
- MFU 敏感的训练优化研究:关注硬件利用率(FLOPs)的性能工程研究
- Google 内部工作流集成:与 XManager、TensorBoard、Vizier(PyGlove) 深度整合,适合已有 Google ML 基础设施的团队
坑与注意
⚠️ TPU 强依赖:Pax 核心设计目标是 TPU,对 NVIDIA GPU 的支持由 NVIDIA 社区维护(NVIDIA/JAX-Toolbox 的 Rosetta 项目),GPU 版本的稳定性和功能完整性不如 TPU 版本,如需 H100 FP8 训练请访问 https://github.com/NVIDIA/JAX-Toolbox/tree/main/rosetta/rosetta/projects/pax
⚠️ 依赖版本地狱:JAX 版本、libtpu 版本、TPU runtime 版本三者必须严格匹配。建议优先使用对应 release branch 的 requirements.txt,而非直接 pip install paxml 的最新依赖图。
⚠️ Python 路径硬编码:README 中大量示例路径是 .local/lib/python3.8/site-packages/paxml/main.py,如果使用不同安装方式(pip install -e . 或 conda),实际路径会不同,建议用 which python3 确认。
⚠️ Multi-slice 需要并发操作:跨 Pod 切片训练时,每个 slice 需要独立终端手动执行,不能在单台机器上用 & 后台运行,这是有意设计(多终端并发同步启动)。
⚠️ Praxis 必须先装:Dev 版本安装流程中 praxis 必须先于 paxml 安装,因为 Pax 层实现依赖 Praxis。
⚠️ 配置系统迁移中:官方文档提到 HParams 正在迁移到 Fiddle 配置系统,实验代码中可能同时存在两套配置方式,兼容性需注意。
⚠️ PyPI 包版本落后 GitHub:PyPI 上的 stable release 可能落后 main 分支数月甚至数个 release,对新模型架构或新 TPU 型号(如 v5)的支持,建议直接用 git clone 指定 branch。
⚠️ GFSA bucket 权限:所有实验输出到 --job_log_dir=gs://<bucket>,需要 TPU VM 有对应 GCS bucket 的写权限(gsutil 或 gcloud auth 配置)。
与同类对比
| 框架 | 定位 | 硬件 | 分布式策略 | 上手难度 |
|---|---|---|---|---|
| Pax | 大规模训练框架 | TPU 为主,GPU 社区支持 | SPMD/pipeline/multi-slice | 高(需要 TPU + Google Cloud) |
| Flax | 层 API | TPU + GPU | 需手写 pjit | 中(推荐研究原型) |
| MaxText | 纯文本大模型训练 | TPU + GPU | SPMD | 中 |
| JAX-Toolbox/Rosetta | Pax 的 NVIDIA GPU 分支 | NVIDIA GPU(H100等) | 同 Pax | 高 |
| PyTorch + FSDP | 通用训练框架 | GPU 为主 | FSDP/DDP | 低(社区最成熟) |
| Megatron-DeepSpeed | 大模型分布式 | GPU 为主 | 张量/流水线并行 | 中 |
核心结论:如果你的战场是 TPU、目标是超大规模(10B+)预训练、已有 Google Cloud 环境,选 Pax;如果在 GPU 上跑、或需要快速实验迭代,选 Flax/MaxText 更合适。
一句话推荐结论
如果你在 Google Cloud TPU 上训练十亿参数以上的Transformer模型、追求接近硬件峰值的 MFU 利用率,且不介意接受 Google 内部工具链(XManager/Fiddle/Praxis)的学习曲线,Pax 是目前最成熟的生产级选择;否则 Flax + 手写分布式 或 MaxText 的组合更灵活轻量。
来源:GitHub README(https://github.com/google/paxml)+ docs/about-pax.md + docs 目录结构页;NVIDIA Rosetta GPU 分支(https://github.com/NVIDIA/JAX-Toolbox/tree/main/rosetta/rosetta/projects/pax)。不确定处已标注 ⚠️。