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)을 활용합니다.

관련

  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트