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 추론을 실행하는 것은 모델 파라미터를 파이프라인과 분리하여 처리하는 기능적 접근 방식을 포함합니다. 주요 구현 단계는 다음과 같습니다:

  1. 모델 로딩: FlaxStableDiffusionXLPipeline.from_pretrained를 사용하여 기본 SDXL 1.0 모델을 로드합니다.
  2. 정밀도 관리: 메모리 사용량을 줄이고 속도를 높이기 위해 모델 파라미터를 bfloat16으로 변환하는 동시에, 저품질 또는 검은색 이미지를 초래하는 정밀도 오류를 방지하기 위해 스케줄러 상태는 float32로 유지합니다.
  3. 입력 준비: prepare_inputs를 사용하여 호출 간에 프롬프트가 일관된 차원을 갖도록 보장하며, 이는 JIT 컴파일에 필수적입니다.
  4. 장치 복제: 사용 가능한 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초입니다.

Sources