70억 파라미터 모델에서 loss.backward() 한 줄은 순전파와 비슷한 시간에 끝납니다. 대략 두 배 정도입니다.
그런데 지난 글 끝에서 봤듯, 그 미분을 「층마다 야코비안을 구해 곱한다」로 곧이곧대로 하면 곱셈이 번대로 불어납니다. 두 배가 아니라 수천 배입니다. 같은 답을 구하는데 왜 이렇게 차이가 날까요.
답은 곱하는 순서 하나입니다. 이 글은 그 순서에 이름을 붙이고, 비용을 실제로 세어 딥러닝이 반드시 뒤에서부터 미분해야 한다는 것을 계산으로 확인합니다.
곱하는 순서와 비용
결합법칙과 비용
지난 글에서 층 개짜리 신경망의 전체 야코비안이
이라는 것을 봤습니다. 행렬곱에는 결합법칙이 성립하므로 어느 쪽부터 곱해도 답은 같습니다. 그런데 비용은 같지 않습니다.
폭이 4096인 층이 100개인 모델로 세어 봅니다. 각 는 입니다.
행렬과 행렬을 곱하는 데 드는 곱셈은 번입니다. 이 사실 하나로 두 순서를 비교합니다.
오른쪽부터 곱하면 매번 짜리 둘을 곱합니다.
왼쪽부터 곱하면 사정이 다릅니다. 손실은 수 하나이므로 맨 왼쪽은 짜리 행 하나입니다. 행 하나에 행렬을 곱하면
이고, 게다가 중간 결과가 계속 짜리 행이라 메모리도 거의 들지 않습니다.
차이는 정확히 4096배이고, 그것은 우연이 아니라 「 대 」라는 지수 차이입니다. 중간 폭이 넓을수록 벌어집니다.
왼쪽부터 곱하는 쪽에서는 야코비안 행렬이 한 번도 통째로 만들어지지 않습니다. 매 걸음에서 나오는 것은 벡터 하나뿐입니다. 이것이 autograd가 하는 일이고, 다음 절에서 그 「벡터 하나」에 이름을 붙입니다.
역전파의 두 배
그렇다면 처음의 「두 배」는 어디서 나올까요. 왼쪽부터 곱하면 곱셈이 순전파와 같은 자릿수여야 할 것 같은데, 실제로는 순전파의 두 배쯤 듭니다. 이유는 층 하나를 되돌아갈 때 해야 하는 곱이 하나가 아니라 둘이기 때문입니다.
선형층 에서 가 이면 순전파의 곱셈은 번입니다. 되돌아올 때는 위에서 내려온 행 (손실을 로 미분한 값)를 받아 두 가지를 만들어야 합니다.
- 아래층으로 넘길 것 — 손실을 로 미분한 . 곱셈 번.
- 이 층이 가져갈 것 — 손실을 로 미분한 . 크기가 인 바깥곱이라 역시 곱셈 번.
앞쪽은 신호를 계속 전달하기 위한 곱이고, 뒤쪽은 파라미터를 고치기 위한 곱입니다. 역전파의 목적은 뒤쪽인데, 뒤쪽을 층마다 얻으려면 앞쪽이 층을 건너 내려와야 합니다. 그래서 곱셈이 두 벌이고 순전파의 두 배입니다. 순전파 1에 되돌아오는 2를 더해 학습 한 걸음이 순전파의 세 배쯤 든다는 경험칙이 여기서 나옵니다. 배치가 개면 세 항 모두 가 곱해질 뿐 비율은 그대로입니다.
층마다 다른 배율
두 배는 선형층의 셈이고 모델 전체의 비율은 층 구성에 따라 움직입니다. 대략 1.5배에서 3배 사이입니다.
- 두 곱 중 하나가 필요 없는 층은 한 배로 내려갑니다. 첫 층의 입력은 데이터라 를 만들 이유가 없고, 가중치를 얼려 둔 층(일부만 파인튜닝할 때)은 를 건너뜁니다. 이런 층이 많을수록 전체 비율이 1.5배 쪽으로 내려옵니다.
- 원소별 활성함수는 곱이 원래 원소 수만큼이라 순전파든 역전파든 가볍습니다. 비율을 거의 안 건드립니다.
- 되돌아올 때 순전파를 다시 해야 하는 설정에서는 올라갑니다. 아래 「손실이 스칼라라는 조건」 절의 체크포인팅이 그것이고, 되돌아오는 쪽에 순전파 한 번이 더 얹혀 3배에 가까워집니다.
그러니 「역전파는 순전파의 두 배」는 법칙이 아니라 선형층이 대부분인 모델의 평균입니다. 중요한 것은 배율이 몇이든 한 자릿수에 머문다는 점이고, 그것을 가능하게 하는 것이 곱하는 순서입니다.
JVP와 VJP의 정의
두 가지 곱
야코비안 가 일 때, 벡터를 곱할 수 있는 자리가 양쪽으로 둘입니다.
정의. 에 대해 을 야코비안-벡터 곱(JVP, Jacobian-vector product)이라 하고, 에 대해 을 벡터-야코비안 곱(VJP, vector-Jacobian product)이라 한다.
앞 절에서 선형층을 되돌아가며 만든 가 바로 VJP입니다. 모양만 보면 두 곱은 대칭이지만 뜻은 정반대입니다.
- JVP는 앞을 봅니다. 는 입력 쪽에서 잡은 방향이고, 는 「입력을 방향으로 흔들면 출력 전체가 어느 방향으로 얼마나 움직이는가」입니다.
- VJP는 뒤를 봅니다. 는 출력 쪽에서 잡은 가중치이고, 는 「출력들을 로 저울질한 값이 각 입력에 얼마나 민감한가」입니다.
지난 글의 예로 확인합니다. 에 를 이은 합성의 야코비안이 점 에서
였습니다. 로 JVP를, 로 VJP를 계산하면
입니다. 여기까지는 를 만들어 놓고 곱한 것이라 아직 절약이 없습니다.
이중수
JVP는 야코비안을 몰라도 계산할 수 있습니다. 수 하나에 꼬리를 하나 달아 꼴로 들고 다니되, 은 0이 아니면서 인 기호로 약속합니다. 이렇게 앞자리 와 꼬리 를 짝으로 묶은 수를 이중수(dual number)라고 합니다. 앞자리가 값이고 꼬리가 그 값의 미분입니다.
왜 꼬리가 미분이 되는지는 곱셈 하나로 보입니다.
꼬리 자리에 곱의 미분법 가 저절로 나타났습니다. 이 2차 이상의 항을 모두 지우므로, 어떤 다항식에 를 넣어도 결과는 입니다. 테일러 전개를 1차에서 자른 것이 정확한 등식이 되는 셈입니다.
점 에서 방향 의 JVP를 이 방식으로 구해 봅니다. 입력마다 꼬리에 의 성분을 달아 , 으로 둡니다.
를 지난 값이 이고, 꼬리 이 곧 입니다. 이것을 그대로 에 넣으면
이 나옵니다. 앞자리 는 합성함수의 값이고 꼬리 는 위에서 를 만들어 얻은 와 같습니다. 값과 미분이 순전파 한 번에 함께 흘러간 것이고, 야코비안은 어디에서도 만들어지지 않았습니다. 연산마다 더 하는 일은 꼬리의 곱셈 몇 개뿐이라 비용도 순전파의 몇 배 안쪽입니다.
방향 여러 개
꼬리를 하나가 아니라 여럿 달 수도 있습니다. , 처럼 꼬리를 길이 2짜리 벡터로 두면, 곱셈 규칙은 그대로이고 꼬리의 두 칸이 각자 따로 굴러갑니다. 순전파 한 번이 끝나면 첫 칸에는 의 JVP인 , 둘째 칸에는 의 JVP인 이 들어 있습니다. 의 두 열이 한꺼번에 채워진 것입니다.
값의 계산은 한 번으로 공유되고 꼬리 계산만 방향 수만큼 늘어나므로, 방향 개를 따로 번 굴리는 것보다 쌉니다. JAX 같은 프레임워크가 여러 방향의 JVP를 한 번에 묶어 도는 것이 이 방식이고, GPU에서는 꼬리의 칸들이 배치 축처럼 나란히 계산됩니다. 다만 셈의 총량은 여전히 방향 수에 비례한다는 사실은 바뀌지 않습니다. 이 점이 뒤의 「전체 야코비안까지의 횟수」에서 비용을 가릅니다.
층으로 나눈 곱
층별 VJP
절약은 곱을 층마다 나눠서 하는 데서 나옵니다. 이므로
로 적을 수 있습니다. 괄호 안이 먼저 끝나고, 그 결과가 벡터 하나라 다음 층에 넘길 것도 벡터 하나입니다.
VJP 쪽을 손으로 따라가 봅니다. , 이므로
앞에서 를 만들어 얻은 값과 같습니다. JVP도 마찬가지로 을 먼저 얻고 로 끝납니다. 이중수로 얻은 꼬리와 똑같은 수입니다.
여기서 두 층의 야코비안을 곱한 적이 한 번도 없습니다. 매번 「행렬 하나 × 벡터 하나」였습니다.
그리고 층에 따라서는 그 행렬조차 만들 필요가 없습니다. 자주 쓰는 층들의 VJP를 보면 대부분 행렬을 건드리지 않습니다.
| 층 | 야코비안 | VJP | 비용 |
|---|---|---|---|
| · 행렬-벡터 곱 하나 | |||
| · 원소별 곱 하나 | |||
| · 그대로 넘김 | |||
| 의 VJP + 덧셈 |
원소별 활성함수 줄이 특히 극적입니다. 야코비안은 짜리 대각행렬이라 성분이 개인데, VJP는 길이 짜리 원소별 곱 한 번으로 끝납니다. 폭이 4096이면 1,677만 개짜리 행렬을 만드는 대신 4,096번 곱하면 됩니다.
그래서 자동미분 프레임워크에서 각 연산이 등록하는 것은 야코비안이 아니라 VJP 함수 하나입니다. 「받은 벡터를 어떻게 바꿔 아래로 넘길 것인가」만 알면 되고, 행렬은 어디에도 나타나지 않습니다. 지난 글에서 봤던 「값이 곱해지려면 순전파 값이 필요하다」는 사실은 여기서도 그대로라, 를 계산하려면 를 들고 있어야 합니다.
갈래의 합
지금까지의 그래프는 한 줄로 이어진 사슬이었습니다. 실제 모델에서는 한 값이 여러 곳에 쓰입니다. 잔차 연결에서 는 로도 들어가고 덧셈으로도 곧장 건너가며, 어텐션에서는 같은 입력이 질의·키·값 세 곳으로 갈라집니다. 이렇게 값 하나가 둘 이상으로 갈라지는 자리에서는 되돌아오는 VJP들을 더합니다.
작은 예로 봅니다. , 으로 갈라졌다가 로 다시 모이는 그래프입니다. 이므로 답은 이고, 이면 54입니다. 되돌아오며 세면
- 의 VJP는 쪽으로 , 쪽으로 을 보냅니다.
- 갈래는 의 VJP를 지나 이 됩니다.
- 갈래는 의 VJP를 지나 이 됩니다.
에 도착한 둘을 더하면 입니다. 이 덧셈은 연쇄법칙의 다변수 꼴 를 그대로 옮긴 것입니다. 갈라지는 자리를 행렬로 적으면 라는 복사이고, 그 야코비안 의 VJP가 , 곧 합입니다. 앞으로 복사하면 뒤로는 더한다고 기억하면 됩니다.
PyTorch는 이 합을 텐서의 .grad 칸에 더하는 방식으로 구현합니다. 갈래마다 도착한 VJP를 차례로 += 하므로 그래프 안에서는 저절로 합이 맞습니다. 그런데 이 칸은 backward() 호출이 끝나도 비워지지 않아서, 다음 배치에서 다시 부르면 지난 배치의 그래디언트 위에 더해집니다. 학습 루프마다 optimizer.zero_grad()를 부르는 이유가 이것입니다. 거꾸로 이 성질을 일부러 쓰는 것이 그래디언트 누적입니다. 메모리에 안 들어가는 큰 배치를 작은 배치 넷으로 나눠 backward()를 네 번 부르고 한 번만 파라미터를 고치면, 큰 배치 하나의 그래디언트를 얻습니다.
브로드캐스팅의 VJP
갈래의 합은 눈에 덜 띄는 곳에도 숨어 있습니다. 선형층의 편향 는 길이 짜리 벡터인데, 배치 줄짜리 출력 에 더할 때는 모든 줄에 같은 가 더해집니다. 모양이 작은 쪽을 큰 쪽에 맞춰 복사해 연산하는 이 규칙을 브로드캐스팅이라 합니다.
브로드캐스팅은 곧 복사이므로, 그 VJP는 복사한 축을 따라 더하는 것입니다. 위에서 내려온 가 짜리
이면 편향의 그래디언트는 세 줄을 세로로 더한 입니다. 편향 하나가 배치의 모든 예제에 똑같이 영향을 주었으니 모든 예제에서 돌아온 민감도를 모으는 것입니다. 편향의 그래디언트가 배치 축의 합이라는 것, 그래서 손실을 배치 평균으로 잡느냐 합으로 잡느냐에 따라 그 크기가 배치 크기만큼 달라진다는 것이 여기서 나옵니다. 직접 브로드캐스팅을 짜는 코드에서 역방향을 손으로 쓸 때 이 합을 빠뜨리면 모양이 로 남아 오류가 나거나, 더 나쁘게는 다른 곳에서 브로드캐스팅되어 조용히 틀립니다.
전체 야코비안까지의 횟수
전방 모드와 역방향 모드
이제 비용을 셉니다. 의 야코비안은 짜리 표입니다.
JVP 한 번은 그 표에서 열 방향의 정보 하나를 줍니다. (그 자리만 1인 벡터)로 잡으면 는 정확히 번째 열입니다. 그러니 표를 다 채우려면 열 개수만큼, 즉 번 굴려야 합니다. 이중수의 꼬리를 여러 칸 달아 한 번에 돌려도 꼬리 칸이 개 필요한 것은 같습니다.
VJP 한 번은 행 하나를 줍니다. 면 가 번째 행이고, 표를 다 채우려면 번입니다.
두 방식에 이름이 붙어 있습니다. JVP를 굴리는 쪽이 전방 모드(forward mode) 자동미분이고, VJP를 굴리는 쪽이 역방향 모드(reverse mode)입니다. 역전파는 역방향 모드를 신경망에 적용한 것의 이름입니다. 굴리기 한 번의 비용은 둘 다 순전파 한 번과 같은 자릿수입니다 — 전방 모드는 값과 미분을 함께 밀고 가고, 역방향 모드는 값을 한 번 밀어 놓은 뒤 되돌아옵니다.
| 한 번에 얻는 것 | 전체 야코비안까지 | 메모리 | |
|---|---|---|---|
| 전방 모드 (JVP) | 열 하나 | 번 | 순전파와 같음 |
| 역방향 모드 (VJP) | 행 하나 | 번 | 중간값을 전부 보관 |
규칙은 한 줄입니다 — 좁은 쪽을 훑는다. 입력이 적으면( 작음) 전방이 싸고, 출력이 적으면( 작음) 역방향이 쌉니다.
헤시안-벡터 곱
두 모드를 겹칠 수도 있습니다. 손실 의 그래디언트 은 그 자체로 인 벡터함수이고, 이 함수의 야코비안, 곧 2차 편미분 를 모은 표를 헤시안이라 합니다. 곡률을 쓰는 최적화나 손실 지형의 뾰족함을 재는 분석은 헤시안이 필요한데, 이면 성분이 개라 만들 수가 없습니다.
필요한 것이 헤시안 전체가 아니라 벡터 하나와의 곱 라면 이야기가 다릅니다. 는 「그래디언트 함수의 JVP」이므로, 안쪽에서 역방향 모드로 그래디언트를 계산하는 과정을 바깥에서 이중수로 밀면 됩니다. 이 조합을 흔히 forward-over-reverse라 부릅니다.
로 확인합니다. 그래디언트는 이고 점 에서 헤시안은
입니다. 그래디언트 식에 이중수 , 을 넣으면 , 이라 꼬리가 로 같습니다. 헤시안의 성분 넷을 한 번도 적지 않고 곱만 얻은 것이고, 비용은 그래디언트 한 번의 몇 배 안쪽입니다. 켤레 기울기법처럼 만 반복해서 쓰는 알고리즘이 거대한 모델에서도 돌아가는 이유가 이것입니다.
정사각 야코비안
이면 두 모드가 전체 야코비안을 얻는 데 드는 횟수가 똑같이 번입니다. 연산량으로는 가를 수 없고, 이때 갈리는 것은 메모리입니다.
전방 모드는 값과 꼬리를 앞으로 밀기만 하므로, 한 층을 지나면 그 층의 중간값을 버려도 됩니다. 메모리가 층 수와 무관하게 순전파 한 번 분량입니다. 역방향 모드는 되돌아올 때 층마다 순전파 값을 다시 읽어야 하므로 층 개의 중간값을 모두 들고 있어야 합니다. 층이 깊을수록 역방향 쪽 메모리가 층 수에 비례해 커집니다.
입력과 출력의 차원이 같도록 설계하는 정규화 흐름(normalizing flow) 같은 모델에서 이 구별이 실제로 문제가 됩니다. 다만 거기서도 전체 야코비안을 번 굴려 만드는 일은 피하고, 행렬식이 대각 성분의 곱으로 떨어지는 삼각 구조의 층을 골라 횟수 자체를 없앱니다.
손실이 스칼라라는 조건
스칼라 손실
학습의 미분이 어떤 모양인지 적어 봅니다. 손실은 파라미터 전체를 받아 수 하나를 내놓습니다.
입니다. 야코비안은 행이 하나뿐인 짜리 표이고, 역방향 모드로는 한 번에 끝납니다. 로 시작해 VJP를 층마다 넘기면 그 한 줄이 전부 채워지고, 그것이 그래디언트입니다.
같은 것을 전방 모드로 하면 열이 70억 개이므로 70억 번 굴려야 합니다. 순전파 한 번이 0.1초라면
입니다. 역방향 모드로 0.2초 남짓에 끝나는 일입니다.
「손실이 스칼라다」라는 한 문장이 역전파를 강제한 셈입니다. 역전파가 좋은 알고리즘이라서가 아니라, 학습이라는 문제의 모양이 이기 때문입니다.
반대 모양
그래서 반대 모양에서는 결론이 뒤집힙니다.
- 입력 하나를 흔들었을 때 출력 전체가 어떻게 되는지 보는 민감도 분석은 이 작으므로 전방 모드가 낫습니다.
- 야코비안에 벡터 하나를 곱하는 것만 필요한 경우도 전방 모드가 자연스럽습니다. 앞 절의 헤시안-벡터 곱이 그 예입니다.
- 야코비안을 진짜로 통째로 원하는 경우(변수변환의 로그행렬식 같은 것)만 번을 각오합니다. 지난 글에서 「행렬식을 싸게 계산할 수 있는 층을 고른다」고 했던 것이 이 비용을 피하려는 설계입니다.
다중 목적 학습
손실이 늘 하나인 것은 아닙니다. 번역 품질과 길이 제약을 함께 맞추거나, 분류 손실과 보조 손실을 같이 두는 다중 목적 학습에서는 손실이 으로 여럿입니다. 그러면 파라미터에서 손실로 가는 함수가 이 되어 야코비안의 행이 개로 늘고, 손실마다 그래디언트를 따로 알고 싶으면 역방향을 번 굴려야 합니다.
대부분의 코드는 그러지 않습니다. 가중치 를 정해 라는 스칼라를 먼저 만들고 backward()를 한 번만 부릅니다. 이것이 정확히 VJP 하나라는 점이 핵심입니다. 손실 벡터의 야코비안 에 을 곱한 는 행들의 가중합 이고, 그것이 가중합 손실의 그래디언트와 같습니다. 가중합을 먼저 만드는 관행은 편의가 아니라 번을 한 번으로 줄이는 계산입니다. 대신 손실끼리의 그래디언트가 부딪히는지 보려는 기법들은 행을 따로 봐야 하므로 이 절약을 포기하고 번 비용을 냅니다.
그래디언트 체크포인팅
역방향 모드가 치르는 값은 메모리입니다. 되돌아올 때 각 층의 순전파 값이 필요하므로 중간값을 전부 보관해야 하고, 활성화 메모리가 배치 크기와 층 수에 비례해 늘어납니다. 값을 일부 버렸다가 되돌아올 때 다시 계산하는 방법을 그래디언트 체크포인팅이라 합니다. 메모리를 계산으로 되사는 거래입니다.
층이 개일 때 개씩 묶어 구간 개로 나눕니다. 순전파에서는 구간 경계의 값만 남기고 나머지는 버립니다. 되돌아올 때는 한 구간씩, 그 구간 경계에서 순전파를 다시 해 안쪽 값을 복원하고 역방향을 지나간 뒤 버립니다. 한순간 들고 있는 것은 경계 개와 복원 중인 구간 하나의 개라 메모리가 입니다. 이면 100개 대신 경계 10개와 구간 안 10개, 모두 20개입니다.
계산은 얼마나 늘까요. 모든 구간을 한 번씩 다시 도는 것이 순전파 한 번 분량이므로, 앞 절의 비율대로 순전파 1과 역방향 2를 합친 3에 1이 더해져 4가 됩니다. 배입니다. 메모리를 에서 로 줄이는 값으로 학습 시간 33%를 내는 것이고, 큰 모델을 한정된 GPU 메모리에 올릴 때 가장 먼저 켜는 옵션이 이것인 이유입니다.
코드로 확인하기
순서와 비용 검산
mv = lambda M, v: [sum(M[i][j] * v[j] for j in range(len(v))) for i in range(len(M))]
vm = lambda u, M: [sum(u[i] * M[i][j] for i in range(len(u))) for j in range(len(M[0]))]
Jf = [[12, 4], [1, 3]] # f 의 야코비안 (2, 3)에서
Jg = [[1, 1], [11, 12]] # g 의 야코비안 f(2,3) = (12, 11)에서
J = [[sum(Jg[i][k] * Jf[k][j] for k in range(2)) for j in range(2)] for i in range(2)]
print(J) # [[13, 7], [144, 80]]
# ① 층을 나눠 곱해도 답이 같다 — 그러나 J 를 만들지 않는다
v, u = [1, 2], [1, 2]
print(mv(J, v), mv(Jg, mv(Jf, v))) # [27, 304] [27, 304]
print(vm(u, J), vm(vm(u, Jg), Jf)) # [301, 167] [301, 167]
# ② JVP 는 열을, VJP 는 행을 준다
print([mv(J, e) for e in ([1, 0], [0, 1])]) # [[13, 144], [7, 80]] ← 열 둘
print([vm(e, J) for e in ([1, 0], [0, 1])]) # [[13, 7], [144, 80]] ← 행 둘
# ③ 원소별 활성함수의 VJP 는 행렬을 만들지 않는다
n = 4
dphi = [0.5, 2.0, 1.5, 0.25] # φ'(x) 를 미리 구해 둔 값
uu = [1.0, 2.0, 3.0, 4.0]
D = [[dphi[i] if i == j else 0.0 for j in range(n)] for i in range(n)]
print(vm(uu, D)) # [0.5, 4.0, 4.5, 1.0] n² 개를 만들어 곱한 것
print([uu[i] * dphi[i] for i in range(n)]) # [0.5, 4.0, 4.5, 1.0] 원소별 곱 하나
# ④ 곱하는 순서만으로 비용이 4096배 갈린다
w, layers = 4096, 100
right = (layers - 1) * w ** 3 # 행렬 × 행렬 을 99번
left = (layers - 1) * 1 * w ** 2 # 행 × 행렬 을 99번
print(f"{right:.3g} {left:.3g} {right / left:.0f}") # 6.8e+12 1.66e+09 4096
# ⑤ 손실이 스칼라라는 사실이 답을 정한다
n, m, one_pass = 7e9, 1, 0.1 # 파라미터 70억, 출력 1개, 순전파 0.1초
print(f"역방향 {m * one_pass:.1f}초, 전방 {n * one_pass / 3.15e7:.0f}년")
# 역방향 0.1초, 전방 22년
③이 이 글의 요점을 가장 짧게 보여 줍니다. 위아래 두 줄의 답이 같은데, 위는 개짜리 행렬을 만들었고 아래는 곱셈을 번 했습니다. 실제 프레임워크가 등록하는 것은 아래쪽 한 줄입니다. ⑤의 0.1초는 되돌아오는 쪽의 배율을 뺀 굴리기 횟수만 센 것이라, 앞 절의 두 배를 곱하면 0.2초쯤이 됩니다.
갈래 검산
「층으로 나눈 곱」 절의 손계산을 그대로 옮깁니다. 갈래마다 VJP를 따로 계산하고 에서 더한 값이 과 같은지 봅니다. 브로드캐스팅된 편향의 그래디언트가 배치 축의 합인지도 함께 확인합니다.
# 갈래 진 그래프: L = a*b, a = 2x, b = x**2
x = 3.0
a, b = 2 * x, x ** 2
ga, gb = b, a # L = a*b 의 VJP (u = 1)
via_a = ga * 2 # a = 2x 의 VJP
via_b = gb * 2 * x # b = x**2 의 VJP
print(via_a, via_b, via_a + via_b, 6 * x ** 2)
# 18.0 36.0 54.0 54.0
# 브로드캐스팅: Y = X + b, b 를 배치 3줄에 복사
U = [[1, 2], [3, 4], [5, 6]] # 위에서 내려온 dL/dY
print([sum(row[j] for row in U) for j in range(2)])
# [9, 12]
두 갈래가 18과 36으로 따로 도착하고 합이 54라 해석적 미분과 맞습니다. PyTorch에서 .grad에 쌓이는 것이 바로 이 두 수의 합입니다.
이중수 검산
이중수는 곱셈 규칙 하나만 정의하면 됩니다. 덧셈은 앞자리끼리·꼬리끼리 더하고, 곱셈은 앞에서 유도한 입니다. 열 줄 남짓으로 JVP 엔진이 됩니다.
class Dual:
def __init__(s, a, b=0.0): s.a, s.b = a, b
def __add__(s, o): o = o if isinstance(o, Dual) else Dual(o); return Dual(s.a + o.a, s.b + o.b)
__radd__ = __add__
def __mul__(s, o): o = o if isinstance(o, Dual) else Dual(o); return Dual(s.a * o.a, s.a * o.b + s.b * o.a)
__rmul__ = __mul__
def __repr__(s): return f"{s.a:g}+{s.b:g}ε"
f = lambda x1, x2: (x1 * x1 * x2, x1 + 3 * x2)
g = lambda u1, u2: (u1 + u2, u1 * u2)
print(f(Dual(2, 1), Dual(3, 2))) # (12+20ε, 11+7ε) 꼬리 = J_f v
print(g(*f(Dual(2, 1), Dual(3, 2)))) # (23+27ε, 132+304ε) 꼬리 = J v
print([g(*f(Dual(2, e1), Dual(3, e2))) for e1, e2 in ((1, 0), (0, 1))])
# [(23+13ε, 132+144ε), (23+7ε, 132+80ε)] 꼬리 = J 의 두 열
# 헤시안-벡터 곱: L = x1^2 x2 의 그래디언트 (2 x1 x2, x1^2) 에 이중수를 넣는다
grad = lambda x1, x2: (2 * x1 * x2, x1 * x1)
print(grad(Dual(2, 1), Dual(3, 2))) # (12+14ε, 4+4ε) 꼬리 = Hv
꼬리가 손 미분과 전부 맞습니다. 는 , 과 로 굴린 꼬리는 의 두 열 와 , 마지막 줄의 는 헤시안을 만들지 않고 얻은 입니다. 열을 얻는 데 순전파가 두 번 든 것도 눈여겨볼 만합니다. 입력이 둘이라 둘이고, 입력이 70억이면 70억 번입니다.
정리
- 야코비안의 곱은 결합법칙 덕분에 순서를 바꿔도 답이 같지만, 비용은 같지 않다. 폭 4096짜리 100층에서 차이가 정확히 4096배다.
- 역전파가 순전파의 두 배쯤인 것은 선형층마다 입력 쪽 와 가중치 쪽 두 곱을 하기 때문이고, 층 구성에 따라 1.5배에서 3배 사이를 오간다.
- JVP는 , VJP는 다. JVP는 입력 쪽 방향을 앞으로 밀고, VJP는 출력 쪽 가중치를 뒤로 되돌린다.
- 이중수 ()를 넣고 순전파를 한 번 하면 앞자리는 값, 꼬리는 JVP가 된다. 꼬리를 여러 칸 달면 열 여러 개가 한꺼번에 나온다.
- 곱을 층마다 나누면 매 걸음이 행렬 하나 × 벡터 하나라 야코비안이 통째로 만들어지지 않는다. 프레임워크가 연산마다 등록하는 것이 VJP 함수다.
- 값이 갈라지면 되돌아오는 VJP는 더해진다.
.grad누적과zero_grad(), 편향 그래디언트의 배치 축 합이 모두 이 규칙이다. - JVP 한 번은 열 하나, VJP 한 번은 행 하나를 준다. 전체 야코비안까지 각각 번과 번이 들고, 규칙은 좁은 쪽을 훑는 것이다. 이면 메모리가 가른다.
- 안쪽 VJP에 바깥 JVP를 겹치면 헤시안을 만들지 않고 를 얻는다.
- 학습은 이라 이다. 역방향이 한 번, 전방이 70억 번이다. 손실이 여럿이어도 가중합으로 스칼라를 먼저 만들면 VJP 한 번이다.
- 역방향의 대가는 메모리다. 구간 체크포인팅은 메모리를 로 줄이고 계산을 1.33배로 늘린다.
loss.backward()가 순전파의 두 배 정도에 끝나는 이유가 이것입니다. 그 안에서 벌어지는 일은 층마다 벡터 하나를 받아 벡터 두 개(아래로 넘길 것과 이 층이 가져갈 것)를 만드는 것뿐이고, 야코비안이라는 표는 개념으로만 존재합니다. 「역전파는 연쇄법칙 한 줄」이라는 말에 한 줄을 더 붙이면 — 연쇄법칙을 오른쪽이 아니라 왼쪽부터 적용한 것입니다.
여기까지 오면 원리는 닫혔습니다. 남은 것은 손으로 유도할 때의 실무입니다. 논문의 그래디언트 식을 코드로 옮기려 하면 전치가 어디에 붙어야 하는지에서 매번 막히는데, 다음 글이 그 규약을 정하고 shape로 검산하는 방법을 세웁니다.
읽어주셔서 감사합니다. 😊

