수학

MATH / 중급 48번

1차·2차 모멘트에서 Adam을 직접 조립하기

AdamW 한 줄에 든 숫자 다섯 개가 각각 무엇을 조절하는지를 갱신식에서 되짚습니다. 2차 모멘트로 나누는 것이 왜 스케일 불변이 되는지, 편향 보정 1/(1−βᵗ)이 없으면 걸음이 몇 배로 부푸는지, 그리고 L2 정규화와 AdamW의 감쇠가 수식의 어느 자리에서 갈리는지를 유도합니다.

PALDYN Team19 MIN READ

학습 스크립트에서 옵티마이저를 세우는 줄은 거의 언제나 이 모양입니다.

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=3e-4, betas=(0.9, 0.95), eps=1e-8, weight_decay=0.1,
)

숫자가 다섯 개인데 어느 것도 설명 없이 놓여 있습니다. betas의 두 번째가 왜 첫 번째보다 1에 가까운지, eps가 왜 하필 10−810^{-8} 인지, weight_decay가 손실에 더하는 정규화 항과 같은 것인지 다른 것인지. 이 글은 갱신식을 처음부터 조립해서 저 다섯 자리에 무엇이 들어가는지를 하나씩 채웁니다.

지난 글에서 모멘텀까지 왔고, 거기서 남은 문제가 하나 있었습니다. 학습률 η\eta 가 모든 파라미터에 하나라는 것입니다. 어떤 방향은 곡률이 크고 어떤 방향은 작은데 걸음의 크기는 같으니, 안전한 η\eta 를 고르면 완만한 방향이 영영 안 움직이고 완만한 방향에 맞추면 가파른 방향이 발산합니다.

방향마다 걸음을 다르게 주려면 무엇을 재야 하는가

방향별로 걸음을 나누려면 "이 방향은 얼마나 가파른가"를 재는 값이 필요합니다. 곡률을 직접 재려면 2차 미분이 필요하고 그건 비쌉니다(그 계산은 뉴턴법과 2차 방법이 맡습니다). 대신 쓸 수 있는 것은 이미 손에 있는 것, 그래디언트 자체의 최근 이력입니다.

가파른 방향에서는 그래디언트의 절댓값이 크고 스텝마다 부호가 잘 뒤집힙니다. 완만한 방향에서는 절댓값이 작고 한쪽으로만 조금씩 갑니다. 그러니까 부호를 지우고 크기만 남긴 평균을 방향마다 따로 들고 있으면 됩니다. 부호를 지우는 가장 쉬운 방법은 제곱입니다.

vt=β2 vt−1+(1−β2) gt2,v0=0v_t = \beta_2\,v_{t-1} + (1-\beta_2)\,g_t^2, \qquad v_0 = 0

여기서 gtg_t 는 이번 스텝의 그래디언트이고 gt2g_t^2 은 성분마다 제곱한 벡터입니다. 지난 글의 지수이동평균과 완전히 같은 점화식인데 재는 대상이 gg 가 아니라 g2g^2 입니다. 이것을 2차 모멘트라고 부릅니다 — 확률변수의 kk 차 모멘트가 E[Xk]E[X^k] 이고, vtv_t 는 E[g2]E[g^2] 의 추정값이기 때문입니다. 같은 이유로 지난 글의 모멘텀 mtm_t 는 1차 모멘트, 즉 E[g]E[g] 의 추정값입니다.

그래디언트를 vt\sqrt{v_t} 로 나누면 방향마다 규모가 사라집니다.

방향마다 다른 그래디언트 규모가 √v로 나눈 뒤 사라지는 것

왼쪽에서 방향 A는 2.5 안팎, 방향 B는 0.09 안팎으로 28배 차이가 납니다. 각각 자기 v\sqrt{v} 로 나눈 오른쪽에서는 둘 다 ±1\pm 1 근처입니다. 걸음의 크기가 그래디언트의 절대적인 규모와 무관해진 것입니다.

이 성질을 정확히 적어 두는 편이 낫습니다. 어떤 방향의 그래디언트가 통째로 cc 배가 된다고 해 봅시다. 그러면 mtm_t 도 cc 배, vtv_t 는 c2c^2 배, vt\sqrt{v_t} 는 ∣c∣|c| 배가 되므로

