|

모델 학습과 최적화 완전 정복: 경사하강법·역전파·AdamW·Muon·AMP까지

AI 모델은 데이터를 한 번 읽는다고 똑똑해지지 않습니다. 입력으로 예측을 만들고, 정답과의 차이를 숫자로 계산하고, 그 오차가 줄어드는 방향으로 수백만~수천억 개의 파라미터를 조금씩 바꾸는 과정을 반복해야 합니다.

이 반복을 구성하는 핵심은 다음과 같습니다.

데이터 입력
  → 순전파(Forward Pass)
  → 손실 계산(Loss)
  → 역전파(Backpropagation)
  → 기울기 계산(Gradient)
  → 옵티마이저 업데이트(Optimizer Step)
  → 다음 미니배치

기존의 입문 설명은 대개 손실함수, 경사하강법, 역전파, Adam을 각각 따로 설명합니다. 하지만 실전에서는 이 요소들이 하나의 학습 루프 안에서 어떤 순서로 연결되는지가 더 중요합니다.

이 글에서는 다음 내용을 하나의 흐름으로 정리합니다.

  • 손실함수와 목적함수의 정확한 차이
  • 경사하강법이 실제 숫자를 어떻게 바꾸는지
  • Epoch, Batch, Step, Iteration의 차이
  • 역전파와 자동미분이 하는 일
  • 학습률 warmup과 cosine decay가 필요한 이유
  • Adam보다 AdamW가 널리 쓰이는 이유
  • 2026년 새 이슈인 Muon 옵티마이저
  • Gradient Clipping, AMP, BF16, 기울기 누적
  • Activation Checkpointing, torch.compile, 분산학습
  • 바로 재사용할 수 있는 PyTorch 학습 루프

손실함수 종류 자체가 궁금하다면 손실함수 완전 정복을 함께 참고하세요. 이 글은 손실함수의 종류보다 모델을 실제로 학습시키는 과정과 최적화 전략에 집중합니다.


핵심 요약

  1. 순전파는 현재 파라미터로 예측을 계산하는 과정입니다.
  2. 손실함수는 한 샘플 또는 미니배치의 예측 오차를 숫자로 바꿉니다.
  3. 역전파는 연쇄법칙을 이용해 각 파라미터가 손실에 미친 영향을 계산합니다.
  4. 옵티마이저는 계산된 기울기를 사용해 파라미터를 업데이트합니다.
  5. 학습률은 옵티마이저 종류만큼 중요하며, 현대 모델은 보통 warmup 후 cosine decay를 사용합니다.
  6. 범용 기본값으로는 Adam보다 AdamW가 더 적절한 경우가 많습니다.
  7. 2026년에는 행렬 파라미터의 업데이트를 직교화하는 Muon이 대규모 언어모델 학습의 새 선택지로 부상했고, PyTorch 안정 문서에도 torch.optim.Muon이 포함됐습니다.
  8. 메모리가 부족할 때는 배치 크기만 줄이지 말고 AMP, gradient accumulation, activation checkpointing을 조합해야 합니다.
  9. 최적화는 학습 손실만 빨리 낮추는 기술이 아닙니다. 최종 목표는 검증 성능, 안정성, 학습 비용, 재현성을 함께 개선하는 것입니다.

1. 모델 학습은 무엇을 바꾸는가

신경망은 입력을 출력으로 바꾸는 함수입니다.

ŷ = fθ(x)
  • x: 입력 데이터
  • ŷ: 모델의 예측값
  • θ: 모델이 학습해야 할 모든 파라미터

선형회귀라면 θ는 가중치 w와 편향 b 정도입니다.

ŷ = wx + b

하지만 딥러닝 모델에서는 θ가 수백만~수천억 개의 숫자로 구성될 수 있습니다. 학습이란 데이터 자체를 저장하는 과정이 아니라, 주어진 데이터에서 손실이 작아지도록 θ를 반복해서 수정하는 과정입니다.

학습과 추론의 차이

구분학습 Training추론 Inference
목적파라미터를 개선학습된 파라미터로 결과 생성
역전파수행수행하지 않음
기울기 저장필요보통 불필요
Dropout활성화비활성화
BatchNorm배치 통계 사용저장된 통계 사용
메모리 사용상대적으로 작음

PyTorch에서는 이 차이를 코드로 명확히 구분해야 합니다.

# 학습
model.train()

# 검증·추론
model.eval()
with torch.inference_mode():
    output = model(x)

model.eval()만 호출한다고 기울기 계산이 꺼지는 것은 아닙니다. 검증에서는 torch.inference_mode() 또는 torch.no_grad()를 함께 사용해야 불필요한 계산 그래프와 메모리 사용을 줄일 수 있습니다.


2. 한 번의 학습 스텝에서 일어나는 7단계

모델 학습을 이해하려면 먼저 한 번의 optimizer step을 쪼개서 봐야 합니다.

2-1. 미니배치를 가져온다

inputs, targets = next(iter(train_loader))

전체 데이터를 한 번에 처리하지 않고, 보통 일정 크기의 미니배치로 나눠 학습합니다.

2-2. 이전 기울기를 초기화한다

optimizer.zero_grad(set_to_none=True)

PyTorch의 기울기는 기본적으로 누적됩니다. 초기화하지 않으면 이전 미니배치의 기울기가 다음 미니배치에 더해집니다.

set_to_none=True는 기울기를 0 텐서로 채우는 대신 None으로 설정합니다. 많은 경우 메모리 쓰기와 연산을 줄일 수 있지만, .grad == 0을 전제로 한 사용자 코드가 있다면 동작 차이를 확인해야 합니다.

2-3. 순전파한다

outputs = model(inputs)

현재 파라미터로 예측을 계산합니다. 이때 PyTorch의 autograd는 연산 관계를 계산 그래프로 기록합니다.

2-4. 손실을 계산한다

loss = criterion(outputs, targets)

손실은 모델 출력과 정답의 차이를 하나의 스칼라 값으로 바꿉니다.

2-5. 역전파한다

loss.backward()

손실을 각 파라미터로 미분한 값이 parameter.grad에 저장됩니다.

∂Loss / ∂θ

2-6. 필요하면 기울기를 안정화한다

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

기울기가 지나치게 커지는 문제를 완화합니다.

2-7. 파라미터를 업데이트한다

optimizer.step()

옵티마이저가 기울기를 읽고 파라미터를 실제로 변경합니다.

가장 기본적인 PyTorch 학습 루프

model.train()

for inputs, targets in train_loader:
    inputs = inputs.to(device)
    targets = targets.to(device)

    optimizer.zero_grad(set_to_none=True)

    outputs = model(inputs)
    loss = criterion(outputs, targets)

    loss.backward()
    optimizer.step()

이 다섯 줄이 딥러닝 학습의 뼈대입니다.

zero_grad
→ forward
→ loss
→ backward
→ step

3. 손실함수, 비용함수, 목적함수는 어떻게 다른가

3-1. 한 샘플의 손실

한 샘플 i의 손실은 다음처럼 표현할 수 있습니다.

ℓᵢ = ℓ(fθ(xᵢ), yᵢ)

3-2. 미니배치 또는 전체 데이터의 평균 손실

Jdata(θ) = (1/N) Σᵢ ℓ(fθ(xᵢ), yᵢ)

전통적으로 한 샘플의 오차를 loss, 전체 또는 미니배치 평균을 cost라고 구분하기도 합니다. 하지만 현대 라이브러리와 논문에서는 둘을 엄격히 나누지 않고 모두 loss라고 부르는 경우가 많습니다.

3-3. 정규화까지 포함한 목적함수

실제 최적화 대상은 데이터 손실만이 아닐 수 있습니다.

J(θ) = Jdata(θ) + λR(θ)
  • Jdata: 예측 오차
  • R(θ): 가중치 크기, 희소성 등을 제어하는 정규화 항
  • λ: 정규화 강도

따라서 다음 표현을 구분하면 이해가 쉬워집니다.

용어의미
Loss샘플 또는 미니배치에서 계산한 오차
Data Loss데이터 적합도만 반영한 손실
Regularization모델 복잡도에 부과하는 제약
Objective옵티마이저가 실제로 최소화하는 전체 목적함수

