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