Stable Diffusion JAX 및 Flax 통합

Hugging Face는 버전 0.5.1부터 diffusers 라이브러리에 Flax 지원을 통합하여 Stable Diffusion을 Google TPU에서 높은 효율로 실행할 수 있게 했습니다. 이 통합을 통해 사용자는 일반적으로 8개의 가속기가 장착된 TPU 서버의 병렬 처리 능력을 활용하여 한 번에 하나의 이미지를 생성하는 시간에 여러 이미지를 동시에 생성할 수 있습니다.

JAX 및 Flax를 이용한 고속 TPU 추론

TPU에서의 Stable Diffusion 추론은 JAX와 Flax를 사용하여 최적화되며, 표준 GPU 구현에 비해 상당한 속도 향상을 제공합니다. TPU v2-8에서는 초기 컴파일 이후의 추론 실행이 약 7초 정도 소요됩니다.

핵심 기술 최적화 사항은 다음과 같습니다:

  • bfloat16 정밀도: TPU 장치는 bfloat16을 지원하는데, 이는 메모리 오버헤드를 줄이면서 성능을 유지하는 효율적인 반정밀도 부동소수점 타입입니다.
  • JIT 컴파일: Flax 파이프라인에 jit=True를 전달하면 JAX가 모델을 효율적인 형태로 컴파일합니다. 첫 실행은 컴파일 기간이 필요하며(TPU v2-8에서는 1분 이상) 이후 호출은 훨씬 빠릅니다.
  • 무상태 모델: Flax는 함수형 프레임워크이므로 모델은 무상태이며, 파라미터는 모델 외부에 저장됩니다.

SPMD를 통한 병렬화

diffusers Flax 파이프라인은 Single-Program, Multiple-Data (SPMD) 병렬화를 활용하여 TPU 하드웨어 활용도를 극대화합니다. 이는 주로 jax.pmap 함수를 통해 구현됩니다.

병렬화 구현 방법

jax.pmap은 두 가지 핵심 기능을 수행합니다: 코드를 컴파일(jax.jit()와 유사)하고, 컴파일된 코드가 사용 가능한 모든 장치에서 병렬로 실행되도록 보장합니다.

병렬 실행을 위해 파이프라인은 다음 단계들을 따릅니다:

  1. 복제: 모델 파라미터는 flax.jax_utils.replicate를 사용하여 모든 장치에 복제됩니다.
  2. 샤딩: 토큰화된 프롬프트 ID와 같은 입력 데이터는 shard를 사용해 샤딩됩니다. 예를 들어 8개의 장치가 있다면 프롬프트 배열을 나누어 각 장치가 입력의 특정 부분을 받게 됩니다.
  3. PRNG 처리: 생성된 이미지의 재현성과 다양성을 보장하기 위해 난수 생성기(RNG)를 만들고 이를 여러 개의 생성기로 분할합니다—각 장치당 하나씩.

이 아키텍처를 통해 파이프라인은 각 장치가 배치 항목을 독립적으로 처리하므로 동시에 8개의 서로 다른 이미지(또는 동일 이미지의 8복사본)를 생성할 수 있습니다.

모델 접근 및 라이선스

Flax용 Stable Diffusion 가중치는 CompVis/stable-diffusion-v1-4 저장소의 Hugging Face Hub에서 제공됩니다. 접근하려면 CreativeML OpenRAIL-M 라이선스를 수락해야 하며, 이 라이선스는 다음과 같은 조항을 포함합니다:

  • 사용자는 모델을 고의적으로 불법 또는 유해한 콘텐츠를 생성하거나 공유하는 데 사용할 수 없습니다.
  • 사용자는 자신이 생성한 출력물에 대한 권리를 보유하며, 그 사용에 대해 책임을 집니다.
  • 동일한 사용 제한 및 CreativeML OpenRAIL-M 라이선스 사본을 모든 사용자와 공유하는 조건 하에, 가중치의 상업적 사용 및 재배포가 허용됩니다.

Sources