손실이 낮다고 항상 좋은 모델은 아니다

훈련 손실은 낮지만 검증 손실이 높다면 과적합일 수 있습니다. 반대로 정규화, Dropout, 데이터 증강이 강하면 훈련 손실이 검증 손실보다 높게 나타날 수도 있습니다.

따라서 최적화의 목표를 다음처럼 잡아야 합니다.

훈련 손실 최소화만 추적 ❌
검증 성능 + 안정성 + 비용 + 재현성 함께 추적 ✅

4. 경사하강법을 실제 숫자로 이해하기

경사하강법은 손실함수의 기울기와 반대 방향으로 파라미터를 이동합니다.

θt+1 = θt − η∇θJ(θt)
  • η: 학습률 learning rate
  • ∇θJ: 현재 파라미터에서 손실이 가장 빠르게 증가하는 방향
  • −∇θJ: 손실이 감소하는 방향

4-1. 가장 단순한 예제

다음 목적함수를 최소화한다고 가정해 보겠습니다.

J(w) = (w − 3)²

최솟값은 w = 3입니다.

미분하면 다음과 같습니다.

dJ/dw = 2(w − 3)

초깃값 w = 0, 학습률 η = 0.1로 시작하면 첫 업데이트는 다음과 같습니다.

gradient = 2(0 − 3) = −6
wnew = 0 − 0.1 × (−6) = 0.6

두 번째 업데이트는 다음과 같습니다.

gradient = 2(0.6 − 3) = −4.8
wnew = 0.6 − 0.1 × (−4.8) = 1.08

파라미터는 0 → 0.6 → 1.08 → ... → 3으로 이동합니다.

w = 0.0
learning_rate = 0.1

for step in range(10):
    loss = (w - 3.0) ** 2
    gradient = 2.0 * (w - 3.0)
    w -= learning_rate * gradient

    print(f"step={step:02d}, w={w:.6f}, loss={loss:.6f}")

4-2. 왜 정확한 최저점으로 한 번에 가지 않는가

현실의 신경망 손실 표면은 단순한 그릇 모양이 아닙니다.

  • 파라미터가 매우 많음
  • 미니배치마다 기울기가 달라짐
  • 평평한 구간과 가파른 구간이 공존
  • 안장점과 좁은 골짜기가 존재
  • 데이터 노이즈와 라벨 오류가 있음

그래서 한 번의 정확한 해를 구하기보다, noisy gradient를 반복적으로 계산하며 충분히 좋은 영역으로 이동합니다.

4-3. 지역 최솟값보다 더 중요한 것

딥러닝에서는 “나쁜 지역 최솟값에 갇히는가?”보다 다음 문제가 실무적으로 더 중요합니다.

  • 학습률이 너무 커 발산하는가
  • 기울기가 사라지거나 폭발하는가
  • 훈련 손실은 감소하지만 일반화가 나쁜가
  • 제한된 시간과 GPU 비용 안에 목표 성능에 도달하는가

즉, 최적화는 수학적 최솟값만 찾는 문제가 아니라 안정적이고 효율적인 학습 경로를 설계하는 문제입니다.


5. Batch, Mini-batch, Epoch, Step, Iteration 차이

이 용어를 혼동하면 학습률 스케줄러와 로그가 모두 꼬일 수 있습니다.

예제 조건

학습 데이터: 10,000개
batch_size: 100
epochs: 5

한 Epoch에 필요한 미니배치 수는 다음과 같습니다.

10,000 / 100 = 100 steps per epoch

전체 optimizer update 수는 다음과 같습니다.

100 steps × 5 epochs = 500 optimizer steps
용어의미
Sample데이터 한 개
Batch한 번에 처리하는 데이터 묶음
Mini-batch전체보다 작은 배치. 실무에서 흔히 batch라고 부름
Epoch전체 학습 데이터를 한 번 모두 사용
Step보통 한 번의 optimizer update
Iteration문맥에 따라 미니배치 처리 또는 optimizer step

Gradient Accumulation을 쓰면 Step 수가 달라진다

미니배치 16개를 네 번 누적한 뒤 한 번 업데이트한다면 유효 배치 크기는 대략 64가 됩니다.

effective_batch_size
= micro_batch_size × accumulation_steps × number_of_devices

예를 들어 다음과 같습니다.

micro batch = 8
accumulation = 4
GPU = 2
유효 배치 크기 = 8 × 4 × 2 = 64

그러나 기울기 누적이 물리적 큰 배치와 항상 완전히 같은 것은 아닙니다.

  • BatchNorm은 각 micro-batch의 통계를 사용
  • Dropout 마스크가 micro-batch마다 달라짐
  • 데이터 증강 결과가 달라질 수 있음
  • 학습률 스케줄을 micro-batch마다 진행하면 실제 update 수와 불일치

따라서 scheduler는 보통 optimizer가 실제로 step한 횟수를 기준으로 진행해야 합니다.


6. 역전파는 무엇을 계산하는가

역전파는 옵티마이저가 아닙니다. 역전파는 손실을 각 파라미터로 미분한 기울기를 효율적으로 계산하는 방법입니다.

옵티마이저는 그 기울기를 사용해 파라미터를 업데이트합니다.

역전파: gradient 계산
옵티마이저: gradient를 사용해 parameter 변경

6-1. 연쇄법칙

간단한 계산을 생각해 보겠습니다.

h = wx
ŷ = vh
L = (ŷ − y)²

w가 손실에 미친 영향은 중간 연산을 따라 계산합니다.

∂L/∂w
= ∂L/∂ŷ × ∂ŷ/∂h × ∂h/∂w

이것이 연쇄법칙입니다.

6-2. 숫자로 계산하기

x = 2
w = 0.5
v = 0.3
y = 1

순전파는 다음과 같습니다.

h = wx = 1.0
ŷ = vh = 0.3
L = (0.3 − 1.0)² = 0.49

기울기는 다음과 같습니다.

∂L/∂ŷ = 2(ŷ − y) = −1.4
∂ŷ/∂v = h = 1.0
∂L/∂v = −1.4

∂ŷ/∂h = v = 0.3
∂h/∂w = x = 2.0
∂L/∂w = −1.4 × 0.3 × 2.0 = −0.84

학습률이 0.1이면 다음처럼 업데이트됩니다.

vnew = 0.3 − 0.1 × (−1.4) = 0.44
wnew = 0.5 − 0.1 × (−0.84) = 0.584

6-3. PyTorch autograd

import torch

x = torch.tensor(2.0)
y_true = torch.tensor(1.0)

w = torch.tensor(0.5, requires_grad=True)
v = torch.tensor(0.3, requires_grad=True)

h = w * x
y_pred = v * h
loss = (y_pred - y_true) ** 2

loss.backward()

print(f"loss: {loss.item():.4f}")       # 0.4900
print(f"dL/dw: {w.grad.item():.4f}")   # -0.8400
print(f"dL/dv: {v.grad.item():.4f}")   # -1.4000

PyTorch의 torch.autograd는 스칼라 목적함수에 대한 자동미분을 제공합니다. 사용자가 모든 미분식을 손으로 구현하지 않아도, 순전파에서 기록된 계산 그래프를 따라 역방향으로 기울기를 계산합니다.

6-4. 기울기는 자동으로 누적된다

loss1.backward()
loss2.backward()

이렇게 두 번 호출하면 두 기울기가 더해집니다. 이것은 gradient accumulation을 구현할 때 유용하지만, 의도하지 않았다면 버그가 됩니다.

optimizer.zero_grad(set_to_none=True)

를 매 update 전 적절한 시점에 호출해야 합니다.

6-5. 그래프를 계속 보존하면 메모리가 증가한다

기본적으로 backward()가 끝나면 중간 그래프는 해제됩니다. 같은 그래프를 다시 역전파하려고 무조건 retain_graph=True를 쓰면 메모리 사용량이 크게 늘어날 수 있습니다.

loss.backward(retain_graph=True)  # 정말 필요한 경우에만

대부분의 일반 학습 루프에는 retain_graph=True가 필요하지 않습니다.


7. 학습률은 가장 중요한 하이퍼파라미터다

