μP: 폭이 달라져도 특징이 같은 속도로 배우게
폭을 바꿀 때마다 학습률을 다시 찾지 않으려면, 폭이 달라져도 똑같이 지켜지는 무언가를 정해 두고 거기에 맞춰 초기화와 학습률을 붙이면 된다. 문제는 무엇을 같게 지키느냐다.
“폭이 달라져도 같게 지켜야 할 「한 걸음의 크기」는 무엇으로 재는가?”
원래의 유도
양과 에드워드 휴(Edward Hu)(2021), 양 외(2022)가 고른 답은 각 층의 특징(은닉 활성값)이 한 스텝에 변하는 크기다. 폭을 무한대로 보내도, 학습 한 스텝이 각 층의 특징을 폭과 무관한 크기만큼 바꿔야 한다. 이 「폭과 무관한 크기」를 Θ(1) 로 적는다(Θ 는 그리스 문자 세타의 대문자다. 폭이 커져도 0 으로 줄지도, 끝없이 커지지도 않는 크기라는 뜻이다). 너무 작으면 특징이 사실상 얼어붙어 신경망이 초기값 근처의 선형 모델처럼 움직이고(커널 영역), 너무 크면 폭발한다. 이 요구를 층마다 초기화와 학습률의 규칙으로 옮긴 것이 μP(maximal update parametrization, 「뮤피」로 읽는다. 특징이 폭발하지 않는 한에서 가장 크게 움직이게 하는 설정)이고, μP 로 좁은 모델에서 찾은 학습률을 넓은 모델에 그대로 옮겨 쓰는 일을 μTransfer 라 부른다.
이 장 첫 절의 예(폭 n 인 층 하나에서 한 스텝의 출력 변화가 −lr·(y − t)·n)가 그 계산의 한 조각이다. 정렬된 합 때문에 변화가 입력 차원(fan_in)배로 커지니, 학습률을 1/fan_in 로 줄인다.
Tensor Programs V의 표 3이 Adam과 SGD에 대한 처방이다(괄호 안은 표준 설정).
| 입력층 가중치와 모든 편향 | 은닉층 가중치 | 출력층 가중치 | |
|---|---|---|---|
| 초기화 분산 | 1/fan_in | 1/fan_in | 1/fan_in² (표준 1/fan_in) |
| Adam 학습률 | 1 | 1/fan_in (표준 1) | 1/fan_in (표준 1) |
| SGD 학습률 | fan_out (표준 1) | 1 | 1/fan_in (표준 1) |
같은 논문의 표 8은 출력층에 1/fan_in 곱셈 계수를 달고 초기화 분산을 1, Adam 학습률을 1로 두는 동등한 형태를 준다. 두 표의 항목을 섞어 쓰면 안 된다. 어느 형태로 쓰는지 정하고 그 표를 통째로 따른다. SGD와 Adam의 처방이 다른 까닭도 한 스텝의 모양에서 보인다. SGD의 한 스텝은 기울기 크기에 비례하고, Adam의 한 스텝은 지수이동평균이 쌓이면 기울기를 제 크기로 나누기 때문에 좌표마다 lr 정도로 맞춰져 있다. 같은 "특징 변화 Θ(1)"을 맞추는 데 필요한 lr 의 폭 의존성이 달라진다.
양, 제임스 사이먼(James Simon), 번스타인(2023)은 이 조건을 스펙트럼 노름 하나로 묶었다.
스펙트럼 노름은 "벡터를 최대 몇 배로 늘리는가"다. 입력 벡터가 크기 √fan_in 이고 출력 벡터가 크기 √fan_out 이어야 한다면(좌표마다 Θ(1)), 가중치와 그 업데이트가 늘리는 배율이 √(fan_out/fan_in) 이면 된다.
피셔 계량과의 관계 — 해석
μP를 "피셔 거리를 일정하게 만드는 처방"이라고 말하고 싶어진다. 이것은 원래 유도가 아니라 해석이다. μP의 유도는 피셔 계량에서 출발하지 않았다. 가중치 위의 피셔는 출력 공간의 자를 끌어온 것이고, "출력(과 각 층의 특징)이 한 스텝에 폭과 무관한 크기만큼 변한다"는 μP의 요구는 "그 끌어온 자로 잰 한 걸음이 폭과 무관하다"에 가깝다. 정확히 같은 말은 아니다. μP는 출력만이 아니라 모든 은닉층의 특징 변화를 맞춘다.
μTransfer를 눈으로
2층 은닉 MLP(ReLU)를 Adam으로 300스텝 학습한다. 폭 32, 128, 512에서 학습률을 2−14부터 2−4까지 훑었다. μP 쪽은 표 3대로 은닉·출력층 학습률과 출력 초기화를 폭 32 기준으로 줄였다.
불러오는 중…
파이썬
위젯의 곡선을 만든 코드다(학습률 구간만 줄였다). CPU로 몇 분 걸린다.
import numpy as np
d, base, B, steps = 16, 32, 128, 300
g = np.random.default_rng(0)
X = g.standard_normal((1024, d))
Y = np.tanh(X @ g.standard_normal((d, 32)) / 4) @ g.standard_normal(32) / 2
def train(n, lr, mup):
r = np.random.default_rng(1)
W = [r.standard_normal((d, n)) / np.sqrt(d), # 입력층
r.standard_normal((n, n)) / np.sqrt(n), # 은닉층
r.standard_normal(n) / np.sqrt(n)] # 출력층
s = base / n if mup else 1.0
if mup: W[2] *= s # μP: 출력 초기화를 s 배로 작게 (0 에 가까워도 된다)
lrs = [lr, lr * s, lr * s] # μP(Adam): 은닉·출력 lr ∝ 1/fan_in
m = [np.zeros_like(w) for w in W]; v = [np.zeros_like(w) for w in W]
for t in range(1, steps + 1):
i = r.integers(0, len(X), B); x, y = X[i], Y[i]
h1 = np.maximum(x @ W[0], 0); h2 = np.maximum(h1 @ W[1], 0); e = h2 @ W[2] - y
g2 = h2.T @ e * 2 / B
d2 = np.outer(e, W[2]) * (h2 > 0) * 2 / B
g1 = h1.T @ d2
d1 = d2 @ W[1].T * (h1 > 0)
g0 = x.T @ d1
for k, gk in enumerate([g0, g1, g2]):
m[k] = 0.9 * m[k] + 0.1 * gk; v[k] = 0.999 * v[k] + 0.001 * gk**2
W[k] -= lrs[k] * (m[k] / (1 - 0.9**t)) / (np.sqrt(v[k] / (1 - 0.999**t)) + 1e-8)
if not np.isfinite(e).all() or np.abs(e).max() > 1e6: return np.nan
out = np.maximum(np.maximum(X @ W[0], 0) @ W[1], 0) @ W[2]
return float(np.mean((out - Y)**2))
widths = [32, 128, 512]
ks = list(range(-10, -3))
for mup in [0, 1]:
for n in widths:
row = [train(n, 2.0**k, mup) for k in ks]
best = ks[int(np.nanargmin(row))]
print('μP' if mup else 'SP', n, best, [None if not np.isfinite(v) else round(v, 4) for v in row], flush=True)
# SP 32 -6 [0.2174, 0.1626, 0.1189, 0.0715, 0.0588, 0.075, 0.114]
# SP 128 -7 [0.0539, 0.0154, 0.0053, 0.0052, 0.0184, 0.0385, 0.088]
# SP 512 -9 [0.001, 0.0008, 0.001, 0.0098, 0.0362, 0.0949, 0.2646]
# μP 32 -6 [0.2174, 0.1626, 0.1189, 0.0715, 0.0588, 0.075, 0.114]
# μP 128 -6 [0.1783, 0.1153, 0.0423, 0.0106, 0.006, 0.0216, 0.0322]
# μP 512 -6 [0.1624, 0.1094, 0.0326, 0.0057, 0.0019, 0.0057, 0.0131]
SP에서 최적 log₂ lr 은 −6, −7, −9 로 폭을 넓힐수록 왼쪽으로 밀린다. μP에서는 세 폭 모두 −6 이고, 넓을수록 같은 학습률에서 손실이 더 낮다.
수확
“표준 설정에서는 최적 학습률이 폭을 따라 움직이고, μP에서는 제자리에 있다. μP의 원래 유도는 '특징이 Θ(1)로 변하게’이고, 스펙트럼 조건 ‖ΔW‖₂ ∝ √(fan_out/fan_in) 으로 요약된다. '피셔 거리를 일정하게’는 그것을 이 책의 눈으로 본 해석이다.”
문제 9. 두 노름으로 잰 은닉층
폭 n 인 은닉층 W(n × n)를 성분마다 분산 1/n 인 정규분포로 초기화한다. (가) n = 256, 1024, 4096 에서 프로베니우스 노름 ‖W‖F(성분 제곱합의 제곱근)와 스펙트럼 노름 ‖W‖₂ 를 구하라. 위의 스펙트럼 조건과 맞는 것은 어느 쪽인가? (나) μP 에서 은닉층으로 돌아오는 역전파 신호 δ 는 좌표마다 크기가 1/n 정도이고, 은닉층에 들어오는 특징 h 는 좌표마다 ±1 정도다. SGD 한 스텝 ΔW = −lr·δhᵀ 에서 lr = 1 일 때 성분 하나의 크기와 ‖ΔW‖₂ 는 n 에 따라 어떻게 되는가? (다) 이 장 첫 절의 층(폭 n, 출력 하나)은 표 3 의 어느 칸이고, 그 칸의 SGD 학습률 처방은 그 층에서 오차가 끝없이 커지지 않는 범위 0 < lr < 2/n 과 맞는가?
함께 풀기

