Перейти к содержимому

CUDA out-of-memory при обучении/инференсе

CUDA out of memory — самый обидный способ потерять многочасовой ран: всё шло хорошо, а на длинном батче или длинной последовательности аллокатор уперся в потолок VRAM, и процесс умер. Если это случилось в 3 часа ночи, вы узнаете об этом в 9 утра — минус полдня. Push ловит момент падения и сразу подсказывает что уменьшить: batch size, seq length или включить gradient checkpointing.

Три места, где это можно поймать: в самом Python-коде, по exit-логам процесса и по свободной памяти карты заранее.

Вариант 1: ловим OOM в коде обучения/инференса

Заголовок раздела «Вариант 1: ловим OOM в коде обучения/инференса»

PyTorch кидает torch.cuda.OutOfMemoryError (наследник RuntimeError). Оборачиваем шаг и вытаскиваем контекст — размеры тензора, который не влез:

import os, torch, requests
NOTIFLY_URL = os.environ["NOTIFLY_URL"]
NOTIFLY_TOKEN = os.environ["NOTIFLY_TOKEN"]
def notify(title, msg, prio):
requests.post(f"{NOTIFLY_URL}/message",
params={"token": NOTIFLY_TOKEN},
json={"title": title, "message": msg, "priority": prio},
timeout=5)
def run_step(model, batch, run_name):
try:
return model(**batch)
except torch.cuda.OutOfMemoryError as e:
# что было в момент падения
bs = batch["input_ids"].shape[0]
seq_len = batch["input_ids"].shape[1]
free, total = torch.cuda.mem_get_info()
notify("💥 CUDA OOM",
f"{run_name}\nbatch_size={bs}, seq_len={seq_len}\n"
f"VRAM: свободно {free/1e9:.1f}/{total/1e9:.1f} ГБ\n"
f"Уменьшить batch/seq или включить gradient checkpointing.",
prio=9)
torch.cuda.empty_cache()
raise

priority=9 — падение стоит вам GPU-часов и, возможно, испорченного рана, поэтому громко. В batch_size/seq_len — самое ценное: обычно достаточно поделить одно из них надвое.

Иногда OOM убивает процесс целиком (OOM-killer ядра, CUDA_ERROR_OUT_OF_MEMORY из C++-слоя) — и до Python-обработчика дело не доходит. Тогда сторожим сам процесс снаружи: запускаем тренировку через обёртку, которая ловит смерть и грепает лог:

#!/usr/bin/env bash
# train_guard.sh — запускает тренировку и алёртит при OOM
LOG=/var/log/train.log
python train.py "$@" > "$LOG" 2>&1
code=$?
if [[ $code -ne 0 ]]; then
# выцепляем последнюю OOM-строку и контекст
oom=$(grep -iE "out of memory|CUDA_ERROR_OUT_OF_MEMORY|Killed process" "$LOG" | tail -3)
title="💥 Тренировка упала (exit $code)"
[[ -n "$oom" ]] && title="💥 CUDA OOM (exit $code)"
msg=$(printf 'Процесс: train.py\nExit: %s\n\n%s' "$code" "${oom:-см. лог}")
curl -fsS "$NOTIFLY_URL/message?token=$NOTIFLY_TOKEN" \
-H "Content-Type: application/json" \
-d "$(jq -n --arg t "$title" --arg m "$msg" '{title:$t, message:$m, priority:9}')"
fi
exit $code

dmesg/journalctl покажут Out of memory: Killed process ... для OOM-killer ядра — тот же греп ловит и его.

Вариант 3: превентивный watchdog по свободной VRAM

Заголовок раздела «Вариант 3: превентивный watchdog по свободной VRAM»

Лучше не падать вовсе. Отдельный таймер-скрипт следит за свободной памятью и предупреждает до OOM, пока ещё можно спасти ран:

import subprocess
def free_vram_mib():
out = subprocess.check_output(
["nvidia-smi", "--query-gpu=memory.free", "--format=csv,noheader,nounits"],
text=True)
return min(int(x) for x in out.split())
if free_vram_mib() < 500: # осталось меньше 500 МиБ — вот-вот OOM
notify("🟠 VRAM почти кончилась",
f"Свободно {free_vram_mib()} МиБ. Скоро OOM — снизьте нагрузку.",
prio=7)

Связка с heartbeat: если после OOM процесс молча умер и перестал слать пинги, Notifly пришлёт второй, независимый алёрт по таймауту.

  • имя процесса/рана и exit-код;
  • batch_size и seq_len в момент падения — главное для быстрого фикса;
  • свободная/общая VRAM на карте;
  • подсказка: уменьшить батч/seq, gradient checkpointing, offload, меньший dtype.