옵티마이저가 좋아도 학습률이 맞지 않으면 학습은 실패합니다.

너무 큰 학습률

  • 손실이 급격히 출렁임
  • 최적 영역을 계속 지나침
  • NaN, Inf 발생
  • 기울기 폭발

너무 작은 학습률

  • 손실 감소가 매우 느림
  • 제한된 epoch 안에 충분히 학습하지 못함
  • 평평한 영역에서 거의 움직이지 않음

7-1. 고정 학습률의 한계

학습 초반에는 비교적 큰 이동이 필요하지만, 후반에는 작은 이동으로 미세 조정해야 합니다. 따라서 하나의 학습률을 끝까지 유지하는 것보다 스케줄러를 사용하는 경우가 많습니다.

7-2. Warmup

학습 시작부터 큰 학습률을 사용하면 초기화된 가중치와 불안정한 통계 때문에 손실이 폭발할 수 있습니다.

Warmup은 초기 몇 step 동안 학습률을 작은 값에서 목표값까지 서서히 올립니다.

0.00003 → 0.00006 → ... → 0.0003

특히 다음 상황에서 자주 사용합니다.

  • Transformer
  • 대규모 배치 학습
  • 사전학습 모델 파인튜닝
  • 매우 깊은 모델
  • 혼합 정밀도 학습

7-3. Cosine Decay

Warmup 이후에는 cosine 곡선을 따라 학습률을 부드럽게 줄일 수 있습니다.

ηt = ηmin + 0.5(ηmax − ηmin)(1 + cos(πt/T))

장점은 다음과 같습니다.

  • 갑작스러운 학습률 변화가 적음
  • 후반 미세 조정에 유리
  • Transformer, ViT, LLM에서 널리 사용

7-4. ReduceLROnPlateau는 언제 쓰는가

검증 지표가 개선되지 않을 때 학습률을 낮추는 방식입니다.

scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
    optimizer,
    mode="min",
    factor=0.5,
    patience=3,
)

scheduler.step(val_loss)

데이터가 작고 epoch 단위 검증이 명확한 문제에는 편리합니다. 반면 대규모 사전학습처럼 총 update 수가 미리 정해진 환경에서는 warmup + cosine schedule이 더 다루기 쉽습니다.

7-5. 스케줄러 호출 순서

대부분의 PyTorch scheduler는 다음 순서를 사용합니다.

optimizer.step()
scheduler.step()

Gradient accumulation을 쓰는 경우에는 micro-batch마다 scheduler.step()을 호출하지 말고, 실제 optimizer.step()이 실행될 때만 진행해야 합니다.


8. 옵티마이저의 계보: SGD에서 AdamW까지

8-1. SGD

기본 확률적 경사하강법은 현재 미니배치의 기울기를 그대로 사용합니다.

θt+1 = θt − ηgt

장점:

  • 단순함
  • 메모리 사용량이 적음
  • 잘 튜닝하면 좋은 일반화 성능

단점:

  • 학습률과 스케줄에 민감
  • 좁고 긴 골짜기에서 진동
  • 초기 수렴이 느릴 수 있음

8-2. Momentum

이전 업데이트 방향을 누적해 관성을 만듭니다.

vt = βvt−1 + gt
θt+1 = θt − ηvt

일반적으로 β = 0.9가 흔한 시작점입니다.

SGD + Momentum은 여전히 이미지 분류 등에서 강력한 기준선입니다. 다만 충분한 튜닝과 긴 학습 스케줄이 필요할 수 있습니다.

8-3. RMSprop

기울기 제곱의 이동평균을 사용해 파라미터별 학습률을 조절합니다.

st = βst−1 + (1 − β)gt²
θt+1 = θt − ηgt / (√st + ε)

RNN 학습과 일부 비정상적 손실 지형에서 역사적으로 많이 사용됐습니다.

8-4. Adam

Adam은 다음 두 통계를 함께 사용합니다.

  • 기울기의 1차 모멘트: 방향과 관성
  • 기울기 제곱의 2차 모멘트: 파라미터별 스케일 조절
mt = β1mt−1 + (1 − β1)gt
vt = β2vt−1 + (1 − β2)gt²

장점:

  • 초기 수렴이 빠른 편
  • 하이퍼파라미터에 비교적 덜 민감
  • 희소하거나 스케일이 다른 기울기에 대응

단점:

  • SGD보다 optimizer state 메모리가 큼
  • weight decay를 단순 L2 항으로 넣으면 의도한 감쇠와 다르게 작동할 수 있음
  • 최종 일반화가 항상 SGD보다 좋은 것은 아님

9. Adam과 AdamW의 핵심 차이

현대 Transformer, ViT, LLM 학습에서는 일반 Adam보다 AdamW가 기본 선택으로 자주 사용됩니다.

9-1. L2 정규화와 Weight Decay는 항상 같은가

일반 SGD에서는 목적함수에 L2 항을 더하는 것과 가중치를 직접 감쇠하는 것이 사실상 같은 형태로 연결됩니다.

하지만 Adam처럼 파라미터별 적응형 스케일을 사용하는 옵티마이저에서는 L2 gradient도 적응형 정규화의 영향을 받습니다. 그 결과 가중치를 직접 일정 비율 줄이는 weight decay와 동작이 달라집니다.

AdamW는 weight decay를 gradient update와 분리합니다.

θ ← θ − η × AdamUpdate
θ ← θ − ηλθ

또는 한 줄로 다음처럼 볼 수 있습니다.

θ ← (1 − ηλ)θ − η × AdamUpdate

PyTorch의 AdamW 문서도 weight decay가 momentum과 variance에 누적되지 않는다고 설명합니다.

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=3e-4,
    weight_decay=0.01,
)

9-2. 모든 파라미터에 weight decay를 적용해야 하는가

실무에서는 보통 다음 파라미터를 분리합니다.

  • 선형층·합성곱층의 가중치: decay 적용
  • bias: decay 제외
  • LayerNorm·BatchNorm의 scale과 bias: decay 제외
  • embedding: 모델과 레시피에 따라 다름
decay_params = []
no_decay_params = []

for name, param in model.named_parameters():
    if not param.requires_grad:
        continue

    if param.ndim >= 2:
        decay_params.append(param)
    else:
        no_decay_params.append(param)

optimizer = torch.optim.AdamW(
    [
        {"params": decay_params, "weight_decay": 0.01},
        {"params": no_decay_params, "weight_decay": 0.0},
    ],
    lr=3e-4,
)

이 규칙은 좋은 시작점이지 절대 법칙은 아닙니다. 사전학습 모델이 사용한 원래 레시피가 있다면 그 설정을 우선 확인해야 합니다.

9-3. AdamW의 기본값을 그대로 쓰면 되는가

PyTorch AdamW의 문서 기본값은 대략 다음과 같습니다.

lr = 0.001
betas = (0.9, 0.999)
eps = 1e-8
weight_decay = 0.01

하지만 모델과 작업에 따라 최적값은 크게 달라집니다.

  • 작은 MLP: 1e-3이 잘 작동할 수 있음
  • Transformer 파인튜닝: 1e-5 ~ 5e-4 범위에서 자주 탐색
  • 대규모 사전학습: 모델 크기, batch, token budget에 맞춘 별도 scaling 필요

숫자를 외우기보다 학습률·weight decay·warmup 비율을 함께 실험하는 것이 중요합니다.


10. 2026년 새 이슈: Muon 옵티마이저

Muon은 MomentUm Orthogonalized by Newton-Schulz의 약자로, 신경망의 2차원 행렬 파라미터 업데이트를 직교화하는 옵티마이저입니다.

일반 AdamW가 각 좌표별 1차·2차 모멘트를 추적한다면, Muon은 선형층 가중치가 행렬이라는 구조를 활용합니다.

10-1. 핵심 아이디어

  1. 기울기에 momentum을 적용
  2. Newton–Schulz 반복으로 업데이트 행렬을 직교화
  3. 직교화된 방향으로 가중치를 업데이트

직관적으로는 특정 행이나 열 방향에 업데이트가 과도하게 몰리지 않도록, 행렬 전체의 업데이트 방향을 더 균형 있게 만드는 접근입니다.

