크로네커 근사: 한 층의 피셔를 작은 두 행렬로
자연 기울기 절의 파이썬은 출력이 하나인 폭 4096 층에서 4096 × 4096 피셔를 만들어 풀었다. 실제 층은 출력도 수천 개다. 입력과 출력이 모두 4096 인 층 하나만 해도 가중치가 4096² ≈ 1,678만 개이고, 그 피셔는 가중치 수의 제곱인 약 2.8 × 10¹⁴ 칸이다. 이 행렬을 스텝마다 만들어 역을 구할 수는 없다.
“자연 기울기의 장점을 살리면서, 그 거대한 행렬을 직접 다루지 않을 수는 없는가?”
K-FAC
마턴스(James Martens)와 그로스(Roger Grosse)(2015)의 K-FAC(Kronecker-factored Approximate Curvature, 크로네커 곱으로 나눈 근사 곡률. 여기서 곡률은 손실을 두 번 미분한 값을 가리킨다)은 층마다 한 블록만 남기고, 그 블록을 두 작은 행렬의 크로네커 곱(한 행렬의 원소마다 다른 행렬 전체를 곱해 큰 블록 행렬을 만드는 곱)으로 근사한다.
근사의 근거는 지수족의 구조가 아니다. 활성값끼리의 곱과 역전파 신호끼리의 곱이 통계적으로 독립이라는 가정이다. 저자들 스스로 「어떤 현실적인 가정 아래서도 아마 정확해지지 않을 큰 근사」라고 적었다. 그래도 16 × 16 으로 줄인 손글씨 숫자(MNIST)를 분류하는 작은 망(층마다 뉴런 20개)에서 정확한 피셔와 근사를 나란히 그려 보니, 근사가 피셔의 굵은 짜임을 잡고 있었다. 크로네커 곱의 역은 두 작은 행렬의 역의 크로네커 곱이라, 역을 구하는 비용이 확 줄어든다.
Shampoo: 기울기에서 바로 두 조각을
K-FAC 의 두 조각을 만들려면 활성값과 역전파 신호를 층마다 따로 모아야 하고, 층의 종류가 바뀌면 무엇을 조각으로 삼을지부터 다시 정해야 한다. K-FAC 의 곱해진 두 조각은 한 층의 기울기 행렬 G 의 "왼쪽 공분산"과 "오른쪽 공분산"을 닮았다. 굽타(Vineet Gupta) 외(2018)의 Shampoo 는 여기서 한 걸음 더 단순해진다. 기울기 G 로 두 조각을 직접 누적하고, 각각의 −1/4 제곱을 양쪽에 곱한다. 행렬의 차원마다 조각 하나씩이라, 저장할 칸 수가 가로 m, 세로 n 인 층에서 m²n² 이 아니라 m² + n² 이다.
Shampoo 의 짐은 행렬의 −1/4 제곱이다. 원 논문의 구현은 이것을 특이값 분해로 구했다. 특이값 분해는 행렬을 「돌리기 · 축마다 늘이기 · 돌리기」 셋의 곱 U Σ Vᵀ 로 쪼개는 일이고, Σ 의 대각선에 놓인 늘이는 배율이 특이값이다. 저자들은 이 계산이 비싸서 20~100스텝에 한 번만 다시 했다.
누적을 끄고 이번 스텝의 G 하나만 쓰면 놀라운 일이 생긴다. 이 식이 G 의 특이값을 모두 1로 바꾼 행렬 UVᵀ 가 된다. 번스타인과 뉴하우스(2024)가 짚은 사실이고, 왜 그런지는 장 끝 문제 「누적 없는 Shampoo는 UVᵀ 다」에서 직접 계산한다. 이것이 뒤에 볼 Muon의 업데이트다.
정리하면 앞의 무엇이 모자라서 다음이 나왔는지가 한 줄로 이어진다. 자연 기울기는 피셔 전체가 너무 컸다. K-FAC 은 층마다 두 조각으로 줄였지만, 활성값과 역전파 신호를 따로 모아야 했다. Shampoo 는 기울기 하나로 두 조각을 만들었지만, 행렬의 −1/4 제곱이 비쌌다. Muon 은 누적을 끄고 UVᵀ 를 행렬 곱만으로 구한다. 곁가지도 있다. Shampoo 의 한 구현(메타의 Distributed Shampoo)은 2024년 MLCommons AlgoPerf 대회의 외부 튜닝 부문에서 1위를 했고(기준선보다 학습이 28% 빨랐다), 비아스(Nikhil Vyas) 외(2024)의 SOAP(ShampoO with Adam in the Preconditioner’s eigenbasis, Shampoo 전조건 행렬의 고유기저 안에서 Adam 을 돌린다)는 Shampoo 의 추가 하이퍼파라미터와 계산 부담을 줄이는 쪽으로 이 계보를 이었다.
수확
“피셔 전체는 너무 커서, 층마다 두 작은 조각의 크로네커 곱으로 근사한다. K-FAC 은 활성값과 역전파 신호로, Shampoo 는 기울기 하나로 두 조각을 만든다. K-FAC 의 근사는 두 신호의 곱이 독립이라는 가정에 기대고, 누적을 끈 Shampoo 는 기울기의 특이값을 모두 1로 바꾼다.”
문제 6. 평균 손님 × 평균 지출
작은 가게의 하루 매출은 손님 수 × 1인당 지출이다. 평일 하루는 손님 100명이 1만 원씩, 주말 하루는 손님 40명이 3만 원씩 쓴다. (가) 평일 하루와 주말 하루의 평균 매출을 정확히 구하고, 「평균 손님 수 × 평균 지출」로 어림한 값과 견주어라. (나) 주말에 손님이 140명으로 늘고 지출은 3만 원 그대로라면 어떻게 되는가?
함께 풀기

