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

mx.grad와 함수 변환 가족

~16 min · autograd, mx.grad, vmap, value-and-grad

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

기록 테이프 대신 함수를 조합해

PyTorch의 자동 미분은 연산하는 동안 테이프를 기록한 뒤 그 길을 거꾸로 따라가. MLX와 JAX는 다른 길을 골랐어. 함수를 받아 새 함수를 내놓는 변환을 제공해. mx.grad(f)는 나중에 f의 기록을 미분하는 게 아니라, 호출할 때 기울기를 계산하는 새 함수를 돌려줘.

사소한 구분이 아니야. 기울기가 다른 변환과 자유롭게 조합되는 일급 값이 된다는 뜻이거든. 다른 변환으로 감싸고, JIT 컴파일하고, 배치에 vmap을 적용할 수 있어. 특별한 "학습 모드"도 살아 있는 테이프도 필요 없어.

가장 자주 쓰는 셋

mx.grad(f) — 스칼라를 내놓는 함수를 받아 첫 번째 인자의 기울기를 계산하는 함수를 돌려줘. 손실값은 필요 없고 기울기만 원할 때 써.

mx.value_and_grad(f) — 한 번의 호출로 손실값과 기울기를 함께 돌려줘. 학습 반복문에서는 거의 항상 손실도 기록하니 이쪽을 쓰게 될 거야.

mx.vmap(f) — 예제 하나를 처리하는 함수를 배치 함수로 바꿔. grad와 자연스럽게 조합되고, 직접 브로드캐스팅을 짜지 않아도 효율적인 배치 계산을 할 수 있어.

핵심은 조합할 수 있다는 것

각 변환이 함수를 받아 함수를 돌려주니 층층이 쌓을 수 있어. mx.vmap(mx.grad(loss_fn))은 배치의 예제별 기울기를, mx.grad(mx.grad(f))는 2차 미분을 줘. mx.compile(mx.value_and_grad(loss_fn))은 기울기 계산을 JIT 컴파일해. 다른 프레임워크에서도 가능한 기술이지만 MLX는 그 구조를 API에 훤히 드러내.

nn.value_and_grad는 모델을 알아

매개변수를 가진 nn.Module에서는 보통 함수의 첫 인자가 아니라 모델 매개변수의 기울기가 필요해. nn 모듈은 학습의 표준 방식인 nn.value_and_grad(model, loss_fn)을 제공해. 레슨 6에서 바로 쓸 거야.

Code

mx.grad — 한 줄짜리·python
import mlx.core as mx

def square(x):
    return x ** 2

grad_sq = mx.grad(square)
print(grad_sq(mx.array(3.0)))   # → array(6, dtype=float32)   d/dx(x^2) = 2x = 6
value_and_grad — 학습 반복문의 기본 재료·python
import mlx.core as mx

def loss(w, x, y):
    pred = w * x
    return ((pred - y) ** 2).mean()

vg = mx.value_and_grad(loss)

w0 = mx.array(1.5)
xs = mx.array([1.0, 2.0, 3.0, 4.0])
ys = mx.array([2.0, 4.0, 6.0, 8.0])      # true relation y = 2x

v, g = vg(w0, xs, ys)
print('loss value:', float(v))            # → 1.875
print('gradient w.r.t. w:', float(g))     # → -7.5  (you'd subtract this in SGD)
vmap — 배치 함수 호출·python
import mlx.core as mx

def f(x):
    return x ** 2 + 1

# Vectorize f over the leading axis
batched = mx.vmap(f)

xb = mx.array([1.0, 2.0, 3.0, 4.0])
print(batched(xb))   # → array([2, 5, 10, 17], dtype=float32)

# Composes with grad: per-example gradient across a batch
batched_grad = mx.vmap(mx.grad(f))
print(batched_grad(xb))   # → array([2, 4, 6, 8], dtype=float32)   (= 2x for each)

External links

Exercise

인자 두 개를 받는 작은 스칼라 함수를 골라. 예를 들면 def f(x, y): return (x - y) ** 2 + x야. mx.grad(f, argnums=0)mx.grad(f, argnums=1)xy의 편미분을 각각 구하고 손으로 계산한 답과 맞춰봐. 이어서 (x, y) 쌍의 배치에 mx.vmap을 조합해. argnums와 함수 조합만으로 별도 역전파 코드를 짜지 않고도 완전히 통제할 수 있다는 감각을 잡는 게 목표야.

Progress

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

댓글 0

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

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