torch.nn.utils.clip_grad_norm_ — 노름 기준 클리핑

torch.nn.utils.clip_grad_norm_ — 노름 기준 클리핑

PyTorch의 clip_grad_norm_은 텐서들의 그래디언트 전체 노름(norm) 을 계산해, 설정한 max_norm을 넘으면 전체를 비율로 축소해요. '그래디언트 방향은 유지하되 크기만 제한'하는 방식이라 RNN 학습에 자주 쓰여요.

동작 원리

  • 모든 파라미터의 그래디언트를 모아 L2 노름 total_norm을 계산해요.
  • total_norm > max_norm이면 스케일 max_norm / total_norm만큼 각 그래디언트에 곱해요.
  • 노름이 한도를 넘지 않으면 그래디언트를 그대로 두어요.

예제 코드

import torch.nn as nn
outputs = model(data)
loss = criterion(outputs, target)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

사용 팁

  • max_norm는 보통 0.5~1.0 근처에서 시작해 실험으로 조정해요.
  • global 노름 기준이므로 모델 전체에 걸쳐 일관된 크기 제한이 적용돼요.

더 알아보기