[논문 리뷰] Exploring the limits of Concurrency in ML Training on Google TPUs
이 논문은 모델 병렬화, 통신 최적화 및 분산 평가를 통해 4,096개의 TPU-v3 코어로 딥러닝 학습을 확장하는 기법을 제시하며, 네 개의 MLPerf 모델에서 기록적인 학습 시간 16~28초를 달성한다. 이 방법은 Google TPU Multipod에서 거의 완벽한 확장성을 실현하기 위해 통신, 데이터 파이프라인 및 최적화기 분할의 병목 현상을 해결한다.
Recent results in language understanding using neural networks have required training hardware of unprecedentedscale, with thousands of chips cooperating on a single training run. This paper presents techniques to scaleML models on the Google TPU Multipod, a mesh with 4096 TPU-v3 chips. We discuss model parallelism toovercome scaling limitations from the fixed batch size in data parallelism, communication/collective optimizations,distributed evaluation of training metrics, and host input processing scaling optimizations. These techniques aredemonstrated in both the TensorFlow and JAX programming frameworks. We also present performance resultsfrom the recent Google submission to the MLPerf-v0.7 benchmark contest, achieving record training times from16 to 28 seconds in four MLPerf models on the Google TPU-v3 Multipod machine.
연구 동기 및 목표
- 최대 학습 처리량을 확보하기 위해 딥러닝 모델을 전체 4,096개의 Google TPU-v3 Multipod로 확장하는 것.
- 배치 크기가 고정되어 있는 데이터 병렬화의 한계를 극복하기 위해 BERT, SSD, Transformers와 같은 대규모 모델에서 모델 병렬화를 적용하는 것.
- 지연 시간을 최소화하고 하드웨어 활용도를 극대화하기 위해 스케일링 시 통신, 시스템 수준의 조율 및 입력 파이프라인 성능을 최적화하는 것.
- TensorFlow 및 JAX 프레임워크 모두에서 고성능 학습을 구현하고, 크로스 스택 시스템 및 컴파일러 최적화에 중점을 두는 것.
- 확장성 병목 현상을 분석하고 프레임워크별 이점 평가를 통해 대규모 ML 학습의 최적 실천 방법을 수립하는 것.
제안 방법
- BERT 및 Transformers와 같이 데이터 병렬화가 배치 크기로 제한되는 대규모 모델의 경우, 큰 레이어를 여러 TPU 코어에 분산시켜 모델 병렬화를 적용한다.
- 4,096개의 TPU 메시에서 통신 오버헤드를 줄이기 위해 최적화된 all-reduce 통신 프리미티브를 구현하여, 스케일링 시 BERT에서 총 장치 시간의 27.3%를 차지하는 통신 오버헤드를 감소시킨다.
- 가중치 업데이트 분할과 혼합 정밀도 학습을 활용한 SPMD 파artitioning을 통해 모델 병렬 학습의 효율성을 향상시킨다.
- 호스트 입력 파이프라인과 학습 메트릭의 분산 평가를 최적화하여 호스트 측 병목 현상을 줄이고 종단 간 처리량을 향상시킨다.
- 특히 소규모 배치 또는 빈번한 업데이트 시나리오에서 유리한 다중 클라이언트 실행 모델을 활용해 컴파일 및 시작 오버헤드를 감소시킨다.
- gather/scatter 연산을 einsum으로 대체하고 하이퍼파라미터를 튜닝하는 등 모델 특화 최적화를 적용하여 수렴 시간과 확장성을 향상시킨다.
실험 결과
연구 질문
- RQ1데이터 병렬화의 배치 크기 제한을 극복하기 위해 모델 병렬화를 어떻게 4,096개의 TPU-v3 코어로 효과적으로 확장할 수 있는가?
- RQ24,096노드 규모에서 높은 확장 효율성을 유지하기 위해 필요한 통신 및 시스템 수준 최적화는 무엇인가?
- RQ3다양한 딥러닝 프레임워크(TensorFlow 대비 JAX)는 대규모 학습 워크로드에서 어떻게 성능을 내며, 각각의 강점은 무엇인가?
- RQ4대규모 모델 학습에서 지배적인 성능 병목 현상은 무엇이며, 이를 어떻게 완화할 수 있는가?
- RQ5시스템, 컴파일러 및 프레임워크 수준 최적화를 조율함으로써 종단 간 학습 시간을 얼마나 줄일 수 있는가?
주요 결과
- 4,096개의 TPU-v3 코어를 갖춘 TPU-v3 Multipod는 네 개의 MLPerf 모델에서 기록적인 학습 시간 16~28초를 달성하여 MLPerf-v0.7 경쟁에서 새로운 기준을 설정했다.
- BERT는 16에서 4,096개의 코어까지 높은 확장성을 보였으며, 스케일링 시 통신 오버헤드(모든-감소)가 총 장치 단계 시간의 27.3%를 차지했다.
- 모델 병렬화 덕분에 SSD, MaskRCNN 및 Transformer 모델에서 빠른 성능 향상을 이룩했으며, Transformer 모델의 경우 네 개의 TPU-v3 코어에서 2.3배의 성능 향상을 관측했다.
- 강력한 컴파일러 최적화, 효율적인 기울기 합산 및 분산 평가의 조합이 시스템 수준 병목 현상을 감소시키고 전체 처리량을 향상시켰다.
- 다중 클라이언트 실행 모델 덕분에 JAX가 두 개의 MLPerf-v0.7 벤치마크에서 컴파일 및 시작 오버헤드가 낮아 TensorFlow를 초월했다.
- 연구는 통신 오버헤드와 비효율적인 파artitioning(예: 파artitioning 후 소규모 공간 차원)이 모델 병렬 학습에서 주요 확장성 병목 현상임을 확인했다.
더 나은 연구,지금 바로 시작하세요
논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.
카드 등록 없음 · 무료 플랜 제공
이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.