mtvt  ⟶  c mt∣c∣vt=sign⁡(c)⋅mtvt\frac{m_t}{\sqrt{v_t}} \;\longrightarrow\; \frac{c\,m_t}{|c|\sqrt{v_t}} = \operatorname{sign}(c)\cdot\frac{m_t}{\sqrt{v_t}}

로 크기가 그대로입니다. 갱신량이 그래디언트의 스케일에 영향받지 않는다는 뜻이고, 이것을 스케일 불변성이라고 부릅니다. 손실 함수에 100을 곱해도 Adam의 걸음은 변하지 않습니다. 순수한 경사하강법이라면 걸음이 100배가 됩니다.

vv 만 쓰고 모멘텀을 빼면 그것이 RMSProp입니다 — 그래디언트를 최근 제곱평균제곱근으로 나누기만 하는 방법이고, Adam은 여기에 1차 모멘트를 얹은 것입니다.

초반에 걸음이 부푸는 문제

점화식을 v0=0v_0 = 0 에서 시작했다는 점이 문제를 하나 만듭니다. 지난 글에서 EMA를 펼치면 가중치의 합이 11 이 아니라 1−βt1 - \beta^t 라는 것을 봤습니다. 같은 계산을 vtv_t 에 하면

vt=(1−β2)∑k=0t−1β2 k gt−k2v_t = (1-\beta_2)\sum_{k=0}^{t-1}\beta_2^{\,k}\,g_{t-k}^2

이고, 그래디언트의 제곱이 대체로 일정한 값 g2ˉ\bar{g^2} 근처에 있다고 두면

E[vt]≈g2ˉ (1−β2)∑k=0t−1β2 k=g2ˉ (1−β2 t)E[v_t] \approx \bar{g^2}\,(1-\beta_2)\sum_{k=0}^{t-1}\beta_2^{\,k} = \bar{g^2}\,(1-\beta_2^{\,t})

가 됩니다. 즉 vtv_t 는 참값보다 1−β2 t1-\beta_2^{\,t} 배만큼 작게 나옵니다. tt 가 작을수록 심합니다. 이 어긋남을 없애려고 1−β2 t1-\beta_2^{\,t} 로 나누는 것을 편향 보정이라고 합니다.

m^t=mt1−β1 t,v^t=vt1−β2 t\hat m_t = \frac{m_t}{1-\beta_1^{\,t}}, \qquad \hat v_t = \frac{v_t}{1-\beta_2^{\,t}}

얼마나 심한지 숫자로 보면 이렇습니다. β2=0.999\beta_2 = 0.999 일 때 t=1t=1 에서 v1=0.001 g12v_1 = 0.001\,g_1^2 이므로 v1=0.0316 ∣g1∣\sqrt{v_1} = 0.0316\,|g_1| 입니다. 참값의 3.2% 밖에 안 되는 값으로 나누니 걸음이 그만큼 부풉니다.

편향 보정을 하지 않았을 때 √v가 참값에 얼마나 못 미치는지

vt\sqrt{v_t} 가 참값의 절반까지 오는 데 β2=0.999\beta_2 = 0.999 면 288스텝, 90%까지 오는 데 1,660스텝이 걸립니다. 보정이 없으면 학습 초반 수천 스텝 내내 걸음이 잘못된 크기로 나간다는 뜻입니다.

여기서 mm 에도 같은 편향이 있으니 서로 상쇄되지 않느냐고 물을 수 있습니다. 상쇄되지 않습니다. 두 편향의 지수가 다르기 때문입니다. t=1t=1 에서 보정 없이 계산하면

m1v1=(1−β1)g(1−β2)g2=0.10.0316=3.162\frac{m_1}{\sqrt{v_1}} = \frac{(1-\beta_1)g}{\sqrt{(1-\beta_2)g^2}} = \frac{0.1}{0.0316} = 3.162

로 참값 11 의 3.16배입니다. 보정하면 m^1/v^1=g/∣g∣=1\hat m_1/\sqrt{\hat v_1} = g/|g| = 1 로 정확히 맞습니다. β1\beta_1 쪽 창이 훨씬 짧아 mm 이 먼저 차오르므로, 어긋남은 오히려 몇 스텝 뒤에 더 커집니다.

import numpy as np

def ratio(g=3.0, b1=0.9, b2=0.999, correct=True, upto=4000):
    m = v = 0.0
    out = {}
    for t in range(1, upto + 1):
        m = b1 * m + (1 - b1) * g
        v = b2 * v + (1 - b2) * g * g
        mh, vh = (m / (1 - b1**t), v / (1 - b2**t)) if correct else (m, v)
        if t in (1, 2, 5, 10, 50, 100, 300, 1000, 3000):
            out[t] = mh / (np.sqrt(vh) + 1e-8)
    return out

