Muon의 직교화 다항식을 반복마다 새로 고르면
ICLR 2026 구두 발표 ‘The Polar Express’ 짧게 읽기
The Polar Express: Optimal Matrix Sign Methods and Their Application to the Muon AlgorithmICLR 2026 Oral
한 줄 요약
Muon 옵티마이저는 모멘텀 행렬의 특이값을 모두 1로 바꾼 방향으로 걸음을 옮기는데, 이 논문은 그 계산에 쓰는 다항식을 반복마다 새로 골라 같은 행렬 곱 횟수로 최악 오차를 가장 작게 만든다.
무엇이 달라졌나
Muon은 층마다 모멘텀 행렬 \(M = U\Sigma V^\top\)을 구한 뒤 \(W \leftarrow W - \lambda\,UV^\top\)로 가중치를 고친다(1.1절). \(UV^\top\)를 극 인자(polar factor)라고 부르는데, 특이값 분해(SVD)로 구하면 GPU에서 느리다. 그래서 행렬 곱만으로 계산되는 홀수 다항식 \(p(x)=ax+bx^3+cx^5\)를 몇 번 되풀이해 근사한다. 이 다항식을 행렬에 적용하면 특이값 하나하나에 \(p\)를 적용한 것과 같으므로, 반복할수록 모든 특이값이 1로 모이면 된다.
그동안 쓰던 선택지는 둘이었다. 고전적인 뉴턴-슐츠(Newton-Schulz) 반복은 끝내 수렴하지만 처음에는 거의 나아가지 않는다. 특이값이 \(10^{-6}\)에서 1까지 퍼진 합성 행렬에서 처음 17번 동안 진척이 거의 없었다(4.1절, 그림 3). 조던(Jordan)이 탐색으로 찾은 고정 다항식은 빨리 내려가지만 오차 약 0.3에서 더 줄지 않는다(1.2절).
Polar Express는 이 둘의 약점을 함께 피한다. 반복 다섯 번으로 맞춘 GPT-2 실험에서, 7억 7,400만 파라미터의 GPT-2-Large를 FineWeb 10억 토큰으로 학습하자 최적 학습률에서 검증 손실이 3.340이었다. 조던 방법은 3.398, You의 방법은 3.399였다(그림 1). GPT-2-Small에서도 3.588 대 3.639, 3.629였다(그림 4). 세 방법 모두 반복 한 번에 5차 다항식 하나를 쓰므로 계산량은 같다(4.2절). 이 방법은 NanoGPT 스피드런 코드에도 들어갔다(1.3절).
어떻게 보였나
먼저 행렬 문제를 스칼라 문제로 바꾼다. 특이값이 구간 \([\ell, u]\) 안에 있다고 하면, 스펙트럼 노름으로 잰 최악 오차는 그 구간에서 \(|1-p(x)|\)가 가장 커지는 값과 같다(2절). 그러니 여러 번 합성한 다항식 가운데 구간 전체를 1에 가장 가깝게 보내는 것을 찾으면 된다.
합성 전체를 한꺼번에 최적화하기는 어렵지만, 저자들은 한 단계씩 욕심껏 골라도 된다는 것을 증명했다(정리 3.1). 첫 다항식 \(p_1\)은 \([\ell,u]\)를 1에 가장 가깝게 보내는 5차 홀수 다항식이다. 그러면 특이값은 더 좁은 구간 \([\ell_2, u_2]\)로 옮겨 가고, \(p_2\)는 이 구간에 맞춰 다시 고른다. 이렇게 이어 붙인 합성이 전체로도 최적이고, 오차는 \(1-\ell_{T+1}\)이다. 각 단계의 다항식은 등진동 정리(equioscillation theorem)에 기대는 레메즈(Remez) 알고리즘을 단순하게 고쳐 구한다(3.2절).
구간이 넓은 초반에는 다항식이 작은 특이값을 세게 끌어올리고, 구간이 1 근처로 좁아지면 뉴턴-슐츠와 같은 꼴에 가까워진다. 그래서 초반에도 빠르고 끝에서는 3차 수렴을 지킨다(정리 3.3). 계수는 입력과 상관없이 한 번만 미리 계산해 두고, 가장 작은 특이값의 추정치 \(\ell\)은 bfloat16의 정밀도에 맞춰 \(10^{-3}\)으로 둔다(3.3절). bfloat16에서 값이 튀지 않도록 다항식을 \(p(x/1.01)\)로 조금 늦추고, 초반에는 덜 출렁이는 다항식을 쓰는 보정도 들어간다(3.4절).
눈여겨볼 점
첫째, 학습 토큰을 늘리자 차이가 크게 줄었다. 가중치 감쇠 0.1을 준 GPT-2-Large에서 10억 토큰일 때는 3.344 대 3.401, 3.390이었는데(그림 12), 100억 토큰으로 늘리자 2.913 대 2.921, 2.919로 좁혀졌다(그림 6). GPT-2 실험에는 여러 시드로 반복한 결과가 없어서, 0.01이 안 되는 차이가 실행마다 생기는 편차보다 큰지는 알 수 없다.
둘째, 논문이 최적으로 만든 기준과 학습에 중요한 기준이 서로 다르다. 정리 3.1은 가장 작은 특이값까지 포함한 최악 오차를 줄이는데, 절제 실험에서는 반복을 여섯 번보다 늘리거나 SVD로 정확한 극 인자를 써도 검증 손실이 나아지지 않았다. SVD를 쓰면 학습 한 걸음에 드는 시간만 두 배가 됐다(4.3절, 그림 5). 부록에서는 가장 큰 특이값의 1,000분의 1보다 작은 방향을 0으로 보내든 −1로 보내든 결과가 비슷했다(부록 H.1, 그림 9). 블로그의 해석으로는, Muon에 필요한 것은 큰 특이값들을 빨리 1 근처로 모으는 일이고 작은 특이값을 정확히 맞추는 일은 크게 중요하지 않아 보인다.
열린 질문
- 극 인자 근사의 어떤 오차가 학습 품질을 좌우하는가. 최악 오차 대신 큰 특이값 쪽에 무게를 둔 기준으로 다항식을 고르면 같은 행렬 곱 횟수로 더 나은 옵티마이저가 되는지, 아니면 학습에는 차이가 없는지는 아직 모른다.
이해 확인
Polar Express가 반복마다 다른 다항식을 쓰는 까닭은?
- 확률적 기울기의 잡음이 반복마다 달라서 다항식을 그때그때 다시 학습하기 때문이다
- 반복할 때마다 특이값이 놓인 구간이 좁아지므로, 그 구간을 1에 가장 가깝게 보내는 다항식이 매번 달라지기 때문이다
- 다항식 차수를 반복마다 올려 정확도를 높이기 때문이다
- 입력 행렬의 특이값을 매번 SVD로 구해 거기에 맞추기 때문이다
첫 다항식이 [ℓ, u]를 더 좁은 [ℓ₂, u₂]로 옮기면, 다음 다항식은 이 새 구간에 맞춰 고른다. 정리 3.1은 이렇게 한 단계씩 고른 합성이 전체로도 최적임을 보인다. 계수는 입력과 상관없이 미리 계산해 두고, 차수는 5로 고정하며, SVD는 쓰지 않는다.
논문 정보
- 제목: The Polar Express: Optimal Matrix Sign Methods and Their Application to the Muon Algorithm
- 저자: Noah Amsel, David Persson, Christopher Musco, Robert M. Gower (New York University, Flatiron Institute)
- 학회: ICLR 2026 구두 발표(Oral)
- 링크: arXiv 2505.16932, OpenReview, 코드