본문 바로가기
C.W.K.
Stream
Lesson 08 of 08 · published

합치기: 운영에 쓸 수 있는 학습 반복문

~14 min · production, gpu, logging, tqdm

Level 0텐서 탐구자
0 XP0/62 lessons0/13 achievements
0/120 XP to next level120 XP to go0% complete

한곳에 모두 모은, 복사해서 응용할 수 있는 반복문

이 레슨에서는 트랙에서 배운 모든 방식을 실행 가능한 학습 반복문 하나로 조립해. 한 번 꼼꼼히 읽어 둬. 앞으로도 이 구조를 바탕으로 필요한 부분을 바꾸게 될 거야.

반복문에는 다음이 들어 있어:

  • 장치 선택(CUDA → MPS → CPU 대안)
  • 가중치 감쇠를 분리한 AdamW. 편향과 정규화 매개변수에는 감쇠를 적용하지 않아.
  • 워밍업을 포함한 코사인 학습률 일정
  • 혼합 정밀도. Ampere 이후 GPU와 Apple Silicon에서는 bfloat16을 사용해.
  • 기울기 제한
  • 누적 손실을 보여 주는 tqdm 진행 표시줄
  • 에포크마다 실행하는 검증
  • 가장 좋은 체크포인트 저장
  • NaN 방어

다음 기능은 각각 별도의 트랙이나 레슨에서 다루므로 일부러 넣지 않았어:

  • 분산 학습(DDP / FSDP): 트랙 7
  • torch.compile(): 트랙 7
  • 실험 추적(W&B / MLflow): 트랙 8

Code

처음부터 끝까지 이어지는 전체 반복문·python
import math
import torch
import torch.nn as nn
from torch.amp import autocast
from tqdm import tqdm

# 1. Device
device = (
    "cuda" if torch.cuda.is_available()
    else "mps" if torch.backends.mps.is_available()
    else "cpu"
)

# 2. Model
model = MyModel().to(device)

# 3. Loss
criterion = nn.CrossEntropyLoss()

# 4. Optimizer with weight-decay split
decay, no_decay = [], []
for n, p in model.named_parameters():
    if not p.requires_grad: continue
    if p.dim() < 2 or any(k in n for k in ('bias', 'norm')):
        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-4,
)

# 5. Scheduler — warmup + cosine
total_steps = len(train_loader) * num_epochs
warmup_steps = total_steps // 20    # 5% warmup
def lr_lambda(step):
    if step < warmup_steps:
        return step / warmup_steps
    progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
    return 0.5 * (1.0 + math.cos(math.pi * progress))
scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
이어서: 학습, 검증, 체크포인트·python
best_val = float('inf')
global_step = 0

for epoch in range(num_epochs):
    # ---- TRAIN ----
    model.train()
    pbar = tqdm(train_loader, desc=f"epoch {epoch:02d}")
    running = 0.0

    for batch_x, batch_y in pbar:
        batch_x = batch_x.to(device, non_blocking=True)
        batch_y = batch_y.to(device, non_blocking=True)

        optimizer.zero_grad(set_to_none=True)

        with autocast(device_type=device, dtype=torch.bfloat16) if device != 'cpu' else nullcontext():
            output = model(batch_x)
            loss = criterion(output, batch_y)

        if not torch.isfinite(loss):
            print(f"step {global_step}: non-finite loss; skipping")
            continue

        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()
        scheduler.step()
        global_step += 1

        running = 0.95 * running + 0.05 * loss.item() if global_step > 1 else loss.item()
        pbar.set_postfix(loss=f"{running:.4f}", lr=f"{scheduler.get_last_lr()[0]:.2e}")

    # ---- VALIDATE ----
    model.eval()
    val_loss = 0.0
    val_n = 0
    with torch.inference_mode():
        for vx, vy in val_loader:
            vx, vy = vx.to(device), vy.to(device)
            val_loss += criterion(model(vx), vy).item() * vx.size(0)
            val_n += vx.size(0)
    val_loss /= val_n
    print(f"epoch {epoch:02d}  val_loss={val_loss:.4f}")

    # ---- CHECKPOINT ----
    if val_loss < best_val:
        best_val = val_loss
        torch.save({
            'epoch': epoch,
            'val_loss': val_loss,
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'scheduler_state_dict': scheduler.state_dict(),
        }, 'best.pt')
nullcontext 보완 코드: 자동 형변환이 CPU에서 아무 일도 하지 않는 연산·python
from contextlib import nullcontext

# autocast doesn't apply on CPU. Use nullcontext as a no-op stand-in
# so the same training code runs cleanly on CPU for debugging.

device = "cpu"  # for example
amp_ctx = autocast(device_type=device, dtype=torch.bfloat16) if device != 'cpu' else nullcontext()

with amp_ctx:
    out = model(x)
    loss = criterion(out, y)

External links

Exercise

이 반복문을 MNIST, CIFAR-10, FashionMNIST 중 하나에 맞게 바꾸고 5에포크 실행해 봐. 다음을 확인해. (1) 에포크가 지날수록 검증 손실이 감소하는가, (2) 가장 좋은 체크포인트를 새 모델에 문제없이 불러올 수 있는가, (3) bfloat16 자동 형변환이 현재 하드웨어에서 실제로 활성화되는가. 자동 형변환을 끈 경우와 실행 시간을 비교하면 차이를 확인할 수 있어.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.