Installation
Installation
JAX를 쓰려면 두 개 패키지를 설치해야 해요. 순수 파이썬이고 크로스 플랫폼인 jax, 그리고 컴파일된 바이너리를 담고 있어 운영체제와 가속기마다 다른 빌드가 필요한 jaxlib예요. 일반적인 사용자라면 대부분 아래처럼 설치하면 돼요.
요약
- CPU 전용 (Linux/macOS/Windows)
pip install -U jax
- GPU (NVIDIA, CUDA 13)
pip install -U "jax[cuda13]"
- GPU (AMD, ROCm)
pip install -U "jax[rocm7-local]"
- TPU (Google Cloud TPU VM)
pip install -U "jax[tpu]"
지원 플랫폼
아래 표는 지원되는 모든 플랫폼과 설치 옵션이에요. 자신의 환경이 지원되는지 확인하고, "yes"나 "experimental"이라면 링크를 눌러 상세 설치 방법을 확인해요.
| 플랫폼 | Linux x86_64 | Linux aarch64 | Mac aarch64 | Windows x86_64 | Windows WSL2 x86_64 |
|---|---|---|---|---|---|
| CPU | yes | yes | yes | yes | yes |
| NVIDIA GPU | yes | yes | n/a | no | experimental |
| Google Cloud TPU | yes | n/a | n/a | n/a | n/a |
| AMD GPU | yes | no | n/a | no | experimental |
| Apple GPU | n/a | no | experimental | n/a | n/a |
| Intel GPU | experimental | n/a | n/a | no | no |
CPU
pip 설치: CPU
현재 JAX 팀은 다음 운영체제·아키텍처용 jaxlib wheel을 배포해요.
- Linux, x86_64
- Linux, aarch64
- macOS, Apple ARM 기반
- Windows, x86_64 (experimental)
노트북에서 로컬 개발에 쓸 CPU 전용 버전을 설치하려면:
pip install --upgrade pip
pip install --upgrade jax
Windows라면 Microsoft Visual Studio 2019 Redistributable이 설치돼 있지 않을 때 설치가 필요할 수 있어요.
다른 운영체제·아키텍처는 소스에서 빌드해야 해요. 지원되지 않는 환경에서 pip install 하면 jax는 설치돼도 jaxlib이 함께 설치되지 않아 런타임에 실패할 수 있어요.
NVIDIA GPU
CUDA 12에서 JAX는 SM 버전 5.2(Maxwell) 이상의 NVIDIA GPU를 지원해요. Kepler 시리즈 GPU는 NVIDIA가 소프트웨어 지원을 중단해 더 이상 지원되지 않아요. CUDA 13에서 JAX는 SM 버전 7.5 이상의 GPU를 지원해요.
먼저 NVIDIA 드라이버를 설치해야 해요. NVIDIA에서 제공하는 최신 드라이버를 권장하지만, Linux의 CUDA 12엔 드라이버 버전 525 이상, CUDA 13엔 580 이상이 필요해요.
드라이버를 쉽게 업데이트할 수 없는 클러스터처럼 최신 CUDA 툴킷을 오래된 드라이버와 함께 써야 한다면, NVIDIA가 제공하는 CUDA forward compatibility 패키지를 쓸 수 있어요.
pip 설치: NVIDIA GPU (pip로 설치, 더 쉬움)
pip install --upgrade "jax[cuda13]"
이 설치 방식은 pip가 CUDA 의존성을 함께 설치하므로, 시스템에 CUDA를 따로 설치할 필요가 없어요. 대신 NVIDIA 드라이버는 여전히 필요해요.
더 알아보기
- Quickstart: How to think in JAX — JAX 시작하기
- Key concepts — jit·vmap·grad 변환
- Automatic differentiation — 그라디언트 계산