(가)는 성분 하나가 1/√n 크기고 성분이 n² 개니까 프로베니우스는 √(n² × 1/n) = √n 이요. 256 이면 16, 4096 이면 64. 넓을수록 큰 행렬이에요.

스펙트럼 노름도 그만큼 커지나요? 코드로 재 봐요.

np.linalg.norm(W, 2) 로 재니까 1.97, 1.99, 2.00 이요. 프로베니우스는 네 배가 됐는데 스펙트럼은 2 근처에 멈춰 있어요.

프로베니우스는 모든 방향으로 늘리는 배율을 제곱해 더한 거라 방향 수만큼 커지고, 스펙트럼은 가장 많이 늘리는 한 방향만 봐. 무작위 행렬은 어느 한 방향으로 몰아서 늘리지 않으니까 한 방향은 2 정도에 머물러. 정사각 층이면 √(fan_out/fan_in) = 1 이니까 스펙트럼 조건 Θ(1) 과 맞는 건 스펙트럼 노름이야.

업데이트 쪽은요? 성분 하나가 얼마나 작은지부터 봐요.

성분 하나가 lr × (1/n) × 1 = 1/n 이에요. n = 4096 이면 0.00024. 넓은 층에서는 업데이트가 거의 없는 거나 마찬가지죠.

