AdamW와 역전파로 설계하는 딥러닝 학습 최적화
Gradient-Based Optimization과 Backpropagation, AdamW의 갱신 구조를 바탕으로 학습률·정규화·재현성 관리 방식을 정리한다.
2026-08-14 · 최초 발행 2024-04-29
기울기를 계산하고 갱신하는 학습 루프
딥러닝 학습에서 손실을 낮추는 과정은 기울기를 계산하고, 그 값으로 파라미터를 갱신하는 반복으로 이루어진다. Gradient-Based Optimization, Backpropagation, AdamW는 이 반복의 서로 다른 층을 담당한다. 하나는 갱신의 원리이고, 하나는 미분 계산 방식이며, 다른 하나는 실제 갱신 규칙이다.
Gradient-Based Optimization은 손실 함수의 기울기 정보를 바탕으로 파라미터를 반복 업데이트하는 최적화 방법론이다. SGD, 모멘텀, Adam, AdamW가 대표적인 구성에 속한다. 이때 학습률(learning rate), 모멘텀 또는 적응적 모멘트의 β 파라미터, 가중치 감쇠(weight decay), 스케줄러(scheduler)가 갱신 특성을 좌우한다.
Backpropagation은 연쇄법칙(chain rule)으로 출력 오차를 계산 그래프의 역방향으로 전달해 각 파라미터의 미분값(gradient)을 구하는 알고리즘이다. PyTorch, JAX, TensorFlow 같은 자동미분 프레임워크에서는 계산 그래프 생성, 순전파, 손실 계산, 역전파, 그래디언트 축적, 옵티마이저 업데이트 순으로 이 절차가 수행된다.
AdamW는 Adam의 적응적 모멘트 추정(일차·이차 모멘트)에 weight decay를 디커플링(decoupled)한 알고리즘이다. Adam의 L2 정규화와 달리 파라미터 업데이트와 별개로 weight decay를 적용하며, 대규모 트랜스포머와 비전 모델 파인튜닝에서 사실상 표준으로 사용된다.
계산 그래프의 불안정성을 다루는 방법
자동미분은 계산 그래프를 구성하고 연쇄법칙에 따라 도함수를 계산한다. 이 구조에는 정방향과 역방향 처리 사이의 메모리·연산 트레이드오프가 있다.
학습 중에는 기울기 소실이나 폭주, 비정상 NaN이 발생할 수 있다. 초기화(He/Xavier), 정규화(LayerNorm/BatchNorm), 기울기 클리핑, 혼합정밀도 환경에서의 eps 보강이 대응 수단이 된다. Adam과 AdamW는 배치와 스케일 변화에 강인하고 워밍업과 함께 초기 수렴을 빠르게 가져갈 수 있다. AdamW는 weight decay를 분리하므로 학습률과 감쇠 사이의 상호작용을 줄인다.
학습률은 코사인, 지수, 스텝 스케줄로 조정할 수 있다. 대규모 모델에서는 수백~수천 스텝의 워밍업을 두어 초기 발산을 막는다. 배치 크기를 키울 때는 선형 스케일링 규칙(LR ∝ batch_size)을 적용하고, 분산 학습에서는 글로벌 배치 기준을 관리한다.
가중치 감쇠는 weight에만 적용하고 바이어스와 정규화 파라미터는 제외하는 파라미터 그룹 구성이 권장된다. Dropout, Stochastic Depth 같은 구조적 정규화도 과적합 억제를 위해 함께 사용할 수 있다.
배치에서 다음 갱신까지
모델 유형에 따라 달라지는 설정
대규모 언어 모델(LLM)의 프리트레이닝과 파인튜닝에서는 AdamW, Cosine decay, Linear warmup 조합을 적용할 수 있다. LayerNorm과 bias는 no-decay 파라미터 그룹으로 구성하고, 혼합정밀도와 ZeRO 또는 텐서 병렬로 메모리를 분할해 운영한다.
컴퓨터 비전 파인튜닝에서는 사전학습 가중치를 불러온 뒤 AdamW와 낮은 LR(1e-5~3e-5), weight_decay 0.01을 적용하며, WD를 헤드와 백본에 차등 적용한다. 스케줄은 5% 워밍업과 Cosine 조합을 사용할 수 있고, 라벨 스무딩과 강화학습 혼합은 케이스별로 검토한다.
표형 데이터 딥러닝이나 추천 문제에서는 초기 수렴 속도가 중요한 상황에 AdamW를 우선 적용할 수 있다. 피처 스케일 차이가 클 때는 적응적 LR의 강인성이 도움이 된다. 일반화는 조기 종료(Early stopping)와 k-폴드 검증으로 확인한다.
운영 환경에서는 하이퍼파라미터, SHA, 데이터셋 버전을 실험 메타데이터로 남기고 체크포인트를 주기적으로 저장한다. 성능 회귀(regression) 검출 파이프라인도 함께 구성한다. 분산 학습에서는 글로벌 클립 한도, 동기화 BatchNorm, 혼합 정밀 loss scaling을 자동 관리한다.
학습 안정성과 운영 효율에 미치는 영향
AdamW 도입 시 초기 수렴은 데이터와 모델에 따라 1.2~3.0배 개선 가능하며, 러닝레이트 스케줄과 병행하면 에폭 단축 효과가 있다. 기울기 폭주와 NaN 비율, 재시작 횟수도 줄일 수 있다. FP16 환경에서는 eps를 1e-5로 조정해 수치 안정성을 개선한다.
동일 자원 대비 검증 성능이 +0.2~1.0pp 상승한 사례가 있으며, no-decay 그룹 분리는 과적합 억제에 활용된다. AMP, 클리핑, 스케줄을 표준화하면 실험 실패율을 낮추고 재현성을 높일 수 있다. 분산 환경으로 확장할 때는 튜닝 결과의 재사용성도 커진다.
PyTorch에서 AdamW 학습 루프 구성하기
전제조건: Python ≥ 3.9, PyTorch ≥ 2.1, CUDA 선택적. CPU에서도 실행 가능.
# pip install torch torchvision
import torch
import torch.nn as nn
from torch.utils.data import DataLoader, TensorDataset
# 재현성
torch.manual_seed(42)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 더미 데이터 (예: 28x28 이미지 → 10 클래스)
N, D_in, H, D_out = 1024, 784, 512, 10
X = torch.randn(N, D_in)
y = torch.randint(0, D_out, (N,))
ds = TensorDataset(X, y)
dl = DataLoader(ds, batch_size=64, shuffle=True, drop_last=True)
# 간단 모델
model = nn.Sequential(
nn.Linear(D_in, H),
nn.ReLU(),
nn.Linear(H, D_out)
).to(device)
# 파라미터 그룹: weight_decay 비적용 대상 분리 (bias/Norm 등)
decay, no_decay = [], []
for name, p in model.named_parameters():
if not p.requires_grad:
continue
if p.dim() == 1 or name.endswith(".bias"):
no_decay.append(p)
else:
decay.append(p)
optimizer = torch.optim.AdamW(
[{"params": decay, "weight_decay": 0.01},
{"params": no_decay, "weight_decay": 0.0}],
lr=1e-3, betas=(0.9, 0.999), eps=1e-8 # FP16 시 eps=1e-5 고려
)
# 스케줄러(예: Cosine)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)
scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())
criterion = nn.CrossEntropyLoss()
model.train()
for epoch in range(5):
for xb, yb in dl:
xb, yb = xb.to(device), yb.to(device)
optimizer.zero_grad(set_to_none=True)
with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):
logits = model(xb)
loss = criterion(logits, yb)
scaler.scale(loss).backward()
# 기울기 폭주 방지
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
scheduler.step()
print("최종 손실 예시:", float(loss))
실무에서는 AdamW를 lr=1e-3으로 두고, 프리트레인 또는 대형 모델 파인튜닝에는 1e-55e-5를 고려한다. betas=(0.9,0.999), weight_decay=0.01, eps=1e-8을 사용하며 AMP에서는 eps=1e-5를 고려한다. 워밍업은 전체 스텝의 25%를 기준으로 둘 수 있고, Cosine 또는 Linear decay를 스케줄로 채택한다.
LayerNorm, Embedding, bias는 no-decay로 분리하고 Conv 및 Linear weight에만 decay를 적용한다. SGD와 모멘텀은 충분히 큰 에폭과 강한 데이터 증강 조건에서 더 나은 일반화를 보일 수 있으므로, 모델과 데이터 특성에 따른 트레이드오프를 검토한다.
옵티마이저별 갱신 특성
| 항목 | SGD | Momentum(SGD+μ) | Adam | AdamW |
|---|---|---|---|---|
| 수렴 속도 | 중간, LR 튜닝 민감 | 빠름, 초기 진동 완화 | 빠름, 배치/스케일 변화 강인 | 빠름, 튜닝 안정성 우수 |
| 일반화 | 우수, 장에폭 유리 | 우수, 과적합 억제 | 경우에 따라 과적합 경향 | Adam 대비 일반화 개선 |
| 안정성 | 폭주/소실 민감 | 진동 감소 | 수치 안정성 양호 | 디커플링으로 안정성 향상 |
| 확장성 | 대규모에서도 단순 | 유사 | 대규모에 적합 | 대규모·분산 표준 |
| 운영 편의 | 하이퍼 튜닝 부담 | 중간 | 기본값으로 용이 | 기본값·no-decay 그룹으로 용이 |
Backpropagation은 계산 그래프에서 기울기를 산출하는 메커니즘이고, Gradient-Based Optimization은 그 기울기로 손실을 줄이는 절차다. AdamW는 weight decay를 분리해 튜닝 안정성과 일반화를 함께 다룬다. no-decay 파라미터 그룹, 워밍업과 Cosine 스케줄, AMP와 클리핑, 시드 및 환경 고정은 이 학습 루프를 운영 가능한 형태로 만드는 구성이다.