Automatic differentiation

Automatic differentiation

현대 머신러닝에서 그라디언트를 계산하는 건 절대 뺄 수 없는 핵심이에요. JAX는 꽤 범용적인 자동 미분(autodiff) 시스템을 갖추고 있어서, jax.grad() 같은 변환으로 어떤 함수의 미분도 쉽게 구할 수 있어요. 이번 페이지에서는 JAX로 그라디언트를 계산하는 기본적인 방법들을 소개할게요.

출처: Automatic differentiation · JAX Documentation

1. jax.grad로 그라디언트 구하기

JAX에서는 스칼라 값 함수를 jax.grad() 변환으로 미분할 수 있어요:

import jax
import jax.numpy as jnp
from jax import grad

jax.grad()는 함수를 받아 함수를 반환해요. 수학 함수 \(f\)를 평가하는 파이썬 함수 f가 있다면, jax.grad(f)는 그라디언트 \(\nabla f\)를 평가하는 파이썬 함수예요. 즉 grad(f)(x)가 \(\nabla f(x)\)의 값이에요.

jax.grad()는 함수에 작동하므로, 그 출력에 다시 적용해 원하는 만큼 여러 번 미분할 수 있어요. 미분을 계산하는 함수도 미분 가능하므로, 고차 도함수도 변환을 쌓는 것만큼 쉬워요.

예를 들어 \(f(x) = x^3 + 2x^2 - 3x + 1\)의 도함수는:

f = lambda x: x**3 + 2*x**2 - 3*x + 1
dfdx = jax.grad(f)

고차 도함수도 jax.grad를 연쇄하면 쉽게 구할 수 있어요:

d2fdx = jax.grad(dfdx)
d3fdx = jax.grad(d2fdx)
d4fdx = jax.grad(d3fdx)

2. 선형 로지스틱 회귀에서 그라디언트 구하기

다음 예시는 선형 로지스틱 회귀 모델에서 jax.grad()로 그라디언트를 계산하는 방법을 보여줘요. 랜덤 모델 계수를 초기화하고, jax.grad()argnums 인자로 위치 인자(positional argument)에 대해 미분할 수 있어요. argnums를 쓰면, f가 수학 함수 \(f\)를 평가하는 파이썬 함수일 때 jax.grad(f, i)는 \(\partial_i f\)를 평가하는 파이썬 함수가 돼요.

3. 중첩 리스트·튜플·딕셔너리에 대해 미분하기

JAX의 PyTree 추상화 덕분에, 표준 파이썬 컨테이너에 대해서도 미분이 자연스럽게 동작해요. jax.grad()뿐 아니라 jax.jit(), jax.vmap() 같은 다른 JAX 변환들도 마찬가지예요.

4. jax.value_and_grad로 함수와 그라디언트 함께 평가

jax.value_and_grad()를 사용하면 함수의 값과 그라디언트를 한 번에 계산할 수 있어요.

import jax
val, grad = jax.value_and_grad(f)(x)

5. 수치 미분으로 확인하기

계산된 그라디언트가 맞는지 수치 미분(유한 차분)과 비교해 확인할 수 있어요. 랜덤 방향으로 유한 차분과 비교해 W_grad를 검산하죠.

정리

  • jax.grad는 함수 → 함수 변환이라 그라디언트 계산을 코드로 자연스럽게 표현해요.
  • 고차 도함수는 grad를 중첩해 구할 수 있어요.
  • PyTree 덕분에 파이썬 컨테이너 단위로도 미분돼요.
  • value_and_grad로 값과 그라디언트를 함께 얻을 수 있어요.

더 알아보기