ΔW 가 바깥곱이라 랭크가 1 이잖아. 랭크 1 행렬은 한 방향만 늘리니까 스펙트럼 노름이 ‖δ‖·‖h‖ = (√n/n)·√n = 1 이야. n 과 상관없이 1. 성분은 1/n 로 작아지지만 그 n² 개가 한 방향으로 줄을 서 있어서, 그 방향으로는 Θ(1) 만큼 늘려.

첫 절에서 본 정렬된 합이네요. 조각은 작아도 같은 방향이면 다 더해지는 거. 그래서 표 3 에서 은닉층 SGD 학습률이 1 이구나.

첫 절의 폭 n 층은 μP의 처방과 어떻게 이어져요?

이건 fan_in = n 인 출력층이니까 SGD 학습률 1/fan_in 이요. 표랑 맞아요. 끝없이 커지지 않는 범위가 2/n 아래였으니, 1/n 은 그 안에서 한 스텝에 오차를 0 으로 보내는 값이고요.

예산 문제의 조원들처럼, 다들 같은 방향으로 한 줄씩만 고쳐도 보고서 전체로는 n 줄이 한쪽으로 바뀌는 거죠. 조장은 조원 수만큼 수정 폭을 줄여야 하고요.
문제 10. 폭 32 에서 찾은 학습률을 폭 2048 로
위 파이썬의 μP 실험에서 기준 폭 32 의 가장 좋은 학습률은 2⁻⁶ 이었다. 같은 코드의 μP 규칙(Adam, 표 3)으로 폭 2048 모델을 만들 때, 입력층·은닉층·출력층에 실제로 들어가는 학습률과 출력층 초기화의 표준편차는 각각 얼마인가? 입력 차원 d = 16 은 그대로다.
함께 풀기

표 3 에서 Adam 의 은닉층 학습률은 1/fan_in 이니까 2⁻⁶ × 1/2048 = 2⁻¹⁷ 이요. 출력층도 2⁻¹⁷ 이고, 입력층은 fan_in 이 입력 차원 16 이라 폭과 상관없으니 2⁻⁶ 그대로고요.

그 계산대로라면 폭 32 일 때 은닉층에 실제로 들어간 학습률은 얼마였어요?

2⁻⁶ × 1/32 = 2⁻¹¹ 이어야 하는데… 코드를 보면 lrs = [lr, lr * s, lr * s] 에 s = 32/n 이라서, 폭 32 에서는 s = 1 이고 은닉층에도 2⁻⁶ 이 그대로 들어갔네요. 표 3 의 1/fan_in 은 폭에 따라 얼마나 줄이느냐는 비례만 말하고, 출발점은 기준 폭에서 찾은 값이었어요.

그럼 폭 2048 은 s = 32/2048 = 1/64 라서 은닉층과 출력층이 2⁻⁶/64 = 2⁻¹², 입력층은 2⁻⁶ 이요.

출력층 초기화도 같은 s 를 써. 표준 설정의 표준편차 1/√2048 ≈ 0.022 에 1/64 을 곱해서 약 0.00035. 위젯 가로축의 lr 은 폭 32 에서 찾은 기준값 하나고, 층마다 실제 값은 μP 가 폭에 맞춰 줄여 주는 거네.

조별 과제에서 「조원 네 명일 때 한 사람 몫」을 기준으로 정해 두면, 조원이 늘어날 때 각자 몫을 그 기준에서 비율로 줄이는 거랑 같네요. 처음부터 「전체 분량 ÷ 사람 수」를 새로 계산하는 게 아니라요.