# 매 스텝 같은 그래디언트가 들어오면 m̂/√v̂ 는 언제나 1 이어야 한다
for name, c in (("보정 함  ", True), ("보정 안 함", False)):
    r = ratio(correct=c)
    print(name, " ".join(f"{t}:{r[t]:.2f}" for t in r))

# 보정 함   1:1.00 2:1.00 5:1.00 10:1.00 50:1.00 100:1.00 300:1.00 1000:1.00 3000:1.00
# 보정 안 함 1:3.16 2:4.25 5:5.80 10:6.53 50:4.50 100:3.24 300:1.96 1000:1.26 3000:1.03

보정을 빼면 열 번째 스텝에서 걸음이 참값의 6.5배까지 부풀고, β2\beta_2 의 창 1,000스텝을 세 번 채운 뒤에야 1로 돌아옵니다. 학습률을 아무리 정성껏 골라도 초반 수천 스텝은 다른 학습률로 도는 셈입니다. 워밍업으로 이 구간을 덮는 관행과도 겹치는데, 그 이야기는 다음 글이 맡습니다.

조립한 갱신식

지금까지의 조각을 순서대로 놓으면 Adam이 나옵니다.

Adam 갱신식의 조립 순서

mt=β1mt−1+(1−β1)gtvt=β2vt−1+(1−β2)gt2m^t=mt/(1−β1 t),v^t=vt/(1−β2 t)θt=θt−1−η m^tv^t+ε\begin{aligned} m_t &= \beta_1 m_{t-1} + (1-\beta_1)g_t \\ v_t &= \beta_2 v_{t-1} + (1-\beta_2)g_t^2 \\ \hat m_t &= m_t/(1-\beta_1^{\,t}), \qquad \hat v_t = v_t/(1-\beta_2^{\,t}) \\ \theta_t &= \theta_{t-1} - \eta\,\frac{\hat m_t}{\sqrt{\hat v_t}+\varepsilon} \end{aligned}

나눗셈과 제곱근은 성분별입니다. 벡터 하나를 다른 벡터로 성분마다 나누는 것은 대각행렬을 곱하는 것과 같으므로, 갱신을 이렇게 다시 쓸 수 있습니다.

θt=θt−1−η Pt m^t,Pt=diag⁡ ⁣(1v^t,1+ε,…,1v^t,n+ε)\theta_t = \theta_{t-1} - \eta\,P_t\,\hat m_t, \qquad P_t = \operatorname{diag}\!\left(\frac{1}{\sqrt{\hat v_{t,1}}+\varepsilon},\dots,\frac{1}{\sqrt{\hat v_{t,n}}+\varepsilon}\right)

그래디언트에 곱해서 방향을 다시 재는 이런 행렬을 프리컨디셔너라고 부릅니다. Adam의 프리컨디셔너는 대각행렬이라 좌표축 방향의 규모만 고칠 수 있고 축이 기울어진 골짜기는 펴지 못합니다. 그래도 대각선인 덕분에 nn 개 파라미터에 메모리가 2n2n, 연산이 O(n)O(n) 으로 끝납니다. 제대로 된 프리컨디셔너를 쓰려면 n2n^2 개짜리 행렬을 뒤집어야 하고, 그 비교는 뉴턴법과 2차 방법이 합니다.

세 손잡이가 각각 조절하는 것

값 조절하는 것 흔한 값 바꾸면
β1\beta_1 방향의 기억 길이. 창은 1/(1−β1)1/(1-\beta_1) 0.9 (창 10) 키우면 진행이 매끄럽고 방향 전환이 늦다
β2\beta_2 규모의 기억 길이 0.999 (창 1000), LLM은 0.95~0.98 작게 하면 최근 규모에 민감, 크게 하면 안정적이나 편향이 오래 간다
ε\varepsilon 걸음의 상한 10−810^{-8} 키우면 Adam이 SGD 쪽으로 넘어간다

