JIT의 효과를 측정하는 일은 단순해 보이지만 함정이 많아. 가장 큰 함정: 비동기 디스패치.
JAX 호출은 기본값으로 비동기로 실행돼. 배열을 반환하지만 실제 계산은 백그라운드에서 진행해. time.time()으로 측정하면 Python이 ndarray를 받은 시각이지, 계산이 끝난 시각이 아니야.
import jax
import jax.numpy as jnp
import time
@jax.jit
def f(x):
return jnp.sum(x ** 2 + jnp.sin(x))
x = jnp.arange(1_000_000.)
# 잘못된 측정
t = time.time()
for _ in range(100):
y = f(x) # async 반환
print(f"잘못된 측정: {time.time()-t:.4f}s")
# 올바른 측정 — block_until_ready
t = time.time()
for _ in range(100):
y = f(x).block_until_ready()
print(f"올바른 측정: {time.time()-t:.4f}s")
차이가 10배 이상 나기도 해. 항상 .block_until_ready() 또는 jax.block_until_ready(...)를 호출해야 해.
컴파일 시간 vs 실행 시간 분리:
# Warm-up: compile + cache
y = f(x).block_until_ready()
# 이제 run 시간만
t = time.time()
for _ in range(1000):
y = f(x).block_until_ready()
elapsed = time.time() - t
print(f"호출 당: {elapsed/1000*1000:.3f}ms")
전형적 결과 (1M 원소, CPU):
- NumPy: ~ 4 ms
- JAX eager (jit 없음): ~ 6 ms, 약간 느림 (오버헤드)
- JAX jit, 첫 호출: ~ 200 ms (컴파일 포함)
- JAX jit, 이후 호출: ~ 0.5 ms, 8배 빠름
GPU에서는 더 극적: 첫 실행 ~ 500ms, 재실행 ~ 50us, 80배 이상이야.
벤치마크 스크립트:
def benchmark(f, *args, n_warmup=3, n_runs=100, label=""):
'''함수의 평균 호출 시간 측정. async dispatch 처리.'''
# warm-up (compile)
for _ in range(n_warmup):
out = f(*args)
if hasattr(out, "block_until_ready"):
out.block_until_ready()
# measure
t = time.time()
for _ in range(n_runs):
out = f(*args)
if hasattr(out, "block_until_ready"):
out.block_until_ready()
elapsed = (time.time() - t) / n_runs * 1000
print(f"{label}: {elapsed:.3f}ms / call")
import numpy as np
x_np = np.random.randn(1_000_000).astype(np.float32)
x_jax = jnp.array(x_np)
benchmark(lambda x: np.sum(x**2 + np.sin(x)), x_np, label="NumPy")
benchmark(jax.jit(lambda x: jnp.sum(x**2 + jnp.sin(x))), x_jax, label="JAX jit")
⚠️ 벤치마크의 함정
(1) 비동기 디스패치에서 block_until_ready를 빠뜨리면 실제보다 빠르게 보여. (2) 첫 실행 vs 재실행, 첫 호출은 컴파일을 포함해. (3) GPU의 lazy launch에서는 실제 GPU 작업이 호스트보다 한 발짝 늦게 시작해. (4) print가 측정 루프 안에 있으면 print가 동기화 지점이 돼. (5) 작은 배열 (< 1000 원소)는 오버헤드가 지배적이라 JAX가 NumPy보다 느릴 수도 있어.
실용 결론: 큰 배열, 무거운 연산, 여러 번 호출, 세 조건을 모두 만족하면 jit 효과가 극적. 한 번 호출 / 작은 데이터 / 단순 연산이면 jit 안 붙여도 돼. 그러나 학습 루프의 스텝 함수는 거의 항상 jit을 적용해.