10-2. 왜 2026년에 주목받는가

Muon은 2024년 공개된 이후 소규모 언어모델 학습 기록과 연구에서 주목받기 시작했습니다. 이후 Moonshot AI는 Kimi K2가 Muon 옵티마이저로 학습됐다고 공개했습니다.

더 중요한 변화는 PyTorch 안정 문서에 torch.optim.Muon API가 포함됐다는 점입니다. PyTorch 2.13 문서는 Muon을 2D hidden-layer 파라미터용 옵티마이저로 설명하며, bias·embedding·기타 비2D 파라미터에는 AdamW 같은 표준 옵티마이저를 사용하라고 안내합니다.

10-3. Muon을 모든 파라미터에 쓰면 안 된다

Muon은 2D 행렬 파라미터를 대상으로 합니다.

  • Linear weight: Muon 후보
  • 2D projection matrix: Muon 후보
  • bias: AdamW
  • LayerNorm parameter: AdamW
  • embedding: AdamW
  • 1D·3D·4D 파라미터: 기본적으로 AdamW 등 별도 처리

PyTorch 문서 예시는 다음 구조를 사용합니다.

muon_params = [
    p for p in model.parameters()
    if p.requires_grad and p.ndim == 2
]

other_params = [
    p for p in model.parameters()
    if p.requires_grad and p.ndim != 2
]

optim_muon = torch.optim.Muon(
    muon_params,
    lr=0.02,
    momentum=0.95,
)

optim_adamw = torch.optim.AdamW(
    other_params,
    lr=3e-4,
    weight_decay=0.01,
)

10-4. 지금 AdamW를 Muon으로 바꿔야 하는가

아직 모든 프로젝트에서 무조건 교체할 단계는 아닙니다.

Muon을 검토할 만한 경우:

  • Transformer나 MLP의 대규모 사전학습
  • 2D 선형층 파라미터가 큰 비중을 차지
  • AdamW 기준선과 충분히 비교할 수 있음
  • 학습 처리량보다 sample efficiency와 최종 비용이 중요

AdamW를 유지하는 편이 나은 경우:

  • 소규모 데이터 파인튜닝
  • 기존 검증된 레시피가 있음
  • 모델 파라미터가 2D 행렬 중심이 아님
  • 재현성과 라이브러리 호환성이 더 중요
  • 하이퍼파라미터 탐색 예산이 부족

Muon은 “AdamW의 완전한 후계자”라기보다, 행렬 구조를 활용하는 새로운 최적화 계열로 보는 편이 정확합니다.

10-5. Lion, Sophia 같은 새 옵티마이저는

최근에는 다음과 같은 옵티마이저도 연구됐습니다.

  • Lion: 부호 기반 momentum update, Adam보다 optimizer state가 적음
  • Sophia: 곡률 정보를 근사해 더 큰 안정적 step을 목표
  • Schedule-Free 계열: 별도 학습률 decay schedule 의존을 줄이려는 접근
  • Muon·NorMuon 계열: 행렬 업데이트 직교화와 적응형 스케일 결합

그러나 새 옵티마이저의 논문 결과가 모든 데이터셋, 모델, 배치 크기에 그대로 재현되는 것은 아닙니다. 실무에서는 다음 순서가 안전합니다.

AdamW 기준선 확보
→ 동일한 데이터·seed·token budget으로 비교
→ 최종 성능뿐 아니라 wall-clock·메모리·안정성 비교
→ 여러 seed에서 재현되는지 확인

10-6. 2026년 상반기 후속 이슈: Muon의 변형과 분산학습

2026년 7월 기준으로 Muon 연구는 단순히 “AdamW보다 빠른가”를 비교하는 단계를 넘어, 정규화·스케줄 제거·파인튜닝 호환성·분산 실행 비용을 해결하는 방향으로 확장되고 있습니다.

Muon+와 정규화 계열

Muon+는 Newton–Schulz 직교화 뒤에 정규화 단계를 하나 더 추가하는 접근입니다. Variance-Adaptive Muon, NorMuon, Muown 같은 변형도 행 또는 뉴런 단위의 스케일 불균형을 줄이는 데 초점을 둡니다.

이 연구들은 Muon의 업데이트가 행렬 전체에서는 균형 있게 보여도, 개별 행·뉴런 수준에서는 크기 차이가 남을 수 있다는 문제를 다룹니다. 다만 대부분 최신 연구 또는 프리프린트 단계이므로, 이름만 보고 기본 옵티마이저를 교체하기보다 공개 코드와 동일 조건 재현 결과를 먼저 확인해야 합니다.

Schedule-Free NorMuon

2026년 5월 공개된 SF-NorMuon은 학습 종료 시점을 미리 정해 cosine decay 같은 스케줄을 정교하게 맞추지 않아도, 학습 도중 여러 시점에서 사용할 수 있는 체크포인트를 얻는 것을 목표로 합니다.

이 접근은 다음 상황에서 의미가 있습니다.

  • 전체 token budget이 자주 바뀌는 실험
  • 중간 체크포인트도 실제 배포 후보가 되는 continual training
  • scheduler 튜닝 비용을 줄이고 싶은 대규모 사전학습

하지만 “schedule-free”는 학습률 설정이 중요하지 않다는 뜻이 아닙니다. base learning rate, weight decay, averaging 방식은 여전히 검증해야 합니다.

Adam 사전학습 모델을 Muon으로 바로 파인튜닝할 때의 불일치

2026년 5월 연구에서는 Adam 계열로 사전학습한 모델을 파인튜닝 단계에서 곧바로 Muon으로 바꾸면 성능이 저하될 수 있는 optimizer mismatch가 보고됐습니다.

따라서 Muon은 현재 기준으로 다음과 같이 구분해 접근하는 편이 안전합니다.

처음부터 사전학습: Muon 실험 가치가 큼
AdamW 사전학습 모델 파인튜닝: AdamW 기준선을 우선 유지
옵티마이저 전환 실험: 낮은 학습률·전환 구간·동일 예산 비교 필요

“사전학습에서 효과가 있었다”는 결과를 일반적인 LoRA·SFT·분류 파인튜닝에 그대로 적용하면 안 됩니다.

DMuon: 분산 환경에서 optimizer step 비용 줄이기

2026년 6월 공개된 DMuon은 Muon의 행렬 직교화가 분산학습에서 만드는 통신·계산 오버헤드를 줄이는 구현입니다. 논문은 기존 분산 학습 파이프라인에 통합하면서 optimizer step 지연을 AdamW에 가까운 수준으로 낮추는 것을 목표로 합니다.

이는 중요한 변화입니다. 대규모 모델에서 옵티마이저는 수렴 속도만 좋아서는 부족하고, 다음 조건도 만족해야 하기 때문입니다.

  • 여러 GPU와 노드에서 효율적으로 shard할 수 있는가
  • 통신량이 과도하게 늘지 않는가
  • FSDP·ZeRO·Megatron 계열과 결합 가능한가
  • optimizer step의 추가 FLOPs가 전체 wall-clock 이득을 상쇄하지 않는가

다만 DMuon 역시 2026년 6월 공개된 초기 연구이므로, 일반 프로젝트의 기본 선택으로 단정하기보다 대규모 사전학습에서 검토할 최신 후보로 보는 것이 적절합니다.


11. 학습을 안정화하는 핵심 기술

11-1. Gradient Clipping

기울기 전체 norm이 임계값을 넘으면 스케일을 줄입니다.

total_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=1.0,
    error_if_nonfinite=True,
)

PyTorch의 clip_grad_norm_은 모든 파라미터 기울기를 하나의 벡터처럼 보고 전체 norm을 계산한 뒤, 기울기를 in-place로 수정합니다.

유용한 상황:

  • RNN·LSTM
  • 긴 sequence Transformer
  • 강화학습
  • 학습 초반 loss spike
  • 혼합 정밀도에서 간헐적 overflow

주의할 점:

  • 너무 작은 max norm은 모든 update를 지나치게 축소
  • clipping이 근본 원인인 잘못된 학습률이나 데이터 오류를 숨길 수 있음
  • AMP의 GradScaler를 쓸 때는 gradient를 unscale한 후 clipping해야 함
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()

