Just-in-time compilation
Just-in-time compilation
JAX가 어떻게 동작하고, 어떻게 하면 빠르게 만들 수 있는지 살펴보는 단계예요. 핵심은 jax.jit() 변환인데, JAX 파이썬 함수를 Just In Time(JIT) 컴파일해서 XLA에서 효율적으로 실행하게 해 줘요. GPU·TPU에서도 거의 같은 코드로 동작하니까, 한 번 JIT의 원리를 이해해 두면 도움이 커요.
JAX 변환이 작동하는 방식
JAX는 각 함수를 primitive 연산의 시퀀스로 환원해 처리해요. 각 primitive는 계산의 한 기본 단위를 나타내죠. 이 과정에서 JAX는 각 인자를 트레이서(tracer) 객체로 감싸요. 트레이서는 함수 호출 중 수행되는 모든 JAX 연산을 기록해요. 그런 다음 JAX는 트레이서 기록을 사용해 전체 함수를 재구성하는데, 그 출력이 바로 jaxpr이에요.
핵심은 이렇습니다. 트레이서가 파이썬 부작용(side-effect)을 기록하지는 않으므로 jaxpr에 나타나지 않아요. 하지만 부작용은 트레이스 도중 실제로 일어나요. 또 jaxpr은 주어진 파라미터로 실행된 함수를 포착하므로, 파이썬 조건문이 있으면 실제로 취한 분기만 알게 돼요.
함수 JIT 컴파일
앞서 언급했듯 JAX는 같은 코드로 CPU/GPU/TPU에서 연산을 실행하게 해 줘요. 그런데 단순히 연산을 하나씩 가속기로 보내면 XLA 컴파일러가 함수를 최적화할 여지가 줄어들어요.
자연스럽게 XLA 컴파일러에게 가능한 많은 코드를 줘서 완전히 최적화하게 만들고 싶어지죠. 이를 위해 JAX는 jax.jit() 변환을 제공해요.
# Pre-compile the function before timing...
selu_jit(x).block_until_ready()
%timeit selu_jit(x).block_until_ready()
665 μs ± 3.05 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
여기서 일어나는 일이에요.
selu_jit을selu의 컴파일된 버전으로 정의해요.selu_jit을x로 한 번 호출해요. 여기가 JAX가 트레이싱하는 지점이에요. 트레이서로 감쌀 입력이 필요하니까요. 그런 다음 jaxpr이 XLA로 컴파일되어 GPU·TPU에 최적화된 매우 효율적인 코드가 돼요.- 마지막으로 컴파일된 코드가 실행되어 호출을 처리해요. 이후
selu_jit호출은 컴파일된 코드를 직접 사용해 파이썬 구현을 아예 건너뜁니다.
(따뜻한 호출을 별도로 하지 않으면 워밍업을 안 하므로 잘못된 벤치마크가 나와요. 벤치마크에 컴파일 시간이 포함되지만 루프를 여러 번 돌리니 더 빠르게 보이는데, 공정한 비교는 아니에요.)
block_until_ready()가 필요한 이유는 JAX의 비동기 디스패치(Asynchronous dispatch) 때문에 결과가 실제로 도착할 때까지 기다려야 하기 때문이에요.
왜 모든 것을 JIT 할 수 없을까
모든 함수에 jax.jit()을 적용하면 될 것 같지만 그렇지 않아요. JIT가 동작하지 않는 경우를 먼저 봐요.
# Condition on value of x.
def f(x):
if x > 0:
return x
else:
return 2 * x
jax.jit(f)(10) # Raises an error
TracerBoolConversionError: Attempted boolean conversion of traced array with shape bool[].
The error occurred while tracing the function f at /tmp/ipykernel_4143/2956679937.py:3 for jit. This concrete value was not available in Python because it depends on the value of the argument x.
# While loop conditioned on x and n.
def g(x, n):
i = 0
while i < n:
i += 1
return x + i
jax.jit(g)(10, 20) # Raises an error
두 경우 모두 런타임 값을 사용해 트레이스 시점의 프로그램 흐름을 통제하려 한 게 문제예요. JIT 안의 트레이스된 값(x처럼)은 shape·dtype 같은 정적 속성으로만 제어 흐름에 영향을 줄 수 있고, 값 자체로는 안 돼요.
그럴 땐 함수 일부만 JIT 컴파일하는 방법도 고려할 수 있어요. 예를 들어 계산 비용이 큰 부분이 루프 안이라면 그 안쪽만 JIT 컴파일할 수 있어요.
# While loop conditioned on x and n with a jitted body.
@jax.jit
def loop_body(prev_i):
return prev_i + 1
def g_inner_jitted(x, n):
i = 0
while i < n:
i = loop_body(i)
return x + i
g_inner_jitted(10, 20)
Array(30, dtype=int32, weak_type=True)
인자를 static 으로 표시하기
입력 값에 조건이 있는 함수를 정말 JIT 컴파일해야 한다면, static_argnums나 static_argnames를 지정해 특정 입력에 대해 덜 추상적인 트레이서를 쓰게 할 수 있어요. 대가로 결과 jaxpr과 컴파일 결과물이 전달된 값에 의존하게 되어, static 입력의 값이 바뀔 때마다 함수를 다시 컴파일해야 해요.
f_jit_correct = jax.jit(f, static_argnums=0)
print(f_jit_correct(10))
10
g_jit_correct = jax.jit(g, static_argnames=['n'])
print(g_jit_correct(10, 20))
30
데코레이터로 쓸 때는 데코레이터 팩토리 패턴을 사용해요.
@jax.jit(static_argnames=['n'])
def g_jit_decorated(x, n):
i = 0
while i < n:
i += 1
return x + i
print(g_jit_decorated(10, 20))
30
JIT과 캐싱
첫 JIT 호출의 컴파일 오버헤드를 고려하면, jax.jit()이 이전 컴파일을 언제 어떻게 캐시하는지 이해하는 게 효과적으로 쓰는 핵심이에요. f = jax.jit(g)라고 정의하면, f를 처음 호출할 때 컴파일되고 그 XLA 코드가 캐시돼요. 이후 f 호출은 캐시된 코드를 재사용해요. 이것이 jax.jit이 컴파일의 선행 비용을 상쇄하는 방식이에요.
static_argnums를 지정하면 캐시된 코드는 static으로 표시된 인자의 값이 같을 때만 재사용돼요. 값이 바뀌면 재컴파일이 일어나요.
주의할 두 가지가 있어요.
- 값이 여러 개라면 프로그램이 연산을 하나씩 실행하는 것보다 컴파일에 더 많은 시간을 쓸 수도 있어요.
- 루프나 다른 파이썬 스코프 안에서 정의된 임시 함수에
jax.jit()을 호출하지 마세요.partial이나lambda로 매번 다른 해시의 함수를 만들면 매번 불필요한 컴파일이 일어나요.
from functools import partial
def unjitted_loop_body(prev_i):
return prev_i + 1
def g_inner_jitted_partial(x, n):
i = 0
while i < n:
# Don't do this! each time the partial returns
# a function with different hash
i = jax.jit(partial(unjitted_loop_body))(i)
return x + i
def g_inner_jitted_lambda(x, n):
i = 0
while i < n:
# Don't do this!, lambda will also return
# a function with a different hash
i = jax.jit(lambda x: unjitted_loop_body(x))(i)
return x + i
def g_inner_jitted_normal(x, n):
i = 0
while i < n:
# this is OK, since JAX can find the
# cached, compiled function
i = jax.jit(unjitted_loop_body)(i)
return x + i
캐싱이 동작하는 g_inner_jitted_normal은 훨씬 빠르지만, partial·lambda 버전은 매 호출마다 재컴파일되어 느려져요.
더 알아보기
- Key concepts — JAX 변환의 기본
- Automatic differentiation — jax.grad 로 미분
- Automatic vectorization — jax.vmap 로 벡터화