Hugging Face PyTorch / XLA TPU Integration

Hugging Face는 사용자가 표준 Hugging Face Trainer 인터페이스를 유지하면서 Cloud TPU에서 트랜스포머 모델을 학습시키고 확장할 수 있도록 PyTorch / XLA를 통합했습니다. 이 통합은 PyTorch / XLA 라이브러리를 활용하여 PyTorch 프레임워크를 Cloud TPU를 포함한 XLA (Accelerated Linear Algebra) 장치와 연결합니다.

PyTorch / XLA 기술적 구현

이 통합은 PyTorch에 xla 장치 유형을 도입하여 TPU 하드웨어에서 텐서를 생성하고 관리할 수 있도록 합니다. Hugging Face Trainer 모듈은 is_torch_tpu_available()이 true일 때 TrainingArguments 데이터 클래스를 사용하여 TPU 장치를 자동으로 감지하고 반환합니다.

Gradient Consolidation and Optimizer Steps

Cloud TPU 장치는 일반적으로 여러 개의 코어(예: 하나의 장치가 8개의 코어를 가질 수 있음)로 구성되므로, 데이터 병렬 복제본 간에 그래디언트(gradient)를 교환해야 합니다. 이 통합은 xm.optimizer_step(optimizer)를 사용하여 그래디언트 통합 및 후속 옵티마이저 단계를 처리하여 TPU 코어 간의 동기화를 보장합니다.

Input Pipelining

호스트 CPU와 TPU 가속기가 서로를 기다리며 유휴 상태가 되는 것을 방지하기 위해, PyTorch / XLA는 입력 파이프라이닝을 구현합니다. pl.MpDeviceLoader를 사용하면 시스템이 $n$ 단계가 실행되는 동안 $n+1$ 단계의 트레이싱(tracing)을 중첩하여 수행할 수 있어, 모델로의 데이터 공급을 최적화할 수 있습니다.

Checkpoint Management

이식성을 보장하고 장치별 로딩 문제를 피하기 위해, 텐서는 체크포인트를 생성하기 전에 CPU로 이동됩니다. xm.save() API는 단 하나의 프로세스(마스터 오디널)만이 저장 위치에 기록하도록 하여 멀티 프로세스 환경에서 파일 손상을 방지합니다.

How PyTorch / XLA Works

Lazy Tensor Execution

연산을 즉시 실행하는 CPU 및 CUDA 텐서와 달리, XLA 텐서는 지연(lazy) 방식입니다. 이들은 결과가 필요할 때까지 연산을 그래프에 기록합니다. 이러한 지연 실행 방식은 XLA 컴파일러가 여러 개의 개별 연산을 하나의 최적화된 연산으로 융합(fuse)할 수 있게 해줍니다.

The Trace-Compile-Execute Cycle

PyTorch / XLA는 TPU 성능을 최적화하기 위해 특정 실행 흐름을 따릅니다:

  1. Tracing: 순전파(forward) 및 역전파(backward) 패스가 실행됨에 따라 중간 표현(IR) 그래프가 즉석에서 트레이싱됩니다.
  2. Truncation: xm.mark_step()이 호출될 때(종종 MpDeviceLoader를 통해 간접적으로 호출됨), 활성 그래프가 절단됩니다.
  3. Compilation: IR 그래프는 XLA Higher Level Operations (HLO)로 낮아지고, TPU 바이너리로 컴파일되어 실행됩니다.
  4. Caching: 재컴파일의 높은 비용을 피하기 위해, 컴파일된 TPU 바이너리는 HLO 그래프의 고유한 해시를 키로 하는 캐시에 저장됩니다.

캐시 적중률을 높이고 컴파일 오버헤드를 최소화하려면 텐서 모양(shape)을 стаtic하게 유지하는 것이 권장됩니다. Hugging Face 모델은 일반적으로 입력 토큰을 적절히 패딩하여 정적적인 모양을 유지합니다.

Performance Benchmarks

WikiText103 데이터셋에서 v3-8 Cloud TPU 시스템(4개의 TPU v3 칩)을 사용하여 bert-large-uncased를 학습시킨 결과는 다음과 같습니다:

Name Global Batch Size Precision Training Time (mins)
bert-large-uncased 64 FP32 178.4
bert-large-uncased 128 BF16 106.4

이 벤치마크는 워크로드가 호스트 CPU에 의해 제한되지 않도록 하기 위해 n1-standard-96 CPU 구성을 사용하여 수행되었습니다.

Sources