11-2. Weight Initialization

초기화는 활성값과 기울기 스케일을 결정합니다.

활성화흔한 초기화
ReLU 계열Kaiming/He Initialization
tanh, sigmoidXavier/Glorot Initialization
Transformer프레임워크·아키텍처 기본 레시피 우선

PyTorch 모듈에는 합리적인 기본 초기화가 들어 있으므로, 특별한 근거 없이 모든 가중치를 임의로 다시 초기화하는 것은 오히려 성능을 해칠 수 있습니다.

11-3. Normalization

  • BatchNorm: 배치 통계 사용, CNN에서 강력
  • LayerNorm: feature 차원 정규화, Transformer 표준
  • RMSNorm: 평균 제거 없이 RMS 기반 정규화, LLM에서 널리 사용

Normalization은 단순한 과적합 방지 기법이 아니라, 활성값과 기울기 스케일을 안정화해 최적화를 쉽게 만듭니다.

11-4. Early Stopping

검증 손실이 개선되지 않으면 학습을 중단합니다.

best_val_loss = float("inf")
patience = 5
bad_epochs = 0

for epoch in range(num_epochs):
    train_loss = train_one_epoch(...)
    val_loss = evaluate(...)

    if val_loss < best_val_loss:
        best_val_loss = val_loss
        bad_epochs = 0
        torch.save(model.state_dict(), "best_model.pt")
    else:
        bad_epochs += 1

    if bad_epochs >= patience:
        print(f"Early stopping at epoch {epoch + 1}")
        break

중요한 점은 마지막 epoch 모델이 아니라 가장 좋은 검증 성능의 체크포인트를 저장하는 것입니다.

11-5. EMA와 SWA

Exponential Moving Average는 학습 중 파라미터의 이동평균을 유지합니다.

θEMA ← βθEMA + (1 − β)θ

Diffusion 모델, 객체 탐지, 이미지 분류 등에서 평가 성능과 안정성을 높이는 데 사용됩니다. SWA는 학습 후반 여러 지점의 가중치를 평균해 더 평평한 해를 찾는 접근입니다.

다만 EMA 모델은 별도 파라미터 복사본을 유지하므로 메모리와 체크포인트 관리가 추가됩니다.


12. 혼합 정밀도: FP32, FP16, BF16, FP8

12-1. 왜 낮은 정밀도를 사용하는가

낮은 정밀도는 다음을 줄일 수 있습니다.

  • 텐서 메모리
  • 메모리 대역폭
  • 행렬곱 시간

현대 GPU의 Tensor Core는 FP16, BF16, FP8 같은 형식에서 높은 처리량을 제공합니다.

12-2. FP16

장점:

  • 널리 지원
  • 메모리와 연산량 절감

단점:

  • 표현 범위가 좁아 작은 gradient가 underflow될 수 있음
  • 큰 값은 overflow할 수 있음
  • 보통 GradScaler 필요

12-3. BF16

BF16은 FP16보다 유효숫자 정밀도는 낮지만 지수 범위가 FP32와 비슷합니다. 지원 GPU에서는 FP16보다 학습 안정성이 좋은 경우가 많고, 일반적으로 별도 loss scaling 의존이 적습니다.

12-4. FP8

FP8은 E4M3와 E5M2 같은 형식을 조합해 더 낮은 정밀도로 학습 처리량을 높이는 기술입니다. NVIDIA Transformer Engine 문서는 H100에서 FP8 지원이 도입됐고, Blackwell 세대에서 MXFP8·NVFP4 같은 더 낮은 정밀도 형식이 확장됐다고 설명합니다.

다만 FP8은 다음 요소에 민감합니다.

  • 하드웨어 지원
  • tensor scaling 전략
  • 연산별 dtype 선택
  • framework와 kernel 구현
  • 모델별 수치 안정성

따라서 일반적인 단일 GPU 학습에서는 BF16 또는 FP16 AMP부터 적용하고, FP8은 대규모 학습 인프라에서 검증하는 편이 안전합니다.

12-5. PyTorch AMP 기본 형태

use_amp = device.type == "cuda"
amp_dtype = (
    torch.bfloat16
    if use_amp and torch.cuda.is_bf16_supported()
    else torch.float16
)

scaler = torch.amp.GradScaler(
    "cuda",
    enabled=use_amp and amp_dtype == torch.float16,
)

with torch.autocast(
    device_type=device.type,
    dtype=amp_dtype,
    enabled=use_amp,
):
    outputs = model(inputs)
    loss = criterion(outputs, targets)

scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

현재 PyTorch 문서는 torch.amptorch.autocast 기반 API를 중심으로 설명합니다.


13. 메모리 부족을 해결하는 네 가지 방법

13-1. Batch Size 줄이기

가장 간단하지만 다음 부작용이 있습니다.

  • gradient noise 증가
  • BatchNorm 통계 불안정
  • GPU 사용률 감소
  • 학습률 재조정 필요

13-2. Gradient Accumulation

여러 micro-batch의 기울기를 누적한 뒤 한 번 update합니다.

accumulation_steps = 4
optimizer.zero_grad(set_to_none=True)

for step, (inputs, targets) in enumerate(train_loader):
    outputs = model(inputs)
    loss = criterion(outputs, targets)
    loss = loss / accumulation_steps

    loss.backward()

    should_step = (
        (step + 1) % accumulation_steps == 0
        or (step + 1) == len(train_loader)
    )

    if should_step:
        optimizer.step()
        optimizer.zero_grad(set_to_none=True)

loss / accumulation_steps를 하지 않으면 누적 step 수만큼 gradient 크기가 커집니다. 의도적으로 합산하려는 특별한 경우가 아니라면 평균을 맞춰야 합니다.

13-3. Activation Checkpointing

순전파 중간 activation을 모두 저장하지 않고, 역전파 때 필요한 부분을 다시 계산합니다.

메모리 감소
↔ 역전파 재계산으로 연산량 증가

PyTorch 문서는 activation checkpointing을 “compute를 memory와 교환하는 기술”로 설명하며, 현재는 use_reentrant=False 사용을 권장합니다.

from torch.utils.checkpoint import checkpoint

x = checkpoint(
    block,
    x,
    use_reentrant=False,
)

Dropout 등 난수 연산이 포함되면 RNG state 저장·복원 비용과 결정성 문제도 확인해야 합니다.

13-4. Optimizer State 줄이기

AdamW는 파라미터 외에 1차·2차 모멘트를 저장합니다. 대규모 모델에서는 optimizer state가 모델 가중치보다 더 큰 메모리를 차지할 수 있습니다.

대안:

  • 8-bit optimizer state
  • Adafactor처럼 factorized state 사용
  • ZeRO/FSDP로 optimizer state sharding
  • CPU offload
  • Muon·Lion처럼 다른 state 구조를 가진 옵티마이저 검토

단, optimizer state를 줄이면 통신량, 수치 안정성, 구현 복잡도가 달라질 수 있습니다.


14. 속도를 높이는 시스템 최적화

모델이 느린 이유가 옵티마이저 때문이라고 단정하면 안 됩니다. 실제 병목은 데이터 로딩, 커널 호출, 통신, 메모리 복사일 수 있습니다.

14-1. DataLoader 최적화

train_loader = DataLoader(
    dataset,
    batch_size=128,
    shuffle=True,
    num_workers=8,
    pin_memory=True,
    persistent_workers=True,
)

GPU가 데이터를 기다린다면 모델 최적화보다 먼저 데이터 파이프라인을 개선해야 합니다.

14-2. torch.compile

PyTorch는 torch.compile을 통해 Python/PyTorch 연산을 추적하고 최적화된 kernel로 컴파일할 수 있습니다.

model = torch.compile(model)

장점:

  • 코드 변경이 적음
  • 연산 fusion과 kernel 최적화 가능

주의:

  • 첫 실행에 컴파일 시간이 필요
  • 동적 control flow는 graph break를 만들 수 있음
  • 모든 모델에서 동일한 속도 향상을 보장하지 않음
  • 짧은 실험에서는 compile overhead가 이득보다 클 수 있음

