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() raisepriority=9 — падение стоит вам GPU-часов и, возможно, испорченного рана,
поэтому громко. В batch_size/seq_len — самое ценное: обычно достаточно
поделить одно из них надвое.
Вариант 2: сторож по exit-коду и логам
Заголовок раздела «Вариант 2: сторож по exit-коду и логам»Иногда OOM убивает процесс целиком (OOM-killer ядра, CUDA_ERROR_OUT_OF_MEMORY
из C++-слоя) — и до Python-обработчика дело не доходит. Тогда сторожим сам
процесс снаружи: запускаем тренировку через обёртку, которая ловит смерть и
грепает лог:
#!/usr/bin/env bash# train_guard.sh — запускает тренировку и алёртит при OOMLOG=/var/log/train.logpython train.py "$@" > "$LOG" 2>&1code=$?
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}')"fiexit $codedmesg/journalctl покажут Out of memory: Killed process ... для OOM-killer
ядра — тот же греп ловит и его.
Вариант 3: превентивный watchdog по свободной VRAM
Заголовок раздела «Вариант 3: превентивный watchdog по свободной VRAM»Лучше не падать вовсе. Отдельный таймер-скрипт следит за свободной памятью и предупреждает до OOM, пока ещё можно спасти ран:
import subprocessdef 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.