기록 테이프 대신 함수를 조합해
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에서 바로 쓸 거야.