File size: 6,430 Bytes
3275441 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 | """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,
}
|