수학

MATH / 중급 51번

뉴턴법과 2차 방법: 왜 안 쓰는가, 요즘은 무엇을 근사하는가

뉴턴 갱신 −H⁻¹∇f를 2차 테일러 근사에서 유도하고, 이차수렴이 실제로 무엇을 뜻하는지 계산으로 봅니다. 그런 다음 7B 모델에서 헤세가 196엑사바이트라는 비용을 계산하고, 헤세를 만들지 않고 헤세-벡터 곱만으로 2차 정보를 쓰는 길, 그리고 Adam·K-FAC·Shampoo·Muon이 각각 그 행렬의 어느 부분을 남긴 것인지 정리합니다.

PALDYN Team18 MIN READ

큰 언어모델 학습에 Shampoo나 Muon 같은 옵티마이저를 썼다는 이야기가 늘었습니다. 소개하는 글에는 "2차 방법이 돌아왔다"는 말이 붙어 있는데, 정작 교과서는 오래전부터 "뉴턴법은 딥러닝에 쓸 수 없다"고 적어 두었습니다. 둘 다 맞는 말이려면 돌아온 것이 뉴턴법 자체는 아니어야 합니다.

이 글은 원본이 무엇인지부터 유도하고, 그것이 왜 그대로는 못 쓰이는지를 숫자로 계산한 다음, 요즘 방법들이 그 원본의 어느 부분을 남긴 것인지 한 표에 정리합니다.

지난 글까지 세 편은 전부 1차 정보만 썼습니다. 곡률은 학습률 상한 2/L2/L 의 LL 로만 등장했고 그마저 우리가 알 수 없는 값이었습니다. 이제 그 곡률을 직접 다뤄 봅니다.

뉴턴 갱신은 어디서 나오는가

경사하강법은 함수를 1차까지만 근사합니다 — 지금 자리의 기울기를 재고 그 반대로 갑니다. 얼마나 갈지는 아무도 안 알려 주므로 학습률이라는 손잡이가 필요합니다. 2차까지 근사하면 그 손잡이가 사라집니다.

테일러 전개와 곡률에서 본 대로 xx 근처에서

f(x+d)≈f(x)+∇f(x)Td+12 dTHdf(x+d) \approx f(x) + \nabla f(x)^{\mathsf T} d + \tfrac12\,d^{\mathsf T} H d

이고 HH 는 2차 편미분을 모은 헤세 행렬입니다. 우변은 dd 에 대한 이차식이므로 최소점을 바로 구할 수 있습니다. dd 로 미분해서 0으로 두면

∇f(x)+Hd=0  ⟹  d=−H−1∇f(x)\nabla f(x) + Hd = 0 \;\Longrightarrow\; d = -H^{-1}\nabla f(x)

이것이 뉴턴 갱신입니다. 근사한 포물면의 바닥까지 곧장 가라는 뜻이고, 걸음의 길이를 곡률이 정하므로 학습률이 없습니다.

이차함수에서 뉴턴법과 경사하강법의 경로

함수가 애초에 이차함수라면 근사가 곧 원함수이므로 한 스텝에 최소점에 닿습니다. 그림의 경사하강법은 두 방향의 곡률이 10배 다른 골짜기에서 최선의 학습률로도 아홉 스텝 뒤에 아직 도착하지 못했습니다. H−1H^{-1} 이 하는 일이 정확히 그 곡률 차이를 지우는 것입니다 — 지난 세 편에서 학습률·모멘텀·v^\sqrt{\hat v} 로 조금씩 흉내 내려 했던 그 일입니다.

이차수렴

일반 함수에서도 최소점 가까이 가면 매우 빠릅니다. 1차원에서 오차를 ek=xk−x∗e_k = x_k - x^* 라 두고 전개하면

