Key concepts

Key concepts

JAX를 본격적으로 쓰려면 배열 연산 함수 외에, JAX 함수를 변형(transform)하는 기능들을 이해하는 게 핵심이에요. JAX의 강력함은 대부분 이 변환들에서 나오기 때문이에요. 이번 페이지에서는 JAX 패키지의 핵심 개념인 변환과 트레이싱(tracing)을 간단히 소개할게요.

출처: Key concepts · JAX Documentation

변환 (Transformations)

배열을 다루는 함수와 함께, JAX는 JAX 함수에 동작하는 여러 변환을 포함해요. 주요 변환은 이렇게 돼요.

그 외에도 여러 변환이 있어요. 변환은 함수를 인자로 받아 새로 변환된 함수를 반환해요. 예를 들어 간단한 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가 아니라, xshape·dtype 같은 필수 속성을 나타내는 Tracer 인스턴스예요. 트레이스된 값으로 함수를 실행해 JAX는 연산이 실제로 실행되기 전에 그 함수가 인코딩하는 일련의 연산을 결정할 수 있어요. 변환들은 바로 이렇게 함수의 구조를 파악해서 최적화·벡터화·미분을 수행해요.

핵심을 정리하면 이래요.

  • jit은 함수를 컴파일해서 가속기에서 효율적으로 실행하게 하고,
  • vmap은 함수를 자동 벡터화해 배치를 처리하게 하고,
  • grad는 함수의 그라디언트를 자동으로 계산해 줘요.

이 세 변환은 서로 자유롭게 조합할 수 있어서, 복잡한 고성능 학습 코드를 간결하게 표현할 수 있어요.

더 알아보기