
학습을 몇 시간 돌리다가 loss가 갑자기 NaN으로 튀는 상황은 PyTorch로 모델을 만들어본 사람이라면 한 번쯤 겪는 문제입니다. 로그를 아무리 들여다봐도 “어느 레이어에서, 어느 연산에서” NaN이 시작됐는지 알기 어렵기 때문에 감으로 학습률만 낮추다 끝나는 경우가 많습니다. 이 글에서는 loss가 NaN으로 튀는 5가지 대표 원인을 체크리스트로 정리하고, register_hook과 register_forward_hook으로 NaN이 발생한 정확한 지점을 찾아내는 실전 디버깅 코드를 다룹니다.
원인 1~2: 학습률 폭주와 손실 함수의 수치 불안정성
가장 흔한 원인은 학습률이 지나치게 높아서 파라미터 업데이트 폭이 커지고, 그 결과 activation 값이 float 표현 범위를 넘어서면서 inf가 생기는 경우입니다. inf끼리 연산되면 inf - inf나 0 * inf 같은 형태가 나타나면서 NaN으로 전이됩니다. 특히 Adam 계열 옵티마이저는 초반 몇 스텝에서 이동평균이 불안정해 학습률이 조금만 높아도 폭주하기 쉽습니다.
두 번째는 손실 함수 자체의 수치 불안정성입니다. CrossEntropyLoss를 직접 구현한다고 log(softmax(x))를 수동으로 계산하면, softmax 출력이 0에 가까워질 때 log(0)이 -inf가 됩니다. PyTorch가 제공하는 nn.CrossEntropyLoss나 nn.functional.log_softmax는 내부적으로 log-sum-exp 트릭을 써서 이런 상황을 방지하므로, 직접 구현보다는 공식 함수를 쓰는 편이 안전합니다. 관련 수치적 배경은 PyTorch 공식 문서의 torch.nn 모듈 설명에서 각 손실 함수의 구현 방식을 확인할 수 있습니다.
체크리스트:
– 학습률을 10분의 1로 낮췄을 때 NaN 발생 시점이 늦춰지는가
– 커스텀 손실 함수에 log, sqrt, division 연산이 들어있는가
– warmup 스케줄러 없이 큰 학습률로 바로 시작하지는 않았는가
원인 3: Mixed Precision(AMP) 언더플로/오버플로
torch.cuda.amp로 float16 혼합정밀도 학습을 쓰면 속도는 빨라지지만, float16의 표현 범위(약 ±65504)를 넘는 gradient가 그대로 inf가 되는 문제가 생깁니다. 반대로 아주 작은 gradient 값은 float16의 최소 표현 범위 아래로 내려가면서 0으로 언더플로되기도 합니다. 이를 막기 위해 GradScaler가 loss를 일정 배율(scale)로 키워서 backward를 수행한 뒤 gradient를 다시 나누는 방식을 씁니다.
scaler = torch.cuda.amp.GradScaler()
for data, target in loader:
optimizer.zero_grad()
with torch.cuda.amp.autocast():
output = model(data)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
GradScaler는 inf/NaN gradient가 감지되면 해당 스텝의 optimizer.step()을 건너뛰고 scale 값을 자동으로 낮춥니다. 그런데도 loss가 NaN으로 찍힌다면, AMP 문제가 아니라 forward 연산 자체(예: exp, log, 큰 값의 나눗셈)에서 오버플로가 났을 가능성이 높습니다. AMP와 float16 범위에 대한 배경은 NVIDIA의 Mixed Precision Training 문서에 상세히 정리되어 있습니다.

원인 4~5: 입력 데이터의 결측치와 그래디언트 폭주
네 번째 원인은 입력 데이터 자체에 NaN이나 Inf가 섞여 들어오는 경우입니다. 정규화(normalization) 과정에서 표준편차가 0인 컬럼을 나누면 0/0이 NaN이 되고, 이 값이 모델에 그대로 들어가 첫 forward부터 NaN이 전파됩니다. 데이터로더 단계에서 다음처럼 검증하는 습관이 필요합니다.
for batch in loader:
x, y = batch
assert not torch.isnan(x).any(), "입력 데이터에 NaN 존재"
assert not torch.isinf(x).any(), "입력 데이터에 Inf 존재"
다섯 번째는 RNN, Transformer처럼 층이 깊거나 순환 구조를 가진 모델에서 흔한 그래디언트 폭주(exploding gradient)입니다. gradient의 노름(norm)이 매 스텝 기하급수적으로 커지다가 어느 순간 float 표현 범위를 넘습니다. torch.nn.utils.clip_grad_norm_으로 gradient norm 상한을 걸어두면 폭주를 상당 부분 막을 수 있습니다.
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
register_hook으로 NaN 발생 지점 정확히 찾기
원인을 좁히는 것과 “정확히 어느 레이어, 어느 연산”인지 찾는 것은 다른 문제입니다. Tensor.register_hook은 해당 텐서의 gradient가 계산될 때마다 콜백을 실행할 수 있게 해주는 autograd 훅으로, 파라미터별로 등록해두면 backward 도중 NaN이 처음 나타나는 파라미터를 바로 특정할 수 있습니다.
def make_nan_hook(name):
def hook(grad):
if torch.isnan(grad).any():
print(f"[NaN 감지] 파라미터 '{name}'의 gradient에서 NaN 발생")
if torch.isinf(grad).any():
print(f"[Inf 감지] 파라미터 '{name}'의 gradient에서 Inf 발생")
return grad
return hook
for name, param in model.named_parameters():
if param.requires_grad:
param.register_hook(make_nan_hook(name))
forward 쪽 activation을 검사하고 싶다면 nn.Module.register_forward_hook을 모든 서브모듈에 걸어두면 됩니다. 이렇게 하면 backward 이전, forward 단계에서부터 NaN이 시작된 레이어를 찾을 수 있어 손실 함수 문제인지 특정 레이어(예: LayerNorm, exp 연산 포함 레이어) 문제인지 구분됩니다.
def forward_nan_hook(module, input, output):
out = output[0] if isinstance(output, tuple) else output
if isinstance(out, torch.Tensor) and torch.isnan(out).any():
print(f"[NaN 감지] {module.__class__.__name__} 의 출력에서 NaN 발생")
for module in model.modules():
module.register_forward_hook(forward_nan_hook)
이 둘을 같이 걸어두고 학습을 재현하면 콘솔에 다음과 같은 형태로 출력이 찍히며, 어느 레이어가 먼저 오염됐는지 순서대로 확인할 수 있습니다.

