함수형 autograd와 자동 배치 변환
torch.func는 예전의 독립 패키지 functorch를 PyTorch 안으로 통합한 기능으로, JAX와 비슷한 함수 변환을 제공해. grad, vmap, jacrev, hessian을 사용할 수 있고, 단일 샘플에 동작하는 함수를 배치 전체에 자동으로 벡터화할 수 있어.
실전에서 의외로 자주 만나는 용도는 두 가지야:
- 샘플별 기울기. 표준
.backward()는 배치에서 합친 손실의 기울기만 줘. 즉, 매개변수마다 기울기가 하나씩 나와. 차등 개인정보 보호, 영향 함수, GradSAM 같은 연구에서는 각 샘플의 손실에 대한 기울기가 필요해.torch.func.vmap(grad(...))를 쓰면 Python 반복문 없이 구할 수 있어. - 고차 기울기. 헤시안-벡터 곱, 2차 최적화, 메타 학습은 모두 기울기를 다시 미분해야 해.
torch.func.grad(grad(f))처럼 변환을 깔끔하게 조합할 수 있어.
관점을 바꿔 보기
표준 PyTorch에서는 텐서가 암묵적인 그래프 상태를 들고 있고 여기에 .backward()를 호출해. torch.func에서는 입력과 매개변수를 받는 순수 함수를 만들고, 변환이 또 다른 순수 함수를 만들어 내. JAX에 조금 더 가까운 관점이라 일반적인 학습에서는 다소 낯설 수 있지만, 앞의 두 경우에는 훨씬 강력해.