BiGRU_T_version / src /bigru_t /training /kill_switch.py
PowerMachine's picture
Upload folder using huggingface_hub
3275441 verified
Raw History Blame Contribute Delete
6.43 kB
"""kill_switch.py — Monitor de saúde do treino + kill automático.
Implementa a regra do usuário:
"monitorar e matar o modelo se não estiver aprendendo"
"monitorar consumo de RAM, consumo de Armazenamento, a loss e a perplexidade"
Critérios de kill:
1. RAM > 90% da disponível (psutil.virtual_memory)
2. Disco < 1 GB livre (psutil.disk_usage)
3. Loss não diminui em N_steps_patience (default 30) — comparado com
min_loss_so_far
4. Loss é NaN ou Inf
5. Perplexidade > 1e6 (explosão)
Reaproveita a filosofia do flexnet/oom_guard.py (thread-based) e
flexnet/memory_monitor.py (leak detection), mas unifica em uma classe
síncrona mais simples para o treino de bug-detection.
"""
from __future__ import annotations
import math
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
import psutil
import torch
@dataclass
class KillSwitchState:
"""Snapshot do estado do kill-switch em um dado step."""
step: int
loss: float
ppl: float
ram_pct: float
ram_used_gb: float
ram_total_gb: float
disk_free_gb: float
disk_total_gb: float
active_modules: int
reason: Optional[str] = None # None = OK, string = motivo do kill
class KillSwitch:
"""Monitor de saúde do treino com kill automático.
Args:
ram_threshold_pct: kill se RAM usage > este pct (default 90)
disk_min_free_gb: kill se disco livre < este valor (default 1.0)
loss_patience: nº de steps sem melhoria antes de matar (default 30)
loss_tolerance: loss considerada "melhoria" se cair mais que isto (default 1e-4)
max_ppl: kill se ppl > este valor (default 1e6)
log_dir: diretório para salvar logs de monitoramento (default /tmp)
Uso:
ks = KillSwitch()
for step, batch in enumerate(loader):
loss = train_step(batch)
state = ks.check(step, loss, active_modules)
if state.reason:
logger.error(f"KILL: {state.reason}")
break
"""
def __init__(
self,
ram_threshold_pct: float = 90.0,
disk_min_free_gb: float = 1.0,
loss_patience: int = 30,
loss_tolerance: float = 1e-4,
max_ppl: float = 1e6,
log_dir: Optional[str] = None,
):
self.ram_threshold_pct = ram_threshold_pct
self.disk_min_free_gb = disk_min_free_gb
self.loss_patience = loss_patience
self.loss_tolerance = loss_tolerance
self.max_ppl = max_ppl
self.log_dir = Path(log_dir) if log_dir else None
# Estado interno
self.min_loss: float = float("inf")
self.steps_since_improvement: int = 0
self.history: list[KillSwitchState] = []
# Snapshot do processo para RAM
self._process = psutil.Process(os.getpid())
def _get_ram_usage(self) -> tuple[float, float, float]:
"""Retorna (ram_pct, ram_used_gb, ram_total_gb)."""
vm = psutil.virtual_memory()
return vm.percent, vm.used / 1e9, vm.total / 1e9
def _get_disk_usage(self) -> tuple[float, float]:
"""Retorna (disk_free_gb, disk_total_gb) para o disco do projeto."""
# Usa o disco onde /home/z/my-project está
path = "/home/z/my-project"
try:
du = psutil.disk_usage(path)
return du.free / 1e9, du.total / 1e9
except Exception:
return float("inf"), float("inf")
def check(
self,
step: int,
loss: float,
active_modules: int = 0,
) -> KillSwitchState:
"""Verifica saúde do treino e retorna estado.
Args:
step: nº do step atual
loss: loss atual (escalar)
active_modules: nº de módulos ativos no modelo (para logging)
Returns:
KillSwitchState com .reason preenchido se kill recomendado.
"""
# Calcula ppl (com clipping para evitar overflow)
if math.isnan(loss) or math.isinf(loss):
ppl = float("inf")
else:
ppl = math.exp(min(20, loss)) if loss < 20 else float("inf")
ram_pct, ram_used, ram_total = self._get_ram_usage()
disk_free, disk_total = self._get_disk_usage()
# Atualiza min_loss e contador de patience
reason = None
if math.isnan(loss) or math.isinf(loss):
reason = f"loss is NaN/Inf at step {step}"
elif loss < self.min_loss - self.loss_tolerance:
self.min_loss = loss
self.steps_since_improvement = 0
else:
self.steps_since_improvement += 1
if self.steps_since_improvement >= self.loss_patience:
reason = (
f"loss not improving for {self.loss_patience} steps "
f"(min={self.min_loss:.4f}, current={loss:.4f})"
)
# Critério de RAM
if ram_pct > self.ram_threshold_pct:
reason = reason or f"RAM {ram_pct:.1f}% > threshold {self.ram_threshold_pct}%"
# Critério de disco
if disk_free < self.disk_min_free_gb:
reason = reason or f"disk free {disk_free:.2f}GB < min {self.disk_min_free_gb}GB"
# Critério de perplexidade
if ppl > self.max_ppl:
reason = reason or f"ppl {ppl:.2e} > max {self.max_ppl:.2e}"
state = KillSwitchState(
step=step,
loss=loss,
ppl=ppl,
ram_pct=ram_pct,
ram_used_gb=ram_used,
ram_total_gb=ram_total,
disk_free_gb=disk_free,
disk_total_gb=disk_total,
active_modules=active_modules,
reason=reason,
)
self.history.append(state)
return state
def summary(self) -> dict:
"""Retorna resumo final do monitoramento."""
if not self.history:
return {}
return {
"total_steps": len(self.history),
"min_loss": min(s.loss for s in self.history if not math.isnan(s.loss) and not math.isinf(s.loss)),
"max_ram_pct": max(s.ram_pct for s in self.history),
"min_disk_free_gb": min(s.disk_free_gb for s in self.history),
"max_ppl": max(s.ppl for s in self.history if not math.isinf(s.ppl)),
"killed": self.history[-1].reason is not None,
"kill_reason": self.history[-1].reason,
}