[NaN 감지] TransformerEncoderLayer 의 출력에서 NaN 발생
[NaN 감지] 파라미터 'encoder.layer.3.self_attn.out_proj.weight'의 gradient에서 NaN 발생
더 정밀하게 “몇 번째 연산에서” NaN이 생겼는지까지 보고 싶다면 torch.autograd.set_detect_anomaly(True)를 학습 루프 진입 전에 호출합니다. 이 모드는 backward 도중 NaN을 만드는 연산을 만나면 그 연산의 스택 트레이스를 포함한 RuntimeError를 즉시 발생시킵니다. 다만 연산마다 추가 검사를 하기 때문에 학습 속도가 눈에 띄게 느려지므로, 디버깅이 끝나면 반드시 꺼야 합니다. 관련 API는 PyTorch 공식 autograd 문서에서 확인할 수 있습니다.
원인별 확인·해결 비교표
| 원인 | 대표 증상 | 확인 방법 | 해결책 |
|---|---|---|---|
| 학습률 과다 | 초반 몇 스텝 안에 loss 급등 후 NaN | 학습률 1/10로 낮춰 재현 | warmup + 학습률 하향 |
| 손실 함수 불안정 | 특정 배치에서만 간헐적으로 NaN | 커스텀 loss 수식 점검 | log_softmax 등 공식 함수 사용 |
| AMP 오버플로/언더플로 | fp16 학습에서만 재현 | fp32로 전환 후 재현 여부 확인 | GradScaler 적용, autocast 범위 재조정 |
| 입력 데이터 결측 | 학습 시작 직후 첫 배치부터 NaN | torch.isnan(x).any() 검사 |
데이터 전처리 단계에서 필터링 |
| 그래디언트 폭주 | 학습이 진행될수록 loss가 점점 커지다 NaN | gradient norm 로깅 | clip_grad_norm_ 적용 |
이 표는 순서대로 점검하면 원인 후보를 빠르게 좁힐 수 있도록 구성한 것으로, 실제로는 두세 가지가 동시에 겹쳐서 나타나는 경우도 많습니다. 예를 들어 학습률이 약간 높은 상태에서 AMP까지 함께 쓰면 오버플로가 훨씬 쉽게 발생합니다.
정리 및 다음 단계
loss가 갑자기 NaN으로 튀는 문제는 학습률, 손실 함수 수식, AMP 설정, 입력 데이터, 그래디언트 폭주라는 5가지 원인 중 하나 또는 조합에서 발생합니다. 감으로 하이퍼파라미터를 바꾸기 전에 register_hook으로 파라미터별 gradient를, register_forward_hook으로 레이어별 activation을 감시해 NaN이 처음 나타나는 지점을 특정하는 것이 훨씬 빠릅니다. 오늘 학습 코드에 위 두 훅을 걸어두고 재현 실행을 한 번 돌려보면, 다음 NaN 발생 시 원인을 바로 좁힐 수 있습니다. 더 정밀한 위치가 필요하면 torch.autograd.set_detect_anomaly(True)를 임시로 켜서 정확한 연산 스택 트레이스를 확인하는 것을 권장합니다.
자주 묻는 질문(FAQ)
Q1. GradScaler를 쓰는데도 loss가 NaN이 됩니다. 왜 그런가요?
GradScaler는 backward 과정에서 발생한 inf/NaN gradient를 감지해 해당 스텝의 파라미터 업데이트만 건너뛰는 역할을 합니다. forward 연산 자체(exp, log, 나눗셈 등)에서 이미 NaN이 만들어졌다면 GradScaler로는 막을 수 없으므로, register_forward_hook으로 forward 단계의 activation을 먼저 점검해야 합니다.
Q2. set_detect_anomaly(True)를 항상 켜두면 안 되나요?
연산마다 추가 검사가 들어가기 때문에 학습 속도가 크게 느려집니다. NaN이 재현되는 상황에서만 짧게 켜서 원인 연산을 특정한 뒤 바로 꺼서 운영하는 것이 일반적인 방식입니다.
Q3. register_hook과 register_forward_hook을 동시에 걸어도 되나요?
네, 서로 다른 시점(backward gradient vs forward activation)을 감시하므로 함께 사용해도 충돌하지 않습니다. 다만 디버깅이 끝나면 handle = module.register_forward_hook(...)로 반환된 핸들을 handle.remove()로 해제해 불필요한 오버헤드를 없애는 것이 좋습니다.