Muon의 직교화 다항식을 반복마다 새로 고르면

ICLR 2026 구두 발표 ‘The Polar Express’ 짧게 읽기

The Polar Express: Optimal Matrix Sign Methods and Their Application to the Muon AlgorithmICLR 2026 Oral

짧은 읽기
이론, 최적화
딥러닝 옵티마이저
Muon 옵티마이저가 쓰는 극 인자 근사에서, 반복마다 최악 오차가 가장 작은 다항식을 골라 같은 행렬 곱 횟수로 더 빨리 수렴하게 만든 논문을 짧게 읽는다.
공개

2026년 9월 29일

한 줄 요약

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절).

00.250.50.751024681012반복 횟수최악 오차Muon에서 흔히 쓰는 횟수뉴턴-슐츠조던: 0.32에서 멈춤Polar Express0.15
특이값 σ가 0.001에서 1 사이에 있을 때, t번 반복한 뒤의 최악 오차 max|1 − p(σ)|를 계산한 그래프. 정확한 극 분해라면 오차는 0이다. 뉴턴-슐츠는 (15x − 10x3 + 3x5)/8을, 조던 방법은 3.4445x − 4.7750x3 + 2.0315x5를 되풀이했고, Polar Express는 논문 부록 A의 코드에 적힌 다항식을 차례로 썼다. 셋 다 5차 다항식이라 한 번 반복에 드는 계산은 같다. 다섯 번 뒤 최악 오차는 뉴턴-슐츠 0.98, 조던 0.53, Polar Express 0.15이고, 여섯 번 뒤 Polar Express는 0.006까지 내려간다. 조던 방법은 여섯 번째부터 0.32에 머문다.

어떻게 보였나

먼저 행렬 문제를 스칼라 문제로 바꾼다. 특이값이 구간 \([\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 근처로 모으는 일이고 작은 특이값을 정확히 맞추는 일은 크게 중요하지 않아 보인다.

열린 질문

  1. 극 인자 근사의 어떤 오차가 학습 품질을 좌우하는가. 최악 오차 대신 큰 특이값 쪽에 무게를 둔 기준으로 다항식을 고르면 같은 행렬 곱 횟수로 더 나은 옵티마이저가 되는지, 아니면 학습에는 차이가 없는지는 아직 모른다.

이해 확인

Polar Express가 반복마다 다른 다항식을 쓰는 까닭은?

  1. 확률적 기울기의 잡음이 반복마다 달라서 다항식을 그때그때 다시 학습하기 때문이다
  2. 반복할 때마다 특이값이 놓인 구간이 좁아지므로, 그 구간을 1에 가장 가깝게 보내는 다항식이 매번 달라지기 때문이다
  3. 다항식 차수를 반복마다 올려 정확도를 높이기 때문이다
  4. 입력 행렬의 특이값을 매번 SVD로 구해 거기에 맞추기 때문이다

첫 다항식이 [ℓ, u]를 더 좁은 [ℓ₂, u₂]로 옮기면, 다음 다항식은 이 새 구간에 맞춰 고른다. 정리 3.1은 이렇게 한 단계씩 고른 합성이 전체로도 최적임을 보인다. 계수는 입력과 상관없이 미리 계산해 두고, 차수는 5로 고정하며, SVD는 쓰지 않는다.

논문 정보