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 노름 기준이므로 모델 전체에 걸쳐 일관된 크기 제한이 적용돼요.