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이 누적 그래디언트에 미치는 영향도 확인해요.