google/paxml

Pax is a Jax-based machine learning framework for training large scale models. Pax allows for advanced and fully configurable experimentation and parallelization, and has demonstrated industry leading model flop utilization rates.

解决的问题

Paxml(或 Pax)提供了一个用于配置和运行大规模机器学习实验的框架,专为在 Jax 之上运行而设计。它简化了在分布式硬件(如 Cloud TPU VM 和 TPU Pod 切片)上管理复杂模型配置和部署的过程。

工作原理

Pax 利用 Jax 处理计算,并通过配置文件以结构化方式定义实验。它支持多种并行策略,包括使用 pjitpmap 的 SPMD(Single Program, Multiple Data)以及流水线并行。该框架处理多主机数据输入,确保数据在不同 Jax 进程之间分片,避免训练期间出现重复批次。它还与多种数据输入类型集成,包括 SeqIO、Lingvo 和自定义实现,以向模型提供训练、评估和解码所需的数据。

适用人群

专为训练超大规模模型(参数量从数十亿到数万亿)的机器学习研究人员和工程师设计,需要在 TPU 或 GPU 基础设施上管理这些实验的稳健系统。

主要亮点

  • TPU 优化:与 Cloud TPU VM 和 Pod 切片深度集成,实现高性能训练。
  • 可扩展性:支持大规模语言模型的弱扩展,允许模型大小随使用芯片数量成比例增长。
  • 灵活的并行:支持 pjit(SPMD)和 pmap 用于分布式执行。
  • 数据管理:内置多主机输入支持,通过 SeqIO 实现评估数据的自动分片/填充。
  • 性能指标:利用 Model FLOPs Utilization(MFU)衡量端到端训练效率。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目