Download src/bigru_t/training/kill_switch.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 6.43 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/training/kill_switch.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/training/kill_switch.py
-
curl -L -o kill_switch.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/training/kill_switch.py
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 | |
| 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, | |
| } | |