β2\beta_2 가 β1\beta_1 보다 1에 가까운 이유는 재는 대상이 다르기 때문입니다. 방향은 지금 어디로 가는지가 중요하니 짧은 창이 맞고, 규모는 안정된 통계량이어야 하니 긴 창이 맞습니다. 언어모델 학습에서 β2\beta_2 를 0.95 근처로 내려 쓰는 것은 손실 스파이크 뒤에 규모 추정이 빨리 따라오게 하려는 것입니다 — 창이 1,000스텝이면 한 번 튄 그래디언트가 1,000스텝 동안 분모에 남습니다.

ε\varepsilon 은 흔히 "0으로 나누는 것을 막는 값"이라고 설명되지만 그것만은 아닙니다. 분모가 v^t+ε\sqrt{\hat v_t} + \varepsilon 이므로 갱신량의 크기에 상한이 걸립니다.

∣η m^tv^t+ε∣≤η ∣m^t∣ε\left|\eta\,\frac{\hat m_t}{\sqrt{\hat v_t}+\varepsilon}\right| \le \frac{\eta\,|\hat m_t|}{\varepsilon}

그래디언트가 아주 작아 v^t≪ε\sqrt{\hat v_t} \ll \varepsilon 인 방향에서는 나눗셈이 사실상 꺼지고 갱신이 ηm^t/ε\eta\hat m_t/\varepsilon, 즉 학습률 η/ε\eta/\varepsilon 짜리 모멘텀 SGD가 됩니다. ε\varepsilon 을 10−810^{-8} 에서 10−410^{-4} 로 올리면 그 문턱이 올라가 더 많은 방향이 SGD처럼 움직입니다. 값 하나로 두 방법 사이를 미끄러지는 손잡이인 셈입니다.

AdamW: 감쇠가 들어가는 자리

마지막 숫자 weight_decay=0.1이 남았습니다. 가중치를 0 쪽으로 당겨 두는 규제인데, 넣는 방법이 두 가지이고 Adam에서는 둘이 같지 않습니다.

첫 번째는 손실에 항을 더하는 L2 정규화입니다.

L′(θ)=L(θ)+λ2∥θ∥2  ⟹  gt=∇L(θ)+λθL'(\theta) = L(\theta) + \tfrac{\lambda}{2}\lVert\theta\rVert^2 \;\Longrightarrow\; g_t = \nabla L(\theta) + \lambda\theta

그래디언트에 λθ\lambda\theta 가 더해졌으니, 이 항은 mtm_t 에도 vtv_t 에도 들어갑니다. 그리고 갱신에서 v^t\sqrt{\hat v_t} 로 나뉩니다. 여기가 문제입니다 — 손실 쪽 그래디언트가 큰 파라미터는 v^t\sqrt{\hat v_t} 도 크므로 감쇠 항까지 함께 작아집니다. 즉 파라미터마다 실제 규제 세기가 다릅니다.

두 번째는 감쇠를 갱신식 밖에서 따로 빼는 것입니다. 이것이 분리형 가중치 감쇠이고, 그렇게 만든 옵티마이저가 AdamW입니다.

θt=θt−1−η m^tv^t+ε−ηλ θt−1\theta_t = \theta_{t-1} - \eta\,\frac{\hat m_t}{\sqrt{\hat v_t}+\varepsilon} - \eta\lambda\,\theta_{t-1}

L2 정규화와 AdamW에서 감쇠 항이 들어가는 자리

두 번째 항에는 v^t\sqrt{\hat v_t} 가 없습니다. 손실 쪽 그래디언트가 크든 작든 파라미터는 스텝마다 (1−ηλ)(1-\eta\lambda) 배로 줄어듭니다. 규제 세기가 λ\lambda 하나로 결정되고 그래디언트 규모와 무관해집니다.

차이는 실제로 큽니다. 손실 쪽 그래디언트의 규모만 바꿔 가며 같은 λ=0.1\lambda = 0.1 로 200스텝을 돌려 봤습니다.

import numpy as np

def run(mode, g_scale, lam=0.1, lr=0.01, b1=0.9, b2=0.999, steps=200):
    theta = 1.0
    m = v = 0.0
    for t in range(1, steps + 1):
        g = g_scale * np.sin(t)          # 손실 쪽 그래디언트, 평균은 0
        if mode == "l2":
            g = g + lam * theta          # 감쇠를 손실에 더한 경우
        m = b1 * m + (1 - b1) * g
        v = b2 * v + (1 - b2) * g * g
        mh, vh = m / (1 - b1**t), v / (1 - b2**t)
        theta -= lr * mh / (np.sqrt(vh) + 1e-8)
        if mode == "adamw":
            theta -= lr * lam * theta    # 감쇠를 갱신 밖에서
    return theta

