"""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, }