TPU 소프트웨어 스택
개요
TPU 소프트웨어 스택은 Google이 개발한 TPU 하드웨어를 활용하기 위한 종합적인 소프트웨어 생태계로, JAX를 핵심으로 하는 유연하고 확장 가능한 프로그래밍 모델을 제공한다. 이 스택은 기계 학습 워크로드의 전체 수명 주기를 지원하며, 연구 단계에서 프로덕션 배포까지 원활한 전환을 가능하게 한다.
TPU 소프트웨어 스택의 핵심 철학은 "느슨하게 결합된 구성 요소(loosely coupled components)"로, 각 구성 요소가 하나의 기능에 특화되어 있다. JAX 자체는 효율적인 배열 연산과 프로그램 변환에 집중하며, 주변 라이브러리들은 이 핵심 위에 다양한 기능을 제공한다. 이러한 아키텍처는 단일 모놀리식 프레임워크와 달리, 구성 요소를 독립적으로 업데이트하고 교체할 수 있는 유연성을 제공한다.
핵심 개념
JAX AI 스택 구성 요소
JAX AI 스택은 크게 네 가지 계층으로 구성된다:
- 코어 라이브러리: JAX, Flax, Optax, Orbax, Grain
- 인프라: XLA 컴파일러, Pathways 분산 런타임
- 고급 개발 도구: Pallas, Tokamax, Qwix
- 애플리케이션 계층: MaxText, MaxDiffusion, Tunix, vLLM
JAX: 고성능 수치 연산 기반
JAX는 Python 기반의 가속기 지향 배열 연산 라이브러리로, NumPy와 유사한 API를 제공하면서도 XLA 컴파일러를 통한 JIT 컴파일, 자동 미분, 벡터화, 병렬화 등의 프로그램 변환 기능을 지원한다.
핵심 변환 기능:
- jit: Python 함수를 최적화된 XLA 실행 가능 코드로 JIT 컴파일
- grad: 자동 미분 지원 (순방향/역방향 모드)
- vmap: 자동 벡터화로 함수 로직 변경 없이 배치 및 데이터 병렬화 지원
- pmap/shard_map: 여러 디바이스(TPU 코어)에 걸친 자동 병렬화
JAX의 함수형 프로그래밍 모델은 순수 함수를 강조하여 프로그램 변환의 tractability와 compositability를 높이며, XLA의 GSPMD(General-purpose SPMD) 모델과의 원활한 통합을 통해 최소한의 코드 변경으로 대규모 TPU Pod에서의 자동 병렬화를 가능하게 한다.
XLA: 하드웨어 독립형 컴파일러
XLA(Accelerated Linear Algebra)는 Google이 개발한 도메인 특화 컴파일러로, TPU, CPU, GPU 등 다양한 하드웨어를 대상으로 최적화된 코드를 생성한다. XLA의 컴파일러 중심 설계(compiler-first design)는 안정된 하드웨어 아키텍처에서 빠르게 진화하는 연구 환경에서 지속적인 이점을 제공한다.
컴파일 파이프라인:
1. JAX 계산 그래프 → HLO(High-Level Optimizer) 표현식 변환
2. HLO에서 하드웨어 독립적 최적화 (연산자 융합, 메모리 관리 등)
3. LLO(Low-Level Optimizer)로 하향 단계적 변환
4. 하드웨어별 최적화된 머신 코드 생성 (TPU의 경우 VLIW 패킷)
XLA의 병렬화 설계는 SPMD(Single Program Multiple Data)를 중심으로 하며, 이를 통해 하나의 프로그램으로 모든 디바이스에서 효율적인 연산을 수행할 수 있다. 더 복잡한 병렬화 패턴의 경우 MPMD(Multiple Program Multiple Data)도 지원된다.
Pathways: 대규모 분산 컴퓨팅 런타임
Pathways는 수만 개의 칩에 걸쳐 분산된 계산을 조율하기 위한 통합 런타임으로, 내장된 결함 허용 및 복구 기능을 제공한다. 이를 통해 연구자들은 단일 강력한 머신을 사용하는 것처럼 프로그래밍할 수 있다.
주요 특징:
- 단일 Python 클라이언트로 멀티 Pod 작업 지원
- DCN(Data Center Network)을 통한 교차 슬라이스 병렬화
- 자동 리소스 선점 복구를 통한 내성(resilience) 확보
- ICI(Chip-to-Chip Interconnect) 및 OCI(On-Chip Interconnect) 패브릭 최적화
동작 원리
코어 라이브러리 생태계
Flax: 유연한 신경망 작성
Flax는 JAX 위에서 신경망을 직관적이고 객체 지향적으로 생성할 수 있게 해주는 라이브러리이다. JAX의 함수형 API는 강력하지만, PyTorch와 유사한 레이어 기반 추상화를 제공하여 개발자 친화적인 인터페이스를 제공한다.
NNX API의 장점:
- 모델 상태를 캡슐화하여 사용자 인지 부담 감소
- Python 스타일의 인터페이스로 모델 계층 구조의 프로그래밍 가능한 순회 및 수정 지원
- LoRA 및 양자화와 같은 기술에 필요한 조작 가능한 모델 정의 제공
Optax: 최적화 전략 라이브러리
Optax는 JAX용 그래프 처리 및 최적화 라이브러리로, 손실 함수와 최적화 알고리즘의 테스트된 구현을 제공한다. 모듈형 체이닝 가능한 변환을 통해 복잡한 최적화 전략을 선언적으로 구성할 수 있다.
설계 특징:
- 순수 함수형 구현으로 JAX의 병렬화 메커니즘과 원활한 통합
- 체이닝을 통한 복잡한 최적화 전략 구성 (예: 그래디언트 클리핑 + RMSProp + 그래디언트 누적)
- 연구 속도 향상과 프로덕션 전환 용이성 강조
Orbax: 대규모 분산 체크포인팅
Orbax는 단일 디바이스에서 대규모 분산 학습까지 지원하는 JAX용 체크포인팅 라이브러리이다. TensorStore를 활용한 효율적인 병렬 읽기/쓰기를 기반으로, 학습 복원에 필수적인 모델 가중치, 최적화기 상태, 데이터 로더 상태를 선택적으로 영속화한다.
성능 특징:
- 비동기 체크포인팅으로 가속기 유휴 시간 최소화
- OCDBT(Optimized Checkpoint Database Technology) 효율적 저장 형식
- 수만 노드에 걸친 대규모 학습 작업 지원
Grain: 결정론적 데이터 파이프라인
Grain은 JAX 모델 학습 및 평가를 위한 데이터 읽기 및 처리 라이브러리로, 결정론적이고 확장 가능한 데이터 파이프라인을 제공한다. 가속기 호스트와 데이터 워커를 동일 위치에 배치하여高效的한 데이터 처리를 달성한다.
핵심 장점:
- 결정론적 데이터 피딩으로 학습 재현성 보장
- 체크포인팅 가능한 이터레이터로 Orbax와 통합된 학습 스냅샷 지원
- ArrayRecord, Bagz, Parquet 등 다양한 데이터 형식 지원
고급 개발 도구
Pallas: 저수준 고성능 커널 작성
Pallas는 JAX의 확장으로, Python으로 저수준 고성능 커널을 작성할 수 있게 해준다. 그리드 기반 병렬화 모델을 통해 사용자 정의 커널 함수를 병렬 워크그룹의 다차원 격자 전체에서 실행할 수 있다.
메모리 계층 관리:
- 느리고 큰 메모리(HBM)와 빠르고 작은 온칩 메모리(VMEM) 간의 텐서 타일링 및 전송 명시적 관리
- 인덱스 맵을 사용하여 격자 위치와 특정 데이터 블록을 연결
- TPU의 Mosaic 또는 GPU의 Triton을 통한 대상 아키텍처별 컴파일
Tokamax: 최신 커널 라이브러리
Tokamax는 Pallas 위에 구축된 최신 가속기 커널 라이브러리로, TPU와 GPU 모두를 지원한다. 자동 튜닝 인프라를 통해 최적의 커널 구현을 선택하고 관리한다.
자동 튜닝 메커니즘:
- 구성 가능한 파라미터(예: 타일 크기)에 대한 전체 범위 스위핑
- 캐시된 결과를 기반으로 최적 설정 결정
- 야간 회귀 테스트로 컴파일러 인프라 변경으로 인한 성능 및 수치 문제 방지
Qwix: 포괄적 양자화 라이브러리
Qwix는 JAI AI 스택용 포괄적 양자화 라이브러리로, 학습(QAT, QT, QLoRA) 및 추론(PTQ) 모든 단계를 지원한다. 비침습적 모델 통합을 통해 양자화 코드를 모델 정의와 완전히 분리한다.
비침습적 통합 메커니즘:
- JAX 함수를 양자화된 대응 함수로 리다이렉트하는 인터셉션 메커니즘
- 모델 수정 없이 양자화 적용 가능
- 다양한 양자화 체계에 대한 하이퍼파라미터 조정 용이
비교/분석
TPU 소프트웨어 스택 vs 경쟁 하드웨어 소프트웨어 스택
| 특성 | TPU 소프트웨어 스택 (JAX) | NVIDIA GPU (CUDA/cuDNN) | AMD GPU (ROCm/HIP) |
|---|---|---|---|
| 핵심 프레임워크 | JAX (함수형, 컴파일러 중심) | CUDA (命令式) | HIP (CUDA 호환) |
| 컴파일러 | XLA (하드웨어 독립) | NVCC (하드웨어 종속) | ROCm/HIP 컴파일러 |
| 모델 작성 | Flax (객체 지향 + 함수형) | PyTorch/TensorFlow | PyTorch/TensorFlow |
| 최적화 | Optax (함수형 체이닝) | 별도 라이브러리 필요 | 별도 라이브러리 필요 |
| 체크포인팅 | Orbax (비동기, 대규모) | 커스텀 구현 필요 | 커스텀 구현 필요 |
| 데이터 파이프라인 | Grain (결정론적) | DALI | 커스텀 구현 |
| 커널 작성 | Pallas/Tokamax | CUDA 코어/커스텀 커널 | HIP 커널 |
| 양자화 | Qwix (비침습적) | TensorRT | MIGraphX |
| 분산 학습 | Pathways/GSPMD | NCCL | RCCL |
| 지원 하드웨어 | TPU v2~v7x | NVIDIA GPU 전반 | AMD GPU 전반 |
JAX AI 스택 구성 요소별 비교
| 구성 요소 | 역할 | 주요 특징 | 대상 사용자 |
|---|---|---|---|
| JAX | 핵심 수치 연산 | JIT, grad, vmap, pmap | 모든 ML 연구자 |
| Flax | 신경망 작성 | NNX API, 객체 지향 | 모델 개발자 |
| Optax | 최적화 전략 | 함수형, 체이닝 가능 | 최적화 연구자 |
| Orbax | 체크포인팅 | 비동기, 대규모 분산 | 대규모 학습 운영자 |
| Grain | 데이터 파이프라인 | 결정론적, 체크포인팅 | 데이터 엔지니어 |
| XLA | 컴파일러 | 하드웨어 독립, SPMD | 시스템 엔지니어 |
| Pathways | 분산 런타임 | 대규모 스케일링 | 인프라 엔지니어 |
| Pallas | 커널 작성 | 저수준 메모리 관리 | 고성능 커널 개발자 |
| Tokamax | 커널 라이브러리 | 자동 튜닝, 최신 커널 | 커널 최적화 전문가 |
| Qwix | 양자화 | 비침습적, 포괄적 | 모델 최적화 전문가 |
장단점
장점
- 유연한 아키텍처: 느슨한 결합으로 구성 요소 독립적 업데이트 및 교체 가능
- 컴파일러 중심 설계: XLA를 통한 하드웨어 독립적 최적화 및 자동 병렬화
- 대규모 스케일링: Pathways를 통한 수만 개 칩까지의 원활한 확장
- 결정론적 재현성: Grain과 Orbax를 통한 학습 과정의 완전한 재현성
- 비침습적 최적화: Qwix를 통한 모델 수정 없이 양자화 적용
- 풍부한 생태계: 학습에서 추론까지 end-to-end 지원
단점
- 학습 곡선: JAX의 함수형 프로그래밍 모델은 명령식 프레임워크 사용자에게 초기 학습 부담
- 하드웨어 종속성: TPU에 최적화되어 있어 다른 가속기에서의 성능 제한 가능
- 커뮤니티 규모: NVIDIA CUDA 생태계에 비해 상대적으로 작은 커뮤니티
- 프로덕션 도구 성숙도: 일부 고급 도구(Tokamax, Pillas)는 아직 상대적으로 젊은 프로젝트
관련 기술
- OpenXLA: XLA 컴파일러의 오픈소스 버전
- MLIR: Multi-Level Intermediate Representation, XLA의 컴파일러 인프라
- StableHLO: XLA의 안정적인 인터미디어트 리프레젠테이션
- TensorStore: 대규모 배열 데이터 효율적 저장/로딩 라이브러리
참고 문헌
- Google Cloud TPU 문서: https://cloud.google.com/tpu/docs
- JAX 공식 문서: https://jax.readthedocs.io/
- JAX AI 스택: https://jaxstack.ai/
- XLA 문서: https://openxla.org/xla
- Pathways 논문: https://arxiv.org/abs/2203.12533
- Orbax 문서: https://orbax.readthedocs.io/
- Grain 문서: https://google-grain.readthedocs.io/
- Pallas 문서: https://jax.readthedocs.io/en/latest/pallas/index.html
- Tokamax: https://github.com/openxla/tokamax
- Qwix 문서: https://qwix.readthedocs.io/
핵심 정리
- TPU 소프트웨어 스택은 JAX를 중심으로 한 느슨하게 결합된 구성 요소들의 생태계로, ML 워크로드의 전체 수명 주기를 지원한다.
- XLA 컴파일러 중심 설계를 통해 하드웨어 독립적 최적화와 자동 병렬화를 달성하며, Pathways를 통해 대규모 분산 컴퓨팅을 가능하게 한다.
- 코어 라이브러리(JAX, Flax, Optax, Orbax, Grain)는 각각 특정 기능에 특화되어 있으며, 고급 도구(Pallas, Tokamax, Qwix)는 성능 최적화를 위한 저수준 제어를 제공한다.
- 이 스택의 주요 강점은 유연한 아키텍처, 결정론적 재현성, 비침습적 최적화, 대규모 스케일링 지원에 있으며, 이를 통해 연구에서 프로DUCTION까지 원활한 전환을 가능하게 한다.