Чекпойнт обучения / прогресс тренировки
Отличие от «дообучение завершилось»: там один финальный push, здесь — промежуточные. Тренировка своей модели идёт часами, и есть два момента, когда хочется знать не в конце, а сейчас:
- веха — прошла эпоха, сохранился чекпойнт, val-loss улучшился;
- авария — NaN/Inf в лоссе, early-stopping, лосс поехал вверх. Тут важна скорость: чем раньше остановите застрявший ран, тем меньше сожжёте GPU-часов.
Вешаем это прямо в цикл обучения через callback.
PyTorch: колбэк на конец эпохи и на аномалию
Заголовок раздела «PyTorch: колбэк на конец эпохи и на аномалию»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)HuggingFace Trainer: тот же смысл через TrainerCallback
Заголовок раздела «HuggingFace Trainer: тот же смысл через TrainerCallback»Если тренируете через 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: «тренировка вообще жива»
Заголовок раздела «Heartbeat: «тренировка вообще жива»»Отдельный сигнал от прогресса — что процесс не завис. Пусть цикл после каждого шага пингует heartbeat; если пинги прекратились (GPU-хост завис, OOM, дёрнули питание) — Notifly сам пришлёт алёрт по таймауту, даже когда код обучения уже ничего не может отправить.
requests.get(f"{NOTIFLY_URL}/heartbeat/H<token>", timeout=3) # раз в шаг/минутуЧто положить в текст алёрта
Заголовок раздела «Что положить в текст алёрта»- имя рана и номер эпохи/шага (для дедупликации);
- train/val loss и текущий best;
- learning rate и размер батча — чтобы понять причину аномалии;
- путь к последнему чекпойнту в S3, чтобы забрать/откатить.