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

Meta-Learning: Grad-of-Grad 패턴

~9 min · advanced, jax, tutorial

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

가장 깊은 합성 가운데 하나는 그래디언트를 다시 미분하는 거야. 메타러닝, 하이퍼파라미터 최적화, 모델 불문 알고리즘에서 등장해.

가장 단순, 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)까지가 흔한 실용적인 범위야.

Code

import jax
import jax.numpy as jnp

# Second-order derivatives are trivial
def f(x):
    return jnp.sin(x) * x ** 2

# First derivative
df = jax.grad(f)
print(df(1.0))  # cos(1)*1 + sin(1)*2 ≈ 2.22

# Second derivative (Hessian for scalar functions)
d2f = jax.grad(jax.grad(f))
print(d2f(1.0))  # ≈ -0.18

# Hessian for vector functions
def g(x):
    return jnp.sum(x ** 3)

hessian = jax.hessian(g)
print(hessian(jnp.array([1.0, 2.0, 3.0])))
# [[6., 0., 0.],
#  [0., 12., 0.],
#  [0., 0., 18.]]
def maml_loss(meta_params, tasks, inner_lr=0.01, inner_steps=5):
    """MAML outer loss: how well do inner-loop-adapted params perform?"""
    total_loss = 0.0

    for task_train, task_test in tasks:
        # Inner loop: adapt to task using gradient descent
        params = meta_params
        for _ in range(inner_steps):
            train_loss = compute_loss(params, *task_train)
            grads = jax.grad(compute_loss)(params, *task_train)
            params = jax.tree.map(
                lambda p, g: p - inner_lr * g, params, grads)

        # Outer loss: evaluate adapted params on test data
        test_loss = compute_loss(params, *task_test)
        total_loss += test_loss

    return total_loss / len(tasks)

# The magic: differentiate through the inner loop!
meta_grads = jax.grad(maml_loss)(meta_params, tasks)
# This computes second-order gradients (gradient through gradient descent)
# In PyTorch, this requires create_graph=True and careful management.
# In JAX, it just works — grad(grad(...)) composes naturally.
# Practical example: memory-efficient Transformer with remat + scan
def create_efficient_transformer(d_model, num_heads, d_ff, num_layers, rngs):
    """Create a transformer that uses scan + remat for efficiency."""
    # Initialize one block's parameters
    def init_block(rngs):
        return TransformerBlock(d_model, num_heads, d_ff, rngs)

    blocks = [init_block(rngs) for _ in range(num_layers)]

    def forward(x):
        for block in blocks:
            # Remat: recompute activations in backward pass (saves memory)
            x = jax.checkpoint(lambda b, x: b(x), block, x)
        return x

    return forward, blocks

External links

Exercise

MAML의 내부 갱신을 30줄 안에 구현해. 내부 SGD 한 스텝을 통과하는 메타 그래디언트를 구하고 두 작업으로 이루어진 합성 문제에서 검증해. grad를 다시 grad하는 패턴이 왜 이 퀘스트의 개념적 정점인지 설명해.

Progress

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

댓글 0

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

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