for gs in (0.01, 1.0, 100.0):
    print(f"그래디언트 규모 {gs:>6}   L2 {run('l2', gs):.4f}   AdamW {run('adamw', gs):.4f}")

# 그래디언트 규모   0.01   L2 0.0169   AdamW 0.7866
# 그래디언트 규모    1.0   L2 0.7284   AdamW 0.7866
# 그래디언트 규모  100.0   L2 0.9584   AdamW 0.7866

L2 쪽은 같은 λ\lambda 인데 남은 값이 0.017부터 0.958까지 흩어집니다. 그래디언트가 작은 파라미터는 감쇠에 짓눌리고 큰 파라미터는 사실상 규제를 안 받습니다. AdamW 쪽은 세 경우 모두 0.7866입니다 — 이론값 (1−0.01×0.1)200=0.8187(1-0.01\times0.1)^{200} = 0.8187 에서 조금 내려간 것은 손실 쪽 갱신이 함께 움직이기 때문이고, 중요한 것은 세 값이 같다는 점입니다.

이것이 weight_decay를 옮기는 것만으로 이름이 Adam에서 AdamW로 바뀐 이유입니다. 두 줄의 코드 차이지만 λ\lambda 라는 손잡이가 뜻대로 동작하느냐 마느냐가 갈립니다. 순수 SGD에서는 v^\sqrt{\hat v} 라는 분모가 없으므로 두 방식이 학습률 배수만 다른 같은 것이 되고, 그래서 이 구분은 Adam 계열에서만 생깁니다. 정규화 기법 전반의 지형은 손실 함수와 정규화가 다루고, 이 글은 갱신식만 맡았습니다.

정리

  • 2차 모멘트 vt=β2vt−1+(1−β2)gt2v_t = \beta_2 v_{t-1} + (1-\beta_2)g_t^2 는 방향마다 그래디언트의 규모를 재는 EMA다. vt\sqrt{v_t} 로 나누면 갱신이 스케일 불변이 되어, 손실에 상수를 곱해도 걸음이 변하지 않는다.
  • v0=0v_0 = 0 에서 시작하므로 E[vt]≈g2ˉ(1−β2 t)E[v_t] \approx \bar{g^2}(1-\beta_2^{\,t}) 로 참값보다 작게 나온다. 편향 보정은 이 1−β2 t1-\beta_2^{\,t} 로 되돌리는 나눗셈이다. 보정을 빼면 β2=0.999\beta_2 = 0.999 에서 걸음이 최대 6.5배까지 부풀고 3,000스텝쯤 지나야 제자리로 온다.
  • 갱신 θ←θ−η m^t/(v^t+ε)\theta \leftarrow \theta - \eta\,\hat m_t/(\sqrt{\hat v_t}+\varepsilon) 는 대각행렬을 곱하는 것과 같다 — 대각선 프리컨디셔너다. 메모리 2n2n, 연산 O(n)O(n) 으로 끝나지만 기울어진 골짜기는 펴지 못한다.
  • β1\beta_1 은 방향의 창, β2\beta_2 는 규모의 창, ε\varepsilon 은 걸음의 상한 η∣m^t∣/ε\eta|\hat m_t|/\varepsilon 을 정한다. ε\varepsilon 을 키우면 Adam이 모멘텀 SGD 쪽으로 미끄러진다.
  • AdamW의 분리형 가중치 감쇠는 −ηλθ-\eta\lambda\theta 를 갱신 밖에서 뺀다. L2 정규화처럼 손실에 더하면 감쇠 항도 v^t\sqrt{\hat v_t} 로 나뉘어 파라미터마다 규제 세기가 달라진다 — 위 실험에서 0.017부터 0.958까지 흩어졌다.

이제 옵티마이저 안쪽은 다 열어 봤습니다. 남은 것은 밖에서 들어오는 값입니다 — gtg_t 자체가 데이터 전체의 그래디언트가 아니라 배치 하나로 잰 추정값이라는 사실입니다. 다음 글에서 그 추정의 분산이 배치 크기로 어떻게 변하는지 계산하고, 거기서 배치를 키울 때 학습률을 얼마나 올려야 하는지, 워밍업이 왜 필요한지, 코사인 감쇠가 후반에 무엇을 하는지를 끌어냅니다.


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

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