Key concepts
Key concepts
JAX를 본격적으로 쓰려면 배열 연산 함수 외에, JAX 함수를 변형(transform)하는 기능들을 이해하는 게 핵심이에요. JAX의 강력함은 대부분 이 변환들에서 나오기 때문이에요. 이번 페이지에서는 JAX 패키지의 핵심 개념인 변환과 트레이싱(tracing)을 간단히 소개할게요.
변환 (Transformations)
배열을 다루는 함수와 함께, JAX는 JAX 함수에 동작하는 여러 변환을 포함해요. 주요 변환은 이렇게 돼요.
- jax.jit(): Just-In-Time(JIT) 컴파일
- jax.vmap(): 벡터화 변환
- jax.grad(): 그라디언트 변환
그 외에도 여러 변환이 있어요. 변환은 함수를 인자로 받아 새로 변환된 함수를 반환해요. 예를 들어 간단한 SELU 함수를 JIT 컴파일하는 방법은 이래요:
import jax
import jax.numpy as jnp
def selu(x, alpha=1.67, lambda_=1.05):
return lambda_ * jnp.where(x > 0, x, alpha * jnp.exp(x) - alpha)
selu_jit = jax.jit(selu)
print(selu_jit(1.0))
1.05
편의를 위해 파이썬 데코레이터 문법으로 변환을 적용하는 경우가 많아요:
@jax.jit
def selu(x, alpha=1.67, lambda_=1.05):
return lambda_ * jnp.where(x > 0, x, alpha * jnp.exp(x) - alpha)
트레이싱 (Tracing)
변환이 가능한 숨은 마법은 Tracer 개념이에요. 트레이서는 배열 객체의 추상적 대체물로, 함수가 인코딩하는 일련의 연산을 추출하기 위해 JAX 함수에 전달돼요.
변환된 JAX 코드 안에서 어떤 배열 값을 print 해 보면 이를 직접 볼 수 있어요. 예를 들어:
@jax.jit
def f(x):
print(x)
return x + 1
x = jnp.arange(5)
result = f(x)
JitTracer(int32[5])
출력된 값은 배열 x가 아니라, x의 shape·dtype 같은 필수 속성을 나타내는 Tracer 인스턴스예요. 트레이스된 값으로 함수를 실행해 JAX는 연산이 실제로 실행되기 전에 그 함수가 인코딩하는 일련의 연산을 결정할 수 있어요. 변환들은 바로 이렇게 함수의 구조를 파악해서 최적화·벡터화·미분을 수행해요.
핵심을 정리하면 이래요.
- jit은 함수를 컴파일해서 가속기에서 효율적으로 실행하게 하고,
- vmap은 함수를 자동 벡터화해 배치를 처리하게 하고,
- grad는 함수의 그라디언트를 자동으로 계산해 줘요.
이 세 변환은 서로 자유롭게 조합할 수 있어서, 복잡한 고성능 학습 코드를 간결하게 표현할 수 있어요.
더 알아보기
- Quickstart: How to think in JAX — JAX 시작하기
- Just-in-time compilation — jax.jit 상세
- Automatic differentiation — jax.grad 상세
- Automatic vectorization — jax.vmap 상세