PyTorch 튜토리얼도 추적이 어려운 코드는 오류보다 graph break로 이어져 최적화 기회를 잃을 수 있다고 설명합니다.

14-3. Foreach와 Fused Optimizer

PyTorch AdamW는 foreachfused 구현을 지원합니다.

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=3e-4,
    fused=True,
)

공식 문서는 일반적으로 fused가 수직·수평 fusion을 모두 사용해 이론적으로 가장 빠르다고 설명하지만, 지원 dtype·device와 안정성을 확인해야 합니다. foreach는 CUDA에서 기본 for-loop보다 빠른 경우가 많지만 중간 tensor list 때문에 peak memory가 늘 수 있습니다.

따라서 다음처럼 판단합니다.

속도 병목 → fused/foreach 벤치마크
메모리 병목 → foreach peak memory 확인
호환성 우선 → 기본값 유지

14-4. 다중 GPU

방식핵심 아이디어적합한 상황
DDPGPU마다 모델 복제, gradient 동기화모델이 GPU 한 장에 들어감
FSDP파라미터·gradient·optimizer state sharding모델이 한 GPU에 안 들어감
Tensor Parallel하나의 연산을 여러 GPU로 분할매우 큰 행렬 연산
Pipeline Parallel레이어 구간을 GPU별로 분할깊고 큰 모델
ZeROoptimizer·gradient·parameter 상태 단계별 분할대규모 분산학습

분산학습에서는 단순히 GPU 수를 늘린다고 선형 가속되지 않습니다.

  • 통신 overhead
  • 작은 per-GPU batch
  • 데이터 로딩 병목
  • 불균형한 연산
  • checkpoint 저장 비용

을 함께 측정해야 합니다.


15. 실전용 PyTorch 학습 루프

다음 코드는 현대적인 학습 구성 요소를 한 번에 연결한 예제입니다.

포함 기능:

  • AdamW parameter group
  • warmup + cosine decay
  • AMP와 FP16 GradScaler
  • gradient accumulation
  • gradient clipping
  • validation
  • best checkpoint
  • early stopping
  • optimizer update 기준 scheduler
from __future__ import annotations

import math
from pathlib import Path
from typing import Iterable

import torch
from torch import nn
from torch.utils.data import DataLoader


def build_adamw(
    model: nn.Module,
    learning_rate: float = 3e-4,
    weight_decay: float = 1e-2,
) -> torch.optim.Optimizer:
    """Bias와 1D 정규화 파라미터를 weight decay에서 제외합니다."""
    decay_params: list[nn.Parameter] = []
    no_decay_params: list[nn.Parameter] = []

    for parameter in model.parameters():
        if not parameter.requires_grad:
            continue

        if parameter.ndim >= 2:
            decay_params.append(parameter)
        else:
            no_decay_params.append(parameter)

    return torch.optim.AdamW(
        [
            {"params": decay_params, "weight_decay": weight_decay},
            {"params": no_decay_params, "weight_decay": 0.0},
        ],
        lr=learning_rate,
    )


def build_scheduler(
    optimizer: torch.optim.Optimizer,
    total_updates: int,
    warmup_ratio: float = 0.05,
    min_lr_ratio: float = 0.01,
) -> torch.optim.lr_scheduler.LRScheduler:
    """Optimizer update 단위의 linear warmup + cosine decay입니다."""
    if total_updates < 2:
        return torch.optim.lr_scheduler.LambdaLR(
            optimizer,
            lr_lambda=lambda _: 1.0,
        )

    warmup_steps = max(1, int(total_updates * warmup_ratio))
    warmup_steps = min(warmup_steps, total_updates - 1)
    cosine_steps = total_updates - warmup_steps

    warmup = torch.optim.lr_scheduler.LinearLR(
        optimizer,
        start_factor=0.1,
        end_factor=1.0,
        total_iters=warmup_steps,
    )

    cosine = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer,
        T_max=max(1, cosine_steps),
        eta_min=optimizer.param_groups[0]["lr"] * min_lr_ratio,
    )

    return torch.optim.lr_scheduler.SequentialLR(
        optimizer,
        schedulers=[warmup, cosine],
        milestones=[warmup_steps],
    )


def evaluate(
    model: nn.Module,
    data_loader: DataLoader,
    criterion: nn.Module,
    device: torch.device,
    use_amp: bool,
    amp_dtype: torch.dtype,
) -> float:
    model.eval()
    total_loss = 0.0
    total_samples = 0

    with torch.inference_mode():
        for inputs, targets in data_loader:
            inputs = inputs.to(device, non_blocking=True)
            targets = targets.to(device, non_blocking=True)

            with torch.autocast(
                device_type=device.type,
                dtype=amp_dtype,
                enabled=use_amp,
            ):
                outputs = model(inputs)
                loss = criterion(outputs, targets)

            batch_size = inputs.shape[0]
            total_loss += loss.item() * batch_size
            total_samples += batch_size

    if total_samples == 0:
        raise ValueError("Validation DataLoader is empty.")

    return total_loss / total_samples


def train_model(
    model: nn.Module,
    train_loader: DataLoader,
    val_loader: DataLoader,
    criterion: nn.Module,
    *,
    epochs: int = 30,
    learning_rate: float = 3e-4,
    weight_decay: float = 1e-2,
    accumulation_steps: int = 1,
    max_grad_norm: float = 1.0,
    patience: int = 5,
    checkpoint_path: str = "best_model.pt",
) -> dict[str, list[float]]:
    if accumulation_steps < 1:
        raise ValueError("accumulation_steps must be at least 1.")

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)

    optimizer = build_adamw(
        model,
        learning_rate=learning_rate,
        weight_decay=weight_decay,
    )

    updates_per_epoch = math.ceil(len(train_loader) / accumulation_steps)
    total_updates = max(1, updates_per_epoch * epochs)
    scheduler = build_scheduler(optimizer, total_updates)

    use_amp = device.type == "cuda"
    amp_dtype = (
        torch.bfloat16
        if use_amp and torch.cuda.is_bf16_supported()
        else torch.float16
    )

    scaler = torch.amp.GradScaler(
        "cuda",
        enabled=use_amp and amp_dtype == torch.float16,
    )

    history: dict[str, list[float]] = {
        "train_loss": [],
        "val_loss": [],
        "learning_rate": [],
    }

    best_val_loss = float("inf")
    bad_epochs = 0
    checkpoint = Path(checkpoint_path)

    for epoch in range(epochs):
        model.train()
        optimizer.zero_grad(set_to_none=True)

        running_loss = 0.0
        seen_samples = 0

        for micro_step, (inputs, targets) in enumerate(train_loader):
            inputs = inputs.to(device, non_blocking=True)
            targets = targets.to(device, non_blocking=True)

            with torch.autocast(
                device_type=device.type,
                dtype=amp_dtype,
                enabled=use_amp,
            ):
                outputs = model(inputs)
                raw_loss = criterion(outputs, targets)
                loss = raw_loss / accumulation_steps

            scaler.scale(loss).backward()

            batch_size = inputs.shape[0]
            running_loss += raw_loss.detach().item() * batch_size
            seen_samples += batch_size

            should_update = (
                (micro_step + 1) % accumulation_steps == 0
                or (micro_step + 1) == len(train_loader)
            )

            if not should_update:
                continue

            if scaler.is_enabled():
                scaler.unscale_(optimizer)

            grad_norm = torch.nn.utils.clip_grad_norm_(
                model.parameters(),
                max_norm=max_grad_norm,
                error_if_nonfinite=True,
            )

            scaler.step(optimizer)
            scaler.update()
            optimizer.zero_grad(set_to_none=True)
            scheduler.step()

        if seen_samples == 0:
            raise ValueError("Training DataLoader is empty.")

        train_loss = running_loss / seen_samples
        val_loss = evaluate(
            model,
            val_loader,
            criterion,
            device,
            use_amp,
            amp_dtype,
        )

        current_lr = optimizer.param_groups[0]["lr"]
        history["train_loss"].append(train_loss)
        history["val_loss"].append(val_loss)
        history["learning_rate"].append(current_lr)

        print(
            f"Epoch {epoch + 1:03d}/{epochs:03d} | "
            f"train={train_loss:.6f} | "
            f"val={val_loss:.6f} | "
            f"lr={current_lr:.3e}"
        )

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            bad_epochs = 0

            torch.save(
                {
                    "epoch": epoch + 1,
                    "model_state_dict": model.state_dict(),
                    "optimizer_state_dict": optimizer.state_dict(),
                    "scheduler_state_dict": scheduler.state_dict(),
                    "scaler_state_dict": scaler.state_dict(),
                    "best_val_loss": best_val_loss,
                },
                checkpoint,
            )
        else:
            bad_epochs += 1

        if bad_epochs >= patience:
            print(
                f"Early stopping: validation loss did not improve "
                f"for {patience} epochs."
            )
            break

    return history

