가장 깊은 합성 가운데 하나는 그래디언트를 다시 미분하는 거야. 메타러닝, 하이퍼파라미터 최적화, 모델 불문 알고리즘에서 등장해.
가장 단순, 2차 미분
import jax
import jax.numpy as jnp
def f(x):
return x ** 4
# 1차: 4x³
print(jax.grad(f)(2.0)) # 32
# 2차: 12x²
print(jax.grad(jax.grad(f))(2.0)) # 48
# 3차: 24x
print(jax.grad(jax.grad(jax.grad(f)))(2.0)) # 48
각 grad가 새 함수를 돌려주므로 자유롭게 합성할 수 있어.
Hessian-벡터 product
큰 모델의 Hessian을 직접 계산하면 매개변수가 N개일 때 N²만큼 메모리가 필요해. 대신, Hessian과 벡터의 곱 (HVP)만 계산:
def hvp(f, x, v):
'''f 의 Hessian 과 vector v 의 곱'''
return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)
# 사용
def loss(x):
return jnp.sum(x ** 4)
x = jnp.array([1., 2., 3.])
v = jnp.array([1., 0., 0.])
print(hvp(loss, x, v)) # H @ v
Newton 방법, Conjugate Gradient, K-FAC 같은 2차 옵티마이저의 기본 재료야.
Meta-learning: MAML
모델 불문 메타러닝은 "여러 작업에서 빠르게 적응할 수 있는 초기 매개변수 찾기"를 목표로 해.
def task_loss(params, task_data):
x, y = task_data
pred = model_apply(params, x)
return jnp.mean((pred - y) ** 2)
def maml_inner_step(params, task, lr=0.1):
'''단일 task 에 대해 1 step 업데이트'''
grads = jax.grad(task_loss)(params, task)
return jax.tree.map(lambda p, g: p - lr * g, params, grads)
def maml_outer_loss(meta_params, tasks):
'''meta-learning objective'''
total = 0.0
for task in tasks:
# support set 에서 1 step 학습
adapted = maml_inner_step(meta_params, task["support"])
# query set 에서 평가 — 이게 meta-loss
total += task_loss(adapted, task["query"])
return total / len(tasks)
# meta-gradient — outer loss 의 meta_params 에 대한 gradient
# 안에 grad (inner) 가 있고, 그 위로 또 grad (outer) — 2 차 미분
@jax.jit
def meta_step(meta_params, tasks, meta_lr=0.001):
grads = jax.grad(maml_outer_loss)(meta_params, tasks)
return jax.tree.map(lambda p, g: p - meta_lr * g, meta_params, grads)
JAX가 이 grad-of-grad 패턴을 추가 코드 한 줄도 없이 자동 처리해. 다른 프레임워크에서는 2차 자동 미분을 명시적으로 작성하고 학습 루프와 그래디언트 훅을 바꿔야 할 수 있어. JAX는 jax.grad(jax.grad(...)).
암묵적 differentiation, 더 깊은 패턴
고정점의 그래디언트에는 별도 기법이 필요해:
# f(x*, theta) = 0 의 고정점 x* (theta 의 함수)
# implicit function theorem: dx*/dtheta = -(df/dx)^-1 (df/dtheta)
def implicit_solver(f, theta, x0):
'''f(x, theta) = 0 의 fixed point — 자동 미분 가능'''
# ... fixed-point 반복으로 x* 찾기
return x_star
# JAX 의 jaxopt 라이브러리가 이 패턴을 깔끔히 wrap
강화학습의 암묵적 보상 모델링, 반복 최적화, 하이퍼파라미터 튜닝도 모두 같은 핵심 기법을 사용해.
🌌 grad-of-grad의 의미
JAX가 메타러닝과 과학 계산에서 사랑받는 이유는 이런 자유로운 합성이 실제로 작동하기 때문이야. 박사 논문이었던 알고리즘이 학생이 하루 안에 실험할 수 있어. jax.grad(jax.grad(jax.grad(f)))가 그냥 작동하는 프레임워크, 다른 데서는 hack 또는 별도 라이브러리가 필요해.
주의, 깊게 중첩한 grad는 메모리 / 시간 비용이 커. 4차 미분쯤 가면 일반적으론 불필요해. 2차 (Hessian, MAML)까지가 흔한 실용적인 범위야.