JAX 본체인 jax와 jaxlib는 의도적으로 작은 범위만 맡아. 함수 변환과 수치 연산, XLA 인터페이스까지가 코어의 책임이야. 실제 머신러닝과 과학 계산에 필요한 상위 기능은 별도 라이브러리 생태계가 제공해.
- Flax (NNX): Google이 개발하는 JAX 네이티브 신경망 라이브러리다.
- Equinox: 모델을 pytree로 표현하는 함수형 신경망 라이브러리다.
- Optax: 옵티마이저와 그래디언트 변환을 제공하며 사실상 표준으로 쓰인다.
- Orbax: 체크포인트 저장과 복원을 담당하는 표준 도구다.
- Diffrax: 미분 가능한 ODE·SDE solver다.
- NumPyro: 베이지안 추론과 확률적 프로그래밍을 지원한다.
- Brax: GPU와 TPU에서 실행되는 미분 가능한 강체 물리 엔진이다.
설치 방법은 사용할 플랫폼에 따라 조금씩 달라.
# CPU 만 (지금 시작하기 가장 쉬움)
pip install -U jax
# NVIDIA GPU (CUDA 12)
pip install -U "jax[cuda12]"
# Cloud TPU
pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
설치가 끝나면 다음 코드로 버전과 장치, 기본 백엔드를 확인해.
import jax
import jax.numpy as jnp
print("JAX version:", jax.__version__)
print("Devices:", jax.devices())
print("Default backend:", jax.default_backend())
x = jnp.array([1.0, 2.0, 3.0])
print(jnp.sum(x)) # 6.0
💡 conda 또는 virtualenv 환경을 따로 만들기
JAX는 플랫폼별 jaxlib 바이너리와 다른 패키지의 버전을 맞추는 일이 까다로울 수 있어. 새 venv나 conda 환경을 하나 만들어 시작하고, 이 퀘스트를 끝낼 때까지 같은 환경을 재사용해.