이 코드에서 주의할 점

  1. criterion의 reduction이 기본 평균이라는 전제입니다.
  2. 분산학습에서는 loss 집계와 sampler를 별도로 처리해야 합니다.
  3. classification accuracy, F1 등 실제 목표 지표도 함께 기록해야 합니다.
  4. 재시작 학습은 model뿐 아니라 optimizer, scheduler, scaler state까지 복원해야 합니다.
  5. torch.compile은 모델과 환경을 확인한 뒤 선택적으로 적용합니다.

16. 학습이 이상할 때 보는 진단표

증상가능한 원인우선 확인할 것
Loss가 처음부터 NaNLR 과대, 잘못된 log/divide, 데이터 NaN입력·label finite 검사, LR 10배 감소
몇 step 후 NaNgradient explosion, FP16 overflowgrad norm 기록, clipping, scaler 확인
Loss가 전혀 감소하지 않음gradient 없음, label 오류, LR 너무 작음.grad 확인, target 범위·dtype 확인
Loss가 크게 진동LR 과대, batch 작음, noisy labelLR 감소, batch/accumulation 증가
Train만 좋아지고 Val 악화과적합, data leakage 반대 방향 점검split, augmentation, weight decay, early stopping
Train Loss가 Val보다 높음Dropout·augmentation·regularization학습·평가 모드 차이 확인
GPU 사용률 낮음DataLoader 병목, 작은 연산profiler, num_workers, pin_memory
OOMactivation·optimizer state 과다AMP, accumulation, checkpointing, FSDP
학습 재시작 후 성능 변화scheduler/scaler/optimizer 미복원전체 state_dict 저장·복원
같은 seed인데 결과 다름비결정적 kernel, 데이터 순서, 분산 통신reproducibility 설정과 환경 기록
Gradient가 모두 0detach, no_grad, saturating activation계산 그래프와 activation 분포 확인
일부 layer만 학습 안 됨frozen parameter, optimizer 누락requires_grad, param group 목록 확인

Gradient norm을 기록하라

Loss만 보면 문제 원인을 알기 어렵습니다.

total_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=float("inf"),
)
print(total_norm.item())
  • norm이 계속 0에 가까움: gradient vanishing 또는 그래프 단절
  • norm이 갑자기 수천~수백만: explosion 또는 이상 배치
  • 특정 시점 spike: 데이터·label·sequence length 이상 가능

파라미터가 실제로 바뀌는지 확인하라

before = model.layer.weight.detach().clone()

loss.backward()
optimizer.step()

after = model.layer.weight.detach()
print((after - before).abs().mean())

업데이트가 0이라면 다음을 확인합니다.

  • optimizer에 해당 parameter가 들어 있는가
  • requires_grad=True인가
  • gradient가 None인가
  • scaler가 계속 step을 skip하는가
  • LR이 0으로 떨어졌는가

17. 모델 유형별 시작 레시피

아래 값은 정답이 아니라 실험 시작점입니다.

17-1. 작은 MLP·표형 데이터

Optimizer: AdamW
LR: 1e-3 전후
Weight Decay: 1e-5 ~ 1e-2 탐색
Scheduler: ReduceLROnPlateau 또는 cosine
Batch: 32 ~ 512
Early Stopping: 적극 사용

표형 데이터는 딥러닝보다 LightGBM·XGBoost가 더 강할 수도 있으므로 반드시 기준선과 비교해야 합니다.

17-2. CNN 이미지 분류

빠른 기준선: AdamW + cosine
강한 전통 기준선: SGD + Momentum + 긴 cosine schedule
AMP: 권장
Augmentation: 필수 비교
Weight Decay: augmentation 강도와 함께 조정

17-3. Transformer 파인튜닝

Optimizer: AdamW
LR: 1e-5 ~ 5e-4 범위 탐색
Warmup: 총 update의 3% ~ 10%부터 탐색
Scheduler: linear 또는 cosine decay
Gradient Clipping: 1.0 시작점
Precision: BF16 가능하면 우선 검토

Pretrained backbone과 새 head에 서로 다른 학습률을 적용할 수도 있습니다.

optimizer = torch.optim.AdamW(
    [
        {"params": model.backbone.parameters(), "lr": 2e-5},
        {"params": model.classifier.parameters(), "lr": 2e-4},
    ],
    weight_decay=0.01,
)

17-4. LLM LoRA·QLoRA 파인튜닝

Trainable Parameter: LoRA adapter 중심
Optimizer: AdamW 또는 8-bit AdamW
Precision: BF16/FP16
Gradient Accumulation: 자주 필요
Gradient Checkpointing: 긴 context에서 유용
Loss Masking: prompt token과 padding 처리 반드시 확인

LLM 파인튜닝에서는 “loss가 감소했다”보다 다음이 중요합니다.

  • padding token이 loss에 포함되지 않는가
  • instruction prompt까지 학습 대상으로 넣을 것인가
  • sequence packing이 label alignment를 깨지 않는가
  • effective batch와 token 수가 얼마나 되는가
  • evaluation prompt가 학습 데이터와 중복되지 않는가

17-5. 대규모 LLM 사전학습

기준선: AdamW + warmup + cosine
비교 후보: Muon + AdamW 혼합
Precision: BF16 또는 검증된 FP8 recipe
Memory: FSDP/ZeRO + activation checkpointing
Metric: loss/token, tokens/sec, MFU, grad norm, skipped step

대규모 학습에서는 epoch보다 token 수와 optimizer update가 더 중요한 단위입니다.

총 학습 토큰
= global batch의 token 수 × optimizer steps

18. 재현 가능한 학습을 위한 체크리스트

  • [ ] train/validation/test 분리가 고정돼 있는가
  • [ ] 데이터 전처리 버전이 기록돼 있는가
  • [ ] random seed가 기록돼 있는가
  • [ ] 모델 코드와 config가 함께 저장되는가
  • [ ] optimizer, scheduler, scaler state를 저장하는가
  • [ ] best 모델과 last 모델을 구분하는가
  • [ ] 학습률과 gradient norm을 로그로 남기는가
  • [ ] 데이터 개수뿐 아니라 실제 학습 token 수를 기록하는가
  • [ ] mixed precision dtype을 기록하는가
  • [ ] PyTorch, CUDA, cuDNN, GPU 모델을 기록하는가
  • [ ] 여러 seed에서 결과 변동을 확인했는가
  • [ ] 테스트 세트를 하이퍼파라미터 선택에 사용하지 않았는가

PyTorch 공식 문서도 완전한 재현성이 릴리스, 플랫폼, CPU/GPU 사이에서 보장되지 않을 수 있다고 안내합니다. 따라서 seed만 저장하는 것으로 충분하지 않으며, 환경과 데이터 순서까지 기록해야 합니다.


19. 자주 묻는 질문

Q1. 초보자는 Adam과 AdamW 중 무엇을 써야 하나요

범용 기준선으로는 AdamW를 먼저 권장할 수 있습니다. 특히 Transformer 계열에서는 AdamW와 warmup/decay 조합이 널리 사용됩니다. 다만 작은 MLP나 이미 검증된 Adam 레시피가 있다면 무조건 변경할 필요는 없습니다.

Q2. SGD는 이제 오래된 옵티마이저인가요

아닙니다. SGD + Momentum은 여전히 강력하며, 충분히 잘 튜닝하면 이미지 분류 등에서 좋은 일반화 성능을 냅니다. 문제는 AdamW보다 긴 학습과 더 세심한 학습률 튜닝이 필요할 수 있다는 점입니다.

