대부분의 머신러닝 프레임워크는 객체를 중심으로 설계돼 있어. 모델 객체를 만들고 데이터를 전달한 다음 메서드를 호출하지. JAX는 다른 출발점을 택해. 핵심은 합성 가능한 함수 변환, 즉 함수를 입력으로 받아 변환된 새 함수를 반환하는 고차 함수야.
JAX의 네 가지 핵심 변환은 다음과 같아.
jax.jit, 컴파일: 함수를 추적한 뒤 XLA 컴파일러로 최적화된 머신 코드를 만든다. 첫 호출은 컴파일 때문에 느리지만 이후 호출은 훨씬 빨라진다.jax.grad, 미분: 스칼라를 반환하는 함수를 받아 그래디언트를 계산하는 새 함수를 돌려준다.jax.vmap, 벡터화: 예제 하나를 처리하는 함수를 배치 전체를 처리하는 함수로 바꾼다.jax.pmap, 병렬화: 여러 장치, 즉 GPU나 TPU에서 함수를 동시에 실행한다.
여기서 가장 중요한 말은 합성 가능하다는 거야. 변환을 원하는 순서로 겹쳐 쓸 수 있어.
import jax
import jax.numpy as jnp
def loss_fn(params, x, y):
predictions = jnp.dot(x, params)
return jnp.mean((predictions - y) ** 2)
# 합성: 미분 → compile
fast_grad = jax.jit(jax.grad(loss_fn))
# 또는: per-example gradient → compile
per_example_grads = jax.jit(jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0)))
💡 왜 이게 중요한가
각 변환은 따로 써도 유용하지만 JAX의 진짜 힘은 합성에서 나와. jit(vmap(grad(f)))처럼 미분하고 벡터화한 함수를 다시 컴파일할 수 있어. PyTorch에서는 같은 결과를 위해 코드를 상당히 다시 구성해야 할 수 있지만 JAX에서는 변환을 한 줄로 조합해. 변환을 일급 값으로 다루는 이 설계가 미분 가능 프로그래밍의 중요한 전환점이야.