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

Чекпойнт обучения / прогресс тренировки

Отличие от «дообучение завершилось»: там один финальный push, здесь — промежуточные. Тренировка своей модели идёт часами, и есть два момента, когда хочется знать не в конце, а сейчас:

  • веха — прошла эпоха, сохранился чекпойнт, val-loss улучшился;
  • авария — NaN/Inf в лоссе, early-stopping, лосс поехал вверх. Тут важна скорость: чем раньше остановите застрявший ран, тем меньше сожжёте GPU-часов.

Вешаем это прямо в цикл обучения через callback.

import os, math, 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)
class NotifyCallback:
def __init__(self, run_name, milestone_every=1):
self.run = run_name
self.every = milestone_every # слать веху раз в N эпох
self.best_val = math.inf
def on_epoch_end(self, epoch, train_loss, val_loss):
# 1) аварии — громко и сразу
if not math.isfinite(train_loss) or not math.isfinite(val_loss):
notify("🔥 NaN/Inf в лоссе",
f"{self.run}: эпоха {epoch}, train={train_loss}, val={val_loss}. "
f"Ран, скорее всего, испорчен — остановить и снизить LR.",
prio=10)
return
# 2) веха — тихо-нормально
if epoch % self.every == 0:
tag = "🟢 новый best" if val_loss < self.best_val else ""
notify(f"📊 Эпоха {epoch} {tag}".strip(),
f"{self.run}\ntrain_loss={train_loss:.4f}\n"
f"val_loss={val_loss:.4f} (best={min(val_loss, self.best_val):.4f})",
prio=4)
self.best_val = min(self.best_val, val_loss)
def on_early_stop(self, epoch, patience):
notify("🛑 Early-stopping сработал",
f"{self.run}: остановка на эпохе {epoch}, "
f"val-loss не улучшался {patience} эпох. Забрать лучший чекпойнт.",
prio=6)

Приоритеты подобраны под важность: веха — 4 (нормальный push, можно глянуть за кофе), early-stopping — 6 (штатное завершение, но требует действия), NaN — 10 (всё, ран горит, тушите).

Встраивание в обычный цикл:

cb = NotifyCallback("llama-lora-run7", milestone_every=1)
for epoch in range(EPOCHS):
tr = train_one_epoch(model, train_loader)
vl = validate(model, val_loader)
torch.save(model.state_dict(), f"ckpt/epoch{epoch}.pt")
cb.on_epoch_end(epoch, tr, vl)

Если тренируете через transformers.Trainer, не надо трогать цикл — есть готовая точка расширения:

from transformers import TrainerCallback
class NotiflyTrainerCallback(TrainerCallback):
def on_evaluate(self, args, state, control, metrics=None, **kw):
loss = (metrics or {}).get("eval_loss")
if loss is not None and not math.isfinite(loss):
notify("🔥 NaN eval_loss", f"step {state.global_step}", 10)
elif loss is not None:
notify("📊 Eval-чекпойнт",
f"step {state.global_step}, eval_loss={loss:.4f}", 4)

Отдельный сигнал от прогресса — что процесс не завис. Пусть цикл после каждого шага пингует heartbeat; если пинги прекратились (GPU-хост завис, OOM, дёрнули питание) — Notifly сам пришлёт алёрт по таймауту, даже когда код обучения уже ничего не может отправить.

requests.get(f"{NOTIFLY_URL}/heartbeat/H<token>", timeout=3) # раз в шаг/минуту
  • имя рана и номер эпохи/шага (для дедупликации);
  • train/val loss и текущий best;
  • learning rate и размер батча — чтобы понять причину аномалии;
  • путь к последнему чекпойнту в S3, чтобы забрать/откатить.