ek+1≈f′′′(x∗)2f′′(x∗) ek2e_{k+1} \approx \frac{f'''(x^*)}{2f''(x^*)}\,e_k^2

로 오차가 제곱됩니다. 맞는 자릿수가 매 스텝 두 배가 된다는 뜻이고, 이것을 이차수렴이라고 합니다. 2\sqrt2 를 f(x)=x2−2f(x)=x^2-2 의 뉴턴법으로 구하면 이렇습니다.

스텝 xkx_k 오차
0 1.5000000000000000 8.6×10−28.6\times10^{-2}
1 1.4166666666666667 2.5×10−32.5\times10^{-3}
2 1.4142156862745099 2.1×10−62.1\times10^{-6}
3 1.4142135623746899 1.6×10−121.6\times10^{-12}
4 1.4142135623730951 0 (배정밀도 한계)

네 스텝에 16자리가 맞습니다. 경사하강법은 같은 정확도에 수천 스텝이 듭니다. 그렇다면 왜 안 쓰는 걸까요.

비용

nn 개 파라미터에 대해 HH 는 n×nn \times n 행렬입니다.

  • 저장: n2n^2 개 항목. 7B 모델이면 4.9×10194.9\times10^{19} 개, float32로 196엑사바이트입니다. 세상의 저장장치를 다 모아도 담지 못합니다.
  • 역행렬: O(n3)O(n^3) 연산. (7×109)3=3.4×1029(7\times10^9)^3 = 3.4\times10^{29} 이고, 초당 101510^{15} 연산을 하는 가속기 한 대로 천만 년입니다.

파라미터 수에 따른 프리컨디셔너 메모리

같은 그래프에서 대각선만 드는 Adam은 8n8n 바이트, 7B에 56GB입니다. 두 곡선의 기울기가 다르다는 것 — 하나는 nn 에, 하나는 n2n^2 에 비례한다는 것 — 이 이 이야기의 전부입니다. 모델이 커질수록 격차가 벌어지므로 "언젠가 하드웨어가 좋아지면"이라는 답은 없습니다.

비용 말고 문제가 하나 더 있습니다. HH 가 양정치가 아닐 수 있습니다. 어떤 방향의 곡률이 음수이면 그 방향의 H−1H^{-1} 성분도 음수라 뉴턴 방향이 손실을 올리는 쪽을 가리킵니다. 더 나쁜 것은 안장점입니다 — 뉴턴법은 "근사한 이차식의 정류점"으로 가는 방법이라 최소점과 안장점을 구별하지 않고, 안장점 근처에서는 오히려 그리로 빨려 들어갑니다. 딥러닝의 손실면은 안장점이 흔한 곳이라 이 성질이 특히 나쁩니다.

그래서 실제로 쓰이는 것은 HH 가 아니라 언제나 양반정치인 대체물입니다. 가장 흔한 것이 피셔 정보 행렬입니다.

F=E ⁣[∇θlog⁡pθ(y) ∇θlog⁡pθ(y)T]F = E\!\left[\nabla_\theta \log p_\theta(y)\,\nabla_\theta \log p_\theta(y)^{\mathsf T}\right]

바깥곱의 기댓값이라 정의상 양반정치이고, 로그가능도의 헤세와 부호만 바꾼 것의 기댓값과 같아서 최소점 근처에서는 HH 와 거의 같습니다. 아래에 나올 방법들이 근사하는 것은 대개 HH 가 아니라 이 FF 입니다.

헤세를 만들지 않고 쓰기

HH 를 통째로 드는 것이 문제였지, 2차 정보 자체가 문제였던 것은 아닙니다. 실제로 필요한 것은 행렬이 아니라 그 행렬을 벡터에 곱한 결과 HvHv 뿐인 경우가 많습니다. 그리고 그것은 HH 없이 얻을 수 있습니다.

Hv=∇x ⁣(∇f(x)Tv)Hv = \nabla_x\!\left(\nabla f(x)^{\mathsf T} v\right)

괄호 안은 스칼라 하나입니다. 그래디언트를 구하고, 고정된 vv 와 내적해 스칼라를 만들고, 그것을 다시 미분하면 됩니다. 자동미분으로 역전파를 두 번 하는 비용이고 메모리는 O(n)O(n) 입니다. 유한차분으로도 됩니다.

Hv≈∇f(x+εv)−∇f(x)εHv \approx \frac{\nabla f(x+\varepsilon v) - \nabla f(x)}{\varepsilon}

헤세를 만드는 길과 만들지 않는 길

HvHv 만 있으면 Hd=−∇fHd = -\nabla f 를 켤레기울기법으로 풀 수 있습니다 — 행렬을 벡터에 곱하는 연산만으로 연립방정식을 푸는 반복법입니다. 이렇게 뉴턴 방향을 구하는 방식을 Hessian-free 방법이라고 부릅니다.

import numpy as np
rng = np.random.default_rng(0)
n, m, lam = 200, 400, 1e-3
A = rng.normal(0, 1, (m, n)) / np.sqrt(n)

def grad(x):
    s = 1 / (1 + np.exp(-A @ x))
    return A.T @ s / m + lam * x

def hess(x):                                   # 비교용으로만 만든다
    s = 1 / (1 + np.exp(-A @ x))
    return (A.T * (s * (1 - s))) @ A / m + lam * np.eye(n)

x = rng.normal(0, 1, n)
v = rng.normal(0, 1, n); v /= np.linalg.norm(v)
exact = hess(x) @ v
for eps in (1e-1, 1e-3, 1e-5, 1e-7, 1e-9):
    ap = (grad(x + eps * v) - grad(x)) / eps
    print(f"ε={eps:<7} 상대오차 {np.linalg.norm(ap - exact) / np.linalg.norm(exact):.2e}")

hvp = lambda u: (grad(x + 1e-5 * u) - grad(x)) / 1e-5      # 헤세를 만들지 않는다
g = grad(x)
d = np.zeros(n); r = -g.copy(); p = r.copy()
ex = np.linalg.solve(hess(x), -g)
for k in range(1, 31):
    Hp = hvp(p)
    a = (r @ r) / (p @ Hp)
    d = d + a * p
    rn = r - a * Hp
    p = rn + (rn @ rn) / (r @ r) * p
    r = rn
    if k in (1, 3, 10, 30):
        print(f"CG {k:>2}회  뉴턴 방향과의 상대오차 {np.linalg.norm(d - ex) / np.linalg.norm(ex):.2e}")

# ε=0.1     상대오차 7.58e-04
# ε=0.001   상대오차 7.57e-06
# ε=1e-05   상대오차 7.56e-08
# ε=1e-07   상대오차 6.91e-08
# ε=1e-09   상대오차 8.01e-06
# CG  1회  뉴턴 방향과의 상대오차 3.43e-01
# CG  3회  뉴턴 방향과의 상대오차 3.76e-02
# CG 10회  뉴턴 방향과의 상대오차 1.09e-05
# CG 30회  뉴턴 방향과의 상대오차 2.62e-07

두 가지를 볼 수 있습니다. 첫째, 유한차분의 ε\varepsilon 에는 최적값이 있습니다 — 너무 크면 근사가 거칠고 너무 작으면 두 그래디언트의 차이가 반올림 오차에 묻힙니다. 여기서는 10−510^{-5} 근처가 가장 좋고 양쪽으로 갈수록 나빠집니다.

둘째, 파라미터가 200개인데 켤레기울기법 10회면 뉴턴 방향을 소수점 다섯 자리까지 맞춥니다. nn 번이 아니라 몇십 번이면 충분하다는 것이 이 방법이 성립하는 이유입니다.

같은 HvHv 로 행렬 거듭제곱과 파워 반복의 방법을 쓰면 HH 의 최대 고윳값도 잽니다. 그 값이 곧 학습률 상한 2/λmax⁡2/\lambda_{\max} 를 정하므로, 헤세를 만들지 않고도 "지금 학습률이 안전한가"를 확인할 수 있습니다.

요즘 방법들이 근사하는 것

Hessian-free 방법도 스텝마다 켤레기울기 수십 회가 들어 대규모 학습에는 비쌉니다. 그래서 실제로 쓰이는 길은 FF 를 통째로 다루지 말고 구조를 가정해 싸게 만드는 것입니다. 어느 칸을 버리느냐가 방법을 가릅니다.

여러 방법이 곡률 행렬의 어느 부분을 남기는가

Adam은 대각선만 남깁니다. v^t\hat v_t 가 성분별 E[g2]E[g^2] 이고 그것은 FF 의 대각 성분이니, Adam의 프리컨디셔너는 diag⁡(F)−1/2\operatorname{diag}(F)^{-1/2} 입니다. 지수가 −1-1 이 아니라 −1/2-1/2 인 것이 눈에 띄는데, 뉴턴법이라면 −1-1 이어야 합니다. 제곱근만 쓰는 쪽이 추정이 틀렸을 때 덜 위험하고 실제로 잘 도는 것이 관행이 된 이유입니다.

K-FAC은 층 하나의 FF 를 크로네커 곱으로 근사합니다. 층의 입력 aa 와 출력 쪽 그래디언트 ss 로 ∇WL=s aT\nabla_W L = s\,a^{\mathsf T} 이므로

F층=E ⁣[(saT)⊗(saT)]≈E[aaT]⊗E[ssT]=A⊗GF_{\text{층}} = E\!\left[(s a^{\mathsf T})\otimes(s a^{\mathsf T})\right] \approx E[aa^{\mathsf T}]\otimes E[ss^{\mathsf T}] = A \otimes G

로 두는 것입니다(두 기댓값이 독립이라는 가정이 들어갑니다). 크로네커 곱의 역은 각각의 역의 크로네커 곱, 즉 (A⊗G)−1=A−1⊗G−1(A\otimes G)^{-1} = A^{-1}\otimes G^{-1} 이므로 m×nm \times n 짜리 층에서 (mn)3(mn)^3 대신 m3+n3m^3 + n^3 만 듭니다.

Shampoo는 가중치를 행렬 WW 그대로 두고 좌우에서 각각 프리컨디셔너를 곱합니다. 그래디언트 GG 에 대해 L=∑GGTL = \sum GG^{\mathsf T}, R=∑GTGR = \sum G^{\mathsf T}G 를 모아 두고 L−1/4GR−1/4L^{-1/4} G R^{-1/4} 를 갱신 방향으로 씁니다. 행 방향과 열 방향의 스케일을 따로 고치는 것이라 대각선보다는 촘촘하고 전체보다는 훨씬 쌉니다.

Muon은 여기서 한 걸음 더 갑니다. 그래디언트 행렬을 특이값분해해서 특잇값을 전부 1로 바꾼 것, 곧 UVTUV^{\mathsf T} 를 갱신 방향으로 씁니다. 이것은 (GGT)−1/2G(GG^{\mathsf T})^{-1/2}G 와 같고, 실제로는 SVD를 하지 않고 뉴턴-슐츠 반복이라는 행렬 곱셈 몇 번으로 얻습니다. 방향만 남기고 크기는 전부 지운다는 점에서 Adam이 성분마다 하던 정규화를 행렬 전체에 대해 하는 것입니다.

방법 근사하는 것 층당 추가 메모리 층당 추가 연산
SGD 없음 0 0
Adam diag⁡(F)−1/2\operatorname{diag}(F)^{-1/2} 2mn2mn O(mn)O(mn)
Shampoo 행·열 방향 L−1/4, R−1/4L^{-1/4},\,R^{-1/4} m2+n2m^2+n^2 O(m3+n3)O(m^3+n^3), 드물게
K-FAC F≈A⊗GF \approx A\otimes G m2+n2m^2+n^2 O(m3+n3)O(m^3+n^3), 드물게
Muon (GGT)−1/2G(GG^{\mathsf T})^{-1/2}G mnmn (모멘텀) 행렬 곱 5~6회
뉴턴 H−1H^{-1} (mn)2(mn)^2 O((mn)3)O((mn)^3)

Shampoo와 K-FAC이 m3m^3 을 감당할 수 있는 것은 매 스텝 하지 않기 때문입니다. 곡률 통계는 천천히 변하므로 프리컨디셔너를 수백 스텝에 한 번만 다시 계산하고 그 사이에는 재사용합니다.

"2차 방법이 돌아왔다"의 정확한 뜻이 이제 보입니다. 돌아온 것은 H−1H^{-1} 이 아니라 한 층 안에서 행·열 구조를 이용한 값싼 프리컨디셔너이고, 그 값어치는 이차수렴이 아니라 같은 스텝 수로 더 나아가는 것입니다.

정리

  • 뉴턴 갱신 d=−H−1∇fd = -H^{-1}\nabla f 는 2차 테일러 근사의 최소점으로 곧장 가는 방향이다. 걸음의 길이를 곡률이 정하므로 학습률이 없고, 이차함수에서는 한 스텝에 도착한다.
  • 최소점 근처에서 오차가 제곱되는 이차수렴이라 2\sqrt2 를 네 스텝에 16자리까지 맞춘다.
  • 그런데 HH 는 n2n^2 개 항목이다. 7B 모델이면 196엑사바이트, 역행렬은 3.4×10293.4\times10^{29} 연산이다. 게다가 HH 는 양정치가 아닐 수 있어 안장점으로 끌리므로, 실제로는 언제나 양반정치인 피셔 정보 FF 를 대신 쓴다.
  • 필요한 것이 HH 가 아니라 HvHv 뿐이면 Hv=∇(∇fTv)Hv = \nabla(\nabla f^{\mathsf T}v) 로 메모리 O(n)O(n) 에 얻는다. 켤레기울기법과 묶으면 파라미터 200개짜리 문제에서 10회면 뉴턴 방향에 소수점 다섯 자리까지 닿는다.
  • 요즘 방법들은 FF 의 어느 칸을 버릴지로 갈린다 — Adam은 대각선만(diag⁡(F)−1/2\operatorname{diag}(F)^{-1/2}), K-FAC은 크로네커 두 조각(A⊗GA\otimes G), Shampoo는 행·열 방향 두 개, Muon은 그래디언트 행렬의 특잇값을 전부 1로. 전부 한 층 안에서만 곡률을 본다.

여기까지가 최적화입니다. 지금까지 손실 L(θ)L(\theta) 는 언제나 "데이터에 대한 평균"이었고, 그 데이터는 우리와 무관하게 주어져 있었습니다. 그런데 「좋은 답을 내게 한다」 같은 목표는 그렇게 적히지 않습니다 — 평가할 답을 모델이 스스로 만들어 내고, 따라서 기댓값을 취하는 분포 자체가 θ\theta 에 의존합니다. 다음 글에서 그 목적식을 세우고, 그 한 가지 차이가 왜 그래디언트를 기댓값 안으로 그냥 넣을 수 없게 만드는지 봅니다.


읽어주셔서 감사합니다. 😊

LATEST

수학의 최신 글

수학2026.09.07

양자화 오차: 격자 사상, 오차 분산, 이상치 채널

실수를 2^b개 격자에 사상할 때 오차의 분산이 왜 Δ²/12인지 유도하고, 그것이 비트당 6.02dB라는 SNR로 번역되는 과정을 실측과 대조했습니다. 이상치 하나가 나머지 값의 유효 비트를 어떻게 먹는지, 그리고 int4에서 성능이 무너지는 지점을 오차 예산으로 미리 계산하는 법까지.

중급18 MIN
수학2026.09.07

수치적으로 안정한 계산 패턴 모음

최댓값 빼기, 로그 공간, log1p·expm1, 분산의 두 공식, 정규화의 ε, fp32 누산, 역행렬 대신 solve — 프레임워크가 몰래 해 주는 일곱 가지를 하나씩 꺼내 각각 어떤 고장을 막는지 직접 재 봤습니다. 수식을 그대로 옮긴 코드가 왜 라이브러리보다 나쁜지에 대한 목록입니다.

중급22 MIN
수학2026.09.07

부동소수점은 어디서 새는가: 반올림, 상쇄, 더하는 순서

0.1 + 0.2가 0.3이 아닌 이유부터 시작해 머신 엡실론을 유도하고, 같은 16비트인데 fp16과 bf16이 서로 다른 지점에서 터지는 이유, 비슷한 수를 뺄 때 유효자리가 사라지는 파괴적 상쇄, 그리고 1,000만 개를 순서만 바꿔 더했을 때 오차가 백만 배 갈리는 실험까지 직접 재 봤습니다.

중급23 MIN