ML / AI 고성능 수치 연산 가이드
JAX 완전 가이드
JAX는 NumPy 스타일 API에 자동 미분, JIT 컴파일, 벡터화와 다중 장치 실행을 결합한 고성능 수치 연산 라이브러리입니다. 설치부터 grad, jit, vmap, 상태 관리와 운영 검증까지 실무 흐름으로 정리합니다.
- JAX NumPy
- 자동 미분
- JIT 컴파일
- vmap
- GPU·TPU
- 수치 검증
JAX 웹 IDE 열기 →
JAX의 핵심은 NumPy와 비슷한 배열 연산을 함수 변환과 결합하는 것입니다.
grad는 자동 미분, jit는 컴파일, vmap은 배치 벡터화를 담당합니다. 실무에서는 순수 함수, 안정적인 shape, PRNG key와 컴파일 비용을 함께 관리해야 합니다.
1. JAX란?
JAX는 NumPy와 유사한 문법으로 배열 연산을 작성하면서 자동 미분, 컴파일, 벡터화와 가속기 실행을 적용할 수 있는 라이브러리입니다. 대규모 모델 연구와 고성능 수치 실험처럼 계산량이 많은 작업에 적합합니다.
| 기능 | 역할 | 주요 사용 목적 |
|---|---|---|
jax.grad |
함수의 gradient를 자동 계산합니다. | 손실 함수 미분과 최적화 |
jax.jit |
함수를 컴파일해 반복 실행 비용을 줄입니다. | 행렬 연산과 학습 단계 가속 |
jax.vmap |
단일 샘플 함수를 배치 함수로 변환합니다. | 수동 반복문 없는 벡터화 |
| 다중 장치 실행 | 여러 가속기에서 계산을 분산합니다. | 대규모 학습과 병렬 수치 연산 |
함수 작성
자동 미분
함수 변환
실행
2. 설치와 실행 확인
먼저 프로젝트 전용 가상환경을 만들고 CPU용 JAX를 설치합니다. GPU 환경은 운영체제, GPU 드라이버와 CUDA 버전에 따라 설치 명령이 달라질 수 있으므로 공식 설치 안내를 기준으로 구성하는 것이 안전합니다.
Bash# 가상환경 생성
uv venv
# Linux / macOS
source .venv/bin/activate
# Windows PowerShell
# .venv\Scripts\Activate.ps1
# CPU 환경 설치
uv pip install --upgrade jax
# 설치와 장치 확인
python -c "import jax; print(jax.__version__); print(jax.devices())"
NVIDIA GPU용 패키지는 CUDA 버전에 따라 설치 옵션이 달라질 수 있습니다. 최신 명령은 JAX 공식 설치 문서에서 확인하세요.
3. JAX 배열 기본
JAX 배열은 jax.numpy를 통해 NumPy와 유사하게 사용할 수 있습니다. 다만 JAX 변환이 적용되는 함수에서는 Python의 가변 상태보다 입력과 반환값이 명확한 함수 형태가 중요합니다.
import jax
import jax.numpy as jnp
x = jnp.array([
[1.0, 2.0],
[3.0, 4.0],
])
w = jnp.ones((2, 3))
result = x @ w
print("shape:", result.shape)
print("dtype:", result.dtype)
print("devices:", jax.devices())
print(result)
- shape: 입력 크기가 바뀌면 JIT 재컴파일이 발생할 수 있습니다.
- dtype: float32와 float64 차이가 결과와 성능에 영향을 줄 수 있습니다.
- device: CPU, GPU, TPU 중 실제 실행 장치를 확인합니다.
- 불변성: 배열을 직접 변경하기보다 새 배열을 반환하는 방식을 사용합니다.
4. 자동 미분
jax.grad는 스칼라 값을 반환하는 함수의 gradient를 계산합니다. 손실 함수와 모델 파라미터를 함수 입력으로 명확하게 전달하면 미분과 JIT 변환을 함께 적용하기 쉬워집니다.
import jax
import jax.numpy as jnp
def loss(w):
return jnp.sum((w - 3.0) ** 2)
grad_loss = jax.grad(loss)
w = jnp.array([1.0, 2.0, 4.0])
print("loss:", loss(w))
print("gradient:", grad_loss(w))
값과 gradient를 한 번에 계산하기
Pythonvalue_and_grad_loss = jax.value_and_grad(loss)
value, gradient = value_and_grad_loss(w)
print("value:", value)
print("gradient:", gradient)
gradient 결과가 예상과 다를 때는 dtype, 입력 범위, NaN과 Inf, 손실 함수의 스케일을 함께 확인해야 합니다. 작은 입력으로 수치 미분 결과와 비교하는 gradient sanity check도 유용합니다.
5. JIT 컴파일
자주 반복되는 순수 함수는 jax.jit로 컴파일해 실행할 수 있습니다. 첫 호출에는 컴파일 시간이 포함되지만 같은 shape와 dtype으로 반복 실행하면 이후 호출에서 컴파일된 결과를 재사용할 수 있습니다.
import jax
import jax.numpy as jnp
@jax.jit
def matmul(a, b):
return a @ b
x = jnp.ones((1024, 1024))
result = matmul(x, x)
result.block_until_ready()
print(result.shape)
block_until_ready()를 호출해 실제 연산 완료 시점까지 기다린 뒤 측정하세요.컴파일과 실행 시간을 분리해서 측정하기
Pythonimport time
start = time.perf_counter()
first = matmul(x, x).block_until_ready()
compile_and_run = time.perf_counter() - start
start = time.perf_counter()
second = matmul(x, x).block_until_ready()
run_only = time.perf_counter() - start
print("compile + run:", compile_and_run)
print("cached run:", run_only)
6. 벡터화
jax.vmap은 단일 샘플을 처리하는 함수를 배치 입력을 처리하는 함수로 변환합니다. Python 반복문 없이 배치 차원을 표현할 수 있어 코드가 간결해지고 JIT 컴파일과 결합하기도 쉽습니다.
import jax
import jax.numpy as jnp
def predict(w, x):
return jnp.dot(w, x)
w = jnp.array([0.2, 0.5, -0.1])
batch_x = jnp.array([
[1.0, 2.0, 3.0],
[0.5, 1.0, 1.5],
[2.0, 0.0, 1.0],
])
batched_predict = jax.vmap(
predict,
in_axes=(None, 0),
)
print(batched_predict(w, batch_x))
JIT와 vmap 결합
Pythonfast_batched_predict = jax.jit(
jax.vmap(
predict,
in_axes=(None, 0),
)
)
result = fast_batched_predict(w, batch_x)
print(result)
7. PRNG와 상태 관리
JAX에서는 난수 상태를 숨겨진 전역 상태로 처리하지 않고 PRNG key를 함수 입력과 출력으로 명시적으로 전달합니다. key를 재사용하지 않고 필요한 만큼 분리하는 규칙이 중요합니다.
Pythonimport jax
key = jax.random.key(42)
key, sample_key = jax.random.split(key)
samples = jax.random.normal(
sample_key,
shape=(3, 4),
)
print(samples)
같은 key를 반복 사용하면 동일한 난수 결과가 만들어질 수 있습니다. 학습 단계, dropout, 데이터 증강 등 서로 다른 연산에는
jax.random.split()으로 분리한 key를 전달하세요.8. JAX 실무 설계
JAX는 순수 함수와 불변 데이터를 전제로 설계할수록 jit, vmap과 다중 장치 확장이 쉬워집니다. model params, optimizer state, random key를 숨기지 말고 함수의 입력과 반환값으로 명확하게 표현하는 것이 좋습니다.
| 결정 지점 | 확인 질문 | 실무 기준 |
|---|---|---|
| 경계 | JAX 코드에서 자주 바뀌는 부분은 어디인가? | 입출력, 설정, 외부 연동과 핵심 수치 연산을 분리합니다. |
| 상태 | 파라미터와 optimizer state는 누가 관리하는가? | 상태를 함수 입력과 반환값으로 드러냅니다. |
| 난수 | PRNG key가 중복 사용되고 있지 않은가? | key 생성, 분리, 전달 규칙을 일관되게 유지합니다. |
| 장애 | 잘못된 shape와 dtype을 어디서 차단하는가? | 입력 경계에서 계약을 검증하고 명확한 오류를 반환합니다. |
상태를 명시적으로 전달하는 함수
Pythonimport jax
import jax.numpy as jnp
def train_step(params, x, y, learning_rate):
def loss_fn(current_params):
prediction = x @ current_params
return jnp.mean((prediction - y) ** 2)
loss, grads = jax.value_and_grad(loss_fn)(params)
new_params = params - learning_rate * grads
return new_params, loss
fast_train_step = jax.jit(train_step)
9. JAX 운영 기준
JAX 운영에서는 컴파일 시간과 실제 실행 시간을 분리해 측정해야 합니다. 입력 shape가 자주 변경되면 재컴파일 비용이 커질 수 있으므로 배치 shape를 안정적으로 유지하는 것이 중요합니다.
- 동일한 작업에서 batch shape와 dtype을 안정적으로 유지합니다.
- PRNG key의 생성, 분리와 전달 규칙을 정합니다.
- 첫 컴파일과 캐시된 실행 시간을 따로 측정합니다.
- 비동기 실행을 고려해 정확한 구간에서
block_until_ready()를 사용합니다. - NaN, Inf와 gradient 크기를 모니터링합니다.
- 메모리 사용량과 가속기 utilization을 함께 확인합니다.
- 반복되는 컴파일이 발생하는지 로그와 프로파일링으로 확인합니다.
간단한 gradient 상태 확인
Pythondef gradient_is_valid(gradient):
leaves = jax.tree.leaves(gradient)
return all(
bool(jnp.all(jnp.isfinite(leaf)))
for leaf in leaves
)
10. JAX 검증 전략
JAX 코드는 일반 기능 테스트와 함께 수치 안정성, dtype 차이, gradient, PRNG 재현성과 컴파일 성능을 검증해야 합니다.
| 품질 축 | 검증 방법 | 완료 기준 |
|---|---|---|
| 정확성 | 정상·경계·실패 입력을 자동화합니다. | 핵심 시나리오가 재현 가능하게 통과합니다. |
| 수치 안정성 | NaN, Inf, dtype과 입력 범위를 확인합니다. | 지원 입력 범위에서 유효한 수치 결과를 반환합니다. |
| Gradient | 작은 입력에서 수치 미분과 자동 미분 결과를 비교합니다. | 정한 허용 오차 범위에서 결과가 일치합니다. |
| 재현성 | 동일 PRNG seed와 key 흐름을 반복 검증합니다. | 같은 조건에서 동일한 결과를 재현할 수 있습니다. |
| 성능 회귀 | 컴파일 시간과 캐시 실행 시간을 구분해 측정합니다. | 기준값 대비 허용 범위를 넘는 성능 저하가 없습니다. |
| 운영성 | 로그, 메트릭과 오류 추적 경로를 확인합니다. | 장애 발생 시 원인과 입력 조건을 추적할 수 있습니다. |
마무리
JAX를 잘 활용하려면 NumPy 스타일 배열 연산뿐 아니라 함수 변환 중심의 사고방식을 익혀야 합니다. 순수 함수와 명시적인 상태 전달을 기본으로 하고 grad, jit, vmap을 단계적으로 결합하면 연구용 코드에서 고성능 수치 연산 파이프라인까지 확장할 수 있습니다.
grad로 함수의 gradient를 계산합니다.jit는 반복 실행되는 순수 함수를 컴파일합니다.vmap으로 단일 샘플 함수를 배치 처리 함수로 확장합니다.- shape와 dtype 변화에 따른 재컴파일 비용을 관리합니다.
- PRNG key, 수치 안정성, gradient와 성능 회귀를 함께 검증합니다.
'개발 가이드 > AI 개발' 카테고리의 다른 글
| [AI 개발] 6. LC LangChain 완전 가이드 (0) | 2026.07.19 |
|---|---|
| [AI 개발] 5. Hugging Face 완전 가이드 (0) | 2026.07.19 |
| [AI 개발] 3. PyTorch 완전 가이드 (0) | 2026.07.19 |
| [AI 개발] 2. C++와 ML/AI 완전 가이드 (0) | 2026.07.19 |
| [AI 개발] 1. Python AI 완전 가이드: 설치부터 AI·FastAPI 개발까지 (0) | 2026.07.18 |
댓글