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 解决了这些问题:

  1. 开箱即用的分布式策略:支持 pmap(单主机多设备)和 pjit/SPMD(多主机跨 Pod 横向扩展),无需手写collective原语
  2. 配置驱动的实验管理:通过 Fiddle 配置类(而非硬编码超参)管理实验变体,配合 XManager 做实验编排
  3. MFU 驱动的性能基准:提供 TPU v4 Pod 弱扩展(weak scaling)基准,展示 1B → 16B → GPT3-XL 在不同规模下的训练效率
  4. 多切片(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 资源。


典型适用场景

  1. 超大规模语言模型预训练:540B PaLM 模型背后训练基础设施的同类架构,适合 1B–540B 参数量级的 Transformer 训练研究
  2. 多模态大模型研究:官方 repo 包含 vision、speech、text 多种模态的配置示例
  3. TPU Pod 级分布式训练:需要横向扩展到数百乃至数千 TPU 核的场景(multi-slice),而非单机多卡
  4. MFU 敏感的训练优化研究:关注硬件利用率(FLOPs)的性能工程研究
  5. 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 的写权限(gsutilgcloud 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)。不确定处已标注 ⚠️。