Лосс уходит в NaN посреди обучения. Каковы причины и какое исправление для каждой?
Прогон нормально обучается примерно две эпохи, затем лосс печатает nan и уже не восстанавливается. Те же данные прекрасно обучаются в fp32 с меньшим learning rate.
Ограничения: mixed precision оставьте включённым, архитектуру сохраните и сделайте сбой обнаружимым на том шаге, где он происходит.
scaler = torch.amp.GradScaler()
for x, y in loader:
opt.zero_grad()
with torch.autocast("cuda", dtype=torch.float16):
logits = model(x)
loss = -(y * torch.log(torch.softmax(logits, dim=-1))).sum(-1).mean()
scaler.scale(loss).backward()
scaler.step(opt)
scaler.update()
Найдите все причины, по которым лосс здесь может уйти в NaN, и устраните их.
Взрыв градиентов — обрезайте по глобальной норме. Переполнение fp16 — loss scaler или bf16. Самописный log(0) — слитый лосс по логитам или эпсилон. Слишком большая скорость — снизьте или добавьте warmup. NaN уже в батче — проверяйте входы. Ищите первый плохой шаг по нормам градиентов.
- ✗Винят данные и никогда не смотрят нормы градиентов или реализацию лосса
- ✗Пишут log(softmax(x)) руками вместо слитого лосса по логитам
- ✗Считают, что GradScaler делает переполнение fp16 невозможным
- →Почему bf16 терпит переполнение лучше, чем fp16, при той же разрядности?
- →Как поймать точный шаг и тензор, где появился первый NaN?
Здесь три отдельные причины, и GradScaler закрывает лишь одну из них.
1. Самописный log(softmax(x)). При уверенном предсказании softmax в fp16 отдаёт ровно 0, а log(0) — это -inf; дальше -inf * 0 даёт nan. Слитый cross_entropy считает то же самое через log_softmax в стабильной форме и никогда не берёт логарифм нуля.
2. Отсутствие обрезания градиентов. Перед scaler.step градиенты нужно расшкалировать (scaler.unscale_) и обрезать по глобальной норме — иначе редкий выброс улетает в inf.
3. Нет проверки на месте. Сбой надо ловить на том шаге, где он случился, а не эпоху спустя.
scaler = torch.amp.GradScaler()
for step, (x, y) in enumerate(loader):
assert torch.isfinite(x).all(), f"шаг {step}: во входах NaN/inf"
opt.zero_grad()
with torch.autocast("cuda", dtype=torch.float16):
logits = model(x)
loss = torch.nn.functional.cross_entropy(logits, y) # слитый, стабильный
assert torch.isfinite(loss), f"шаг {step}: лосс {loss.item()}"
scaler.scale(loss).backward()
scaler.unscale_(opt) # расшкалировать ДО обрезания
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(opt)
scaler.update()
Если nan остаётся и после этого — переключите dtype на torch.bfloat16: его диапазон экспоненты совпадает с fp32, поэтому переполнение практически исчезает, а loss scaler становится не нужен.