본문 바로가기
C.W.K.
Stream
Lesson 02 of 06 · published

핵심 명제: 합성 가능한 함수 변환

~10 min · origins, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

대부분의 머신러닝 프레임워크는 객체를 중심으로 설계돼 있어. 모델 객체를 만들고 데이터를 전달한 다음 메서드를 호출하지. JAX는 다른 출발점을 택해. 핵심은 합성 가능한 함수 변환, 즉 함수를 입력으로 받아 변환된 새 함수를 반환하는 고차 함수야.

JAX의 네 가지 핵심 변환은 다음과 같아.

  • jax.jit, 컴파일: 함수를 추적한 뒤 XLA 컴파일러로 최적화된 머신 코드를 만든다. 첫 호출은 컴파일 때문에 느리지만 이후 호출은 훨씬 빨라진다.
  • jax.grad, 미분: 스칼라를 반환하는 함수를 받아 그래디언트를 계산하는 새 함수를 돌려준다.
  • jax.vmap, 벡터화: 예제 하나를 처리하는 함수를 배치 전체를 처리하는 함수로 바꾼다.
  • jax.pmap, 병렬화: 여러 장치, 즉 GPU나 TPU에서 함수를 동시에 실행한다.

여기서 가장 중요한 말은 합성 가능하다는 거야. 변환을 원하는 순서로 겹쳐 쓸 수 있어.

import jax
import jax.numpy as jnp

def loss_fn(params, x, y):
    predictions = jnp.dot(x, params)
    return jnp.mean((predictions - y) ** 2)

# 합성: 미분 → compile
fast_grad = jax.jit(jax.grad(loss_fn))

# 또는: per-example gradient → compile
per_example_grads = jax.jit(jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0)))

💡 왜 이게 중요한가

각 변환은 따로 써도 유용하지만 JAX의 진짜 힘은 합성에서 나와. jit(vmap(grad(f)))처럼 미분하고 벡터화한 함수를 다시 컴파일할 수 있어. PyTorch에서는 같은 결과를 위해 코드를 상당히 다시 구성해야 할 수 있지만 JAX에서는 변환을 한 줄로 조합해. 변환을 일급 값으로 다루는 이 설계가 미분 가능 프로그래밍의 중요한 전환점이야.

Code

import jax
import jax.numpy as jnp

def loss_fn(params, x, y):
    predictions = jnp.dot(x, params)
    return jnp.mean((predictions - y) ** 2)

# Compose transformations: differentiate, then compile
fast_grad = jax.jit(jax.grad(loss_fn))

# Or: vectorize per-example gradients, then compile
per_example_grads = jax.jit(jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0)))

params = jnp.array([1.0, 2.0])
x = jnp.array([[1.0, 0.5], [0.3, 0.8], [0.9, 0.1]])
y = jnp.array([1.5, 1.0, 0.8])

# Get compiled gradient
print(fast_grad(params, x, y))

# Get per-example gradients (one gradient vector per data point!)
print(per_example_grads(params, x, y))

External links

Exercise

loss_fn 예제에 jax.jit(jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0)))을 적용해. 배치 크기 1, 100, 10000에서 실행 시간을 측정해 봐. 합성 순서가 어떤 차이를 만드는지 확인하려고 jax.vmap(jax.jit(jax.grad(...)))도 실행해 비교해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.