PyTorch — 옵티마이저 스텝과 그래디언트 누적

PyTorch — 옵티마이저 스텝과 그래디언트 누적

PyTorch 공식 레시피는 그래디언트를 언제 누적하고 언제 0으로 초기화(zero_grad) 해야 하는지를 코드로 설명해요. backward()가 그래디언트를 매 스텝 누적하는 동작 원리를 이해하면 누적 기법을 직접 구현할 수 있어요.

핵심 이해

  • PyTorch는 loss.backward()를 호출할 때마다 텐서의 .grad더해(누적) 저요.
  • 그래서 optimizer.zero_grad()를 호출하지 않으면 다음 스텝의 그래디언트가 이전 것에 합쳐져요.
  • 이 특성을 이용하면, 아래처럼 누적만 하고 주기적으로만 zero_grad + step을 하면 돼요.

예제 코드

opt.zero_grad()
for i, (x, y) in enumerate(dataloader):
    loss = model(x, y)
    loss.backward()                      # 그래디언트 누적
    if (i + 1) % accum_steps == 0:       # 누적 주기 도달 시
        opt.step()                        # 옵티마이저 갱신
        opt.zero_grad()                   # 그래디언트 초기화

실전 팁

  • 주기와 데이터셋의 관계를 신경 써서, epoch 끝에서 그래디언트가 남지 않도록 처리해요.
  • 혼합 정밀도(fp16)와 함께 쓸 때 scale이 누적 그래디언트에 미치는 영향도 확인해요.

더 알아보기