Cloud TPU v5e에서 JAX를 사용하여 Stable Diffusion XL 추론 가속화하기
Hugging Face는 Cloud TPU v5e에서 Stable Diffusion XL (SDXL)에 대한 고성능 및 비용 효율적인 추론을 가능하게 하기 위해 Diffusers 라이브러리에 JAX 지원을 통합했습니다. 이 통합은 JAX의 just-in-time (JIT) 컴파일과 XLA 기반 병렬 처리를 활용하여, 이전 모델보다 UNet 크기가 약 3배 더 큰 SDXL의 계산 문제를 해결합니다.
JAX 및 TPU v5e를 통한 기술적 최적화
Cloud TPU v5e에서 SDXL을 서빙하면 JIT 컴파일 및 SPMD 병렬 처리라는 두 가지 주요 소프트웨어 및 하드웨어 메커니즘을 통해 높은 효율성을 달성할 수 있습니다.
정적 형상을 위한 JIT 컴파일
JAX는 초기 실행 중에 코드를 추적하고 이후 호출을 위해 최적화된 TPU 바이너리를 생성하는 just-in-time (JIT) 컴파일을 활용합니다. 이 프로세스에는 정적인 입력, 중간 및 출력 형상이 필요합니다. SDXL은 다음과 같은 이유로 JIT 컴파일과 높은 호환성을 가집니다:
- 일정한 출력 형상: 이미지 생성은 일반적으로 고정된 수의 이미지와 일관된 크기를 사용합니다.
- 고정된 형상의 임베딩: Stable Diffusion 및 SDXL은 텍스트 프롬프트에 대해 고정된 형상의 임베딩 벡터(패딩 포함)를 사용합니다.
초기 컴파일에는 몇 분(제공된 예시에서는 약 3분)이 소요되지만, 이후의 추론 호출은 크게 가속화됩니다.
XLA 병렬 처리 및 처리량
JAX의 pmap은 single-program multiple-data (SPMD) 실행을 가능하게 하여 워크로드를 여러 XLA 장치에 확장할 수 있도록 합니다. 이를 통해 이미지 생성의 선형 확장이 가능합니다. 예를 들어, 8개의 칩을 가진 TPU는 단일 칩이 이미지를 하나 생성하는 데 걸리는 시간 동안 8개의 이미지를 생성할 수 있습니다. Cloud TPU v5e 인스턴스는 초고속 ICI 링크로 연결된 다양한 구성(1개에서 256개 칩까지)으로 제공되어 사용자가 특정 처리량 요구 사항에 따라 확장할 수 있습니다.
JAX에서의 구현 파이프라인
JAX로 SDXL 추론을 실행하는 것은 모델 파라미터를 파이프라인과 분리하여 처리하는 기능적 접근 방식을 포함합니다. 주요 구현 단계는 다음과 같습니다:
- 모델 로딩:
FlaxStableDiffusionXLPipeline.from_pretrained를 사용하여 기본 SDXL 1.0 모델을 로드합니다. - 정밀도 관리: 메모리 사용량을 줄이고 속도를 높이기 위해 모델 파라미터를
bfloat16으로 변환하는 동시에, 저품질 또는 검은색 이미지를 초래하는 정밀도 오류를 방지하기 위해 스케줄러 상태는float32로 유지합니다. - 입력 준비:
prepare_inputs를 사용하여 호출 간에 프롬프트가 일관된 차원을 갖도록 보장하며, 이는 JIT 컴파일에 필수적입니다. - 장치 복제: 사용 가능한 TPU 칩 전체에 파라미터와 입력을 복제하고(예: TPU v5e-4에 대해
replicate사용), 다양한 이미지 출력을 보장하기 위해 각 칩에 고유한 랜덤 시드를 할당합니다. n5. 실행:jit=True와 함께 파이프라인을 호출하여 XLA 컴파일 프로세스를 트리거합니다.
성능 벤치마크
Euler Discrete 스케줄러를 사용하여 20단계로 SDXL 1.0 base에서 수행된 벤치마크 결과, TPU v5e가 TPU v4보다 우수한 비용 효율성을 제공함을 보여줍니다.
| 하드웨어 | 배치 크기 | 지연 시간 | Perf/$ |
|---|---|---|---|
| TPU v5e-4 (JAX) | 4 | 2.33s | 21.46 |
| TPU v5e-4 (JAX) | 8 | 4.99s | 20.04 |
| TPU v4-8 (JAX) | 4 | 2.16s | 9.05 |
| TPU v4-8 (JAX) | 8 | 4.17s | 8.98 |
TPU v5e는 TPU v4에 비해 달러당 성능이 최대 2.4배 더 높습니다. 성능은 처리량(칩당 지연 시간으로 나눈 배치 크기)을 계산하고 이 수치를 하드웨어의 권장 소비자 가격으로 나누어 측정합니다.
배포 아키텍처
현재 구현은 로드 밸런싱 서버를 사용하여 사용자 요청을 사전 할당된 Cloud TPU v5e-4 인스턴스에서 실행되는 백엔드 서버로 무작위로 라우팅합니다. 각 인스턴스는 약 4초(프론트엔드 처리 및 통신 포함) 만에 1024×1024 이미지 4개를 생성하며, 실제 생성 시간은 약 2.3초입니다.