평균끼리 곱하면 대충 맞겠죠. 평균은 극단을 깎으니까 좀 작게 나올 테고요. (가)는 손님 평균 70명, 지출 평균 2만 원이라 어림은 140만 원. 정확한 값은 (100 + 120)/2 = 110만 원.

어, 어림이 더 크네요. 작게 나올 줄 알았는데.

(나)도 같은 쪽으로 틀려요?

(나)는 정확한 값이 (100 + 420)/2 = 260만 원이고 어림은 120 × 2 = 240만 원이야. 이번엔 어림이 작아. 곱의 평균은 평균의 곱에 공분산을 더한 거라서, 손님이 많은 날 지출이 적으면(가) 어림이 크게, 손님이 많은 날 지출도 많으면(나) 어림이 작게 나와. 둘이 따로 놀 때만 맞고.

조별 과제 점수를 「평균 출석 × 평균 기여」로 매기면, 출석은 많은데 기여가 적은 조원 때문에 실제보다 후하게 나오는 거랑 같네요.
문제 7. K-FAC 의 걸음은 크게 틀리는가 작게 틀리는가
가중치 하나짜리 층을 생각한다. 그 가중치의 피셔는 E[a²g²] 이고(a 는 들어오는 활성값, g 는 역전파 신호), K-FAC 은 이것을 E[a²]·E[g²] 로 어림한다. (가) 두 예제의 (a, g) 가 (1, 2), (2, 1) 일 때와 (1, 1), (2, 2) 일 때 정확한 피셔와 K-FAC 의 어림을 구하라. (나) 자연 기울기의 걸음은 기울기를 피셔로 나눈 것이다. 두 경우에 K-FAC 의 걸음은 정확한 자연 기울기 걸음의 몇 배인가? (다) 어느 경우가 학습을 위험하게 만드는가?
함께 풀기

(가)는 앞 문제랑 같은 계산이야. 첫째 경우 정확한 값은 (1·4 + 4·1)/2 = 4, 어림은 E[a²] = 2.5, E[g²] = 2.5 라서 6.25. 둘째 경우는 정확한 값 (1 + 16)/2 = 8.5, 어림은 그대로 6.25.

첫째 경우는 K-FAC 피셔가 실제보다 크니까 걸음도 실제보다 크겠네요. 큰 피셔면 큰 걸음이요.

자연 기울기 식에서 피셔는 기울기에 곱해져요, 나눠져요?

역행렬이니까 나눠지죠… 아, 거꾸로네요. 피셔를 크게 어림하면 걸음은 작아져요. 첫째 경우 4/6.25 = 0.64 배, 둘째 경우 8.5/6.25 = 1.36 배.

큰 활성값이 큰 역전파 신호와 함께 오는 둘째 경우가 위험해. 피셔를 작게 어림해서 걸음이 실제 눈금보다 36% 커지니까. 첫째 경우는 걸음이 줄어서 느려질 뿐이고. 매출 문제에서 손님 수와 지출이 같이 움직이면 어림이 작게 나왔던 것과 같은 방향이야.

그래서 K-FAC 의 독립 가정은 「대충 맞는다」로 끝나는 말이 아니에요. 어느 쪽으로 틀리는지가 두 신호의 관계에 달려 있어요.
문제 8. 폭 8192 층 하나의 피셔
입력과 출력이 모두 8192 인 층 하나를 생각한다. 숫자 하나를 4바이트(32비트 부동소수점)로 저장한다. (가) 이 층의 파라미터 수, 피셔 행렬의 칸 수, 그 피셔를 저장하는 데 드는 바이트 수를 구하라. (나) K-FAC 의 두 조각을 저장하는 데는 얼마가 드는가? (다) n × n 행렬의 역을 구하는 계산이 대략 n³ 에 비례한다면, 피셔 전체의 역과 두 조각의 역은 몇 배 차이인가?
함께 풀기

파라미터는 8192² = 6,711만 개요. 피셔도 가중치랑 같은 8192 × 8192 행렬이니까 6,711만 칸, 4바이트면 268 MB. GPU 에 충분히 들어가는데요.

피셔의 칸 하나는 무엇과 무엇의 짝을 적어요?

파라미터 하나와 파라미터 하나의 짝이요… 그럼 파라미터 수의 제곱이네요. 6,711만² ≈ 4.5 × 10¹⁵ 칸, 4바이트면 1.8 × 10¹⁶ 바이트, 18 페타바이트예요. 저장부터 안 돼요.

K-FAC 두 조각은 활성값 쪽 8192 × 8192 와 역전파 쪽 8192 × 8192 라서 1억 3,422만 칸, 537 MB 야. (다)는 피셔 전체가 (8192²)³ = 8192⁶ 이고 두 조각이 2 × 8192³ 이니까 비가 8192³/2 ≈ 2.7 × 10¹¹ 배.

피셔는 가중치의 모양이 아니라 가중치 「짝」의 모양이었네요. 조원 사이의 의견 차이를 다 적으려면 조원 수가 아니라 조원 쌍의 수만큼 칸이 필요한 거랑 같아요.