Q3. Muon을 바로 사용해도 되나요

대규모 Transformer 사전학습을 실험한다면 비교할 가치가 있습니다. 하지만 2D hidden-layer 파라미터에만 적용하고 나머지는 AdamW로 처리해야 합니다. 소규모 파인튜닝에서는 AdamW 기준선을 먼저 확보하는 것이 안전합니다.

Q4. Batch Size가 클수록 좋은가요

아닙니다. 큰 batch는 처리량과 gradient 안정성을 높일 수 있지만 일반화, 메모리, 통신 비용이 달라집니다. batch를 키우면 학습률과 warmup도 함께 조정해야 합니다.

Q5. Gradient Accumulation은 큰 batch와 완전히 같나요

대략적인 gradient 평균은 비슷하게 만들 수 있지만 BatchNorm, Dropout, 데이터 증강, scheduler 진행 단위 때문에 완전히 같지 않을 수 있습니다.

Q6. Loss가 낮아지는데 정확도가 오르지 않는 이유는

Cross-Entropy는 정답 클래스 순위뿐 아니라 확률의 confidence까지 반영합니다. 예측 클래스가 같아도 confidence가 개선되면 loss는 줄어들 수 있습니다. 반대로 class imbalance나 threshold 문제 때문에 loss 개선이 목표 metric으로 이어지지 않을 수 있습니다.

Q7. Validation Loss가 Train Loss보다 낮아도 정상인가요

정상일 수 있습니다. 학습 중에는 Dropout, 데이터 증강, label smoothing, regularization이 적용되지만 검증에서는 꺼질 수 있기 때문입니다. 단, 두 손실 계산 방식이 동일한지 먼저 확인해야 합니다.

Q8. AMP를 쓰면 정확도가 떨어지나요

적절히 사용하면 많은 모델에서 FP32와 비슷한 품질을 유지할 수 있습니다. 그러나 수치적으로 민감한 연산, custom op, 잘못된 scaling에서는 문제가 생길 수 있으므로 FP32 기준선과 비교해야 합니다.

Q9. Gradient Clipping은 항상 켜야 하나요

항상 필요한 것은 아닙니다. Transformer, RNN, 강화학습처럼 gradient spike가 흔한 모델에서는 좋은 안전장치가 될 수 있습니다. 그러나 clipping이 빈번하게 발생한다면 학습률, 데이터, 초기화 문제를 함께 점검해야 합니다.

Q10. Epoch 수는 어떻게 정하나요

데이터셋 크기와 모델 종류에 따라 달라집니다. 고정된 epoch 숫자를 외우기보다 validation metric, early stopping, 총 optimizer update, 총 token 수를 기준으로 관리하는 것이 좋습니다.

Q11. Scheduler는 epoch마다 호출하나요 step마다 호출하나요

Scheduler 설계에 따라 다릅니다. StepLR처럼 epoch 단위로 쓰는 경우도 있고, warmup/cosine처럼 optimizer update마다 쓰는 경우도 있습니다. 문서를 확인하고, gradient accumulation 사용 시 실제 optimizer step 수와 일치시켜야 합니다.

Q12. Loss가 NaN일 때 가장 먼저 무엇을 보나요

다음 순서를 권장합니다.

1. 입력·label에 NaN/Inf가 있는지
2. 학습률을 10배 낮췄을 때 해결되는지
3. loss 함수 입력 범위와 dtype이 맞는지
4. gradient norm이 폭발하는지
5. FP16 scaler가 step을 계속 건너뛰는지
6. 특정 batch에서만 재현되는지

20. 최종 정리

모델 학습은 다음 한 줄로 요약할 수 있습니다.

예측하고 → 틀린 정도를 계산하고 → 책임을 역으로 나누고 → 파라미터를 조금 바꾼다

수식으로는 다음 흐름입니다.

ŷ = fθ(x)
L = ℓ(ŷ, y)
g = ∇θL
θ ← Optimizer(θ, g)

하지만 현대 딥러닝의 실제 성능은 이 기본식만으로 결정되지 않습니다.

  • 손실함수와 출력층의 조합
  • 학습률과 warmup·decay
  • AdamW의 weight decay
  • gradient clipping과 normalization
  • AMP와 BF16·FP8
  • gradient accumulation과 activation checkpointing
  • fused optimizer와 torch.compile
  • DDP·FSDP·ZeRO 같은 분산 전략
  • 데이터 품질과 평가 설계

가 함께 작동해야 합니다.

처음에는 다음 조합으로 기준선을 만드는 것이 실용적입니다.

AdamW
+ 적절한 learning rate
+ warmup 후 cosine decay
+ validation 기반 checkpoint
+ gradient norm 모니터링
+ BF16/FP16 AMP

그다음 병목이 확인되면 Muon, fused optimizer, activation checkpointing, FSDP 같은 기술을 단계적으로 추가해야 합니다.

가장 중요한 원칙은 이것입니다.

최신 옵티마이저를 쓰는 것보다, 동일한 데이터와 예산에서 재현 가능한 기준선을 만들고 한 요소씩 비교하는 것이 먼저입니다.


같이 보기


참고문헌과 공식 문서

핵심 논문

  1. Rumelhart, Hinton, Williams, Learning representations by back-propagating errors, Nature, 1986.
    https://www.nature.com/articles/323533a0
  2. Kingma, Ba, Adam: A Method for Stochastic Optimization, 2014.
    https://arxiv.org/abs/1412.6980
  3. Loshchilov, Hutter, Decoupled Weight Decay Regularization, 2017/2019.
    https://arxiv.org/abs/1711.05101
  4. Loshchilov, Hutter, SGDR: Stochastic Gradient Descent with Warm Restarts, 2016.
    https://arxiv.org/abs/1608.03983
  5. Chen et al., Symbolic Discovery of Optimization Algorithms — Lion, 2023.
    https://arxiv.org/abs/2302.06675
  6. Micikevicius et al., FP8 Formats for Deep Learning, 2022.
    https://arxiv.org/abs/2209.05433
  7. Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, 2019.
    https://arxiv.org/abs/1910.02054

PyTorch 공식 문서

  • Autograd
    https://docs.pytorch.org/docs/stable/autograd.html
  • AdamW
    https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html
  • Muon
    https://docs.pytorch.org/docs/stable/generated/torch.optim.Muon.html
  • Automatic Mixed Precision
    https://docs.pytorch.org/docs/stable/notes/amp_examples.html
  • Gradient Clipping
    https://docs.pytorch.org/docs/stable/generated/torch.nn.utils.clip_grad_norm_.html
  • Activation Checkpointing
    https://docs.pytorch.org/docs/stable/checkpoint.html
  • DistributedDataParallel
    https://docs.pytorch.org/docs/stable/generated/torch.nn.parallel.DistributedDataParallel.html
  • torch.compile 튜토리얼
    https://docs.pytorch.org/tutorials/intermediate/torch_compile_tutorial.html
  • Reproducibility
    https://docs.pytorch.org/docs/stable/notes/randomness.html

Muon과 최신 저정밀도 학습

  • Keller Jordan, Muon: An optimizer for hidden layers in neural networks
    https://kellerjordan.github.io/posts/muon/
  • Liu et al., Muon is Scalable for LLM Training
    https://arxiv.org/abs/2502.16982
  • Towards Better Muon via One Additional Normalization Step — Muon+
    https://arxiv.org/abs/2602.21545
  • Anytime Training with Schedule-Free Spectral Optimization — SF-NorMuon
    https://arxiv.org/abs/2605.23061
  • Can Muon Fine-tune Adam-Pretrained Models?
    https://arxiv.org/abs/2605.10468
  • DMuon: Efficient Distributed Muon Training with Near-Adam Overhead
    https://arxiv.org/abs/2606.27153
  • Moonshot AI, Kimi K2 공식 저장소
    https://github.com/moonshotai/kimi-k2
  • NVIDIA Transformer Engine, FP8 Primer
    https://docs.nvidia.com/deeplearning/transformer-engine/user-guide/examples/fp8_primer.html

One Comment

댓글 남기기