JAX의 속도는 우연히 얻어지는 게 아니야. 핵심에는 Google이 만든 도메인 특화 컴파일러 XLA(Accelerated Linear Algebra)가 있어. JAX는 함수를 추적해 XLA HLO 중간 표현으로 바꾸고, XLA는 연산 융합과 메모리 배치 최적화, 장치별 코드 생성을 수행해.
- NumPy: C와 Fortran 백엔드를 이용해 CPU에서 실행하며 GPU를 직접 사용하지 않는다.
- jax.numpy: XLA가 실행 방식을 정한다. CPU에서는 LLVM 최적화를, GPU에서는 cuDNN과 cuBLAS를, TPU에서는 TPU 명령을 활용한다.
특히 중요한 최적화가 연산 융합이야. NumPy에서 (x*2 + 1) ** 0.5를 계산하면 중간 결과가 메모리에 세 번 만들어졌다가 사라질 수 있어. XLA는 이 연산들을 하나의 융합 커널로 묶어 중간 메모리 이동을 줄여.
import jax
import jax.numpy as jnp
import time
x_jax = jnp.zeros((2048, 2048))
@jax.jit
def f(x):
return jnp.sin(x @ x) + jnp.cos(x @ x)
f(x_jax).block_until_ready() # warm up
t = time.time()
for _ in range(10):
f(x_jax).block_until_ready()
print(f"JAX: {time.time()-t:.3f}s")
🔬 이렇게 생각해
JAX 코드를 작성할 때는 "나는 NumPy와 비슷한 코드를 쓰고, XLA가 알맞은 하드웨어에서 실행한다"고 생각하면 좋아. jax.devices()를 호출해 실제로 어떤 장치가 선택됐는지 확인해 봐.