File size: 7,235 Bytes
8594de8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
"""memory_cleanup.py — Limpeza agressiva de memória estilo Xavante.

Reaproveita os padrões de:
  - xavante_work/flexnet/advanced_memory_cleanup.py (AdvancedMemoryCleaner)
  - xavante_work/flexnet/oom_guard.py (OomGuard daemon)
  - xavante_work/xavante/utils/timing.py (TimeBudget)

Implementa:
  1. production_cleanup() — context manager com cleanup garantido
  2. aggressive_cleanup() — gc.collect 3 gerações + torch.cuda.empty_cache
  3. TimeBudget — orçamento de tempo por treino/época (streaming com timed steps)
  4. get_rss_mb() — RSS do processo em MB
"""
from __future__ import annotations

import gc
import logging
import os
import threading
import time
from contextlib import contextmanager
from dataclasses import dataclass, field
from typing import Optional

logger = logging.getLogger(__name__)


def get_rss_mb() -> float:
    """Retorna o RSS (Resident Set Size) do processo atual em MB.

    Lê /proc/self/status (Linux). Fallback para psutil se disponível.
    """
    try:
        with open("/proc/self/status", "r") as f:
            for line in f:
                if line.startswith("VmRSS:"):
                    # VmRSS:    12345 kB
                    parts = line.split()
                    return float(parts[1]) / 1024.0
    except (FileNotFoundError, IndexError, ValueError):
        pass
    # Fallback psutil
    try:
        import psutil
        return psutil.Process(os.getpid()).memory_info().rss / (1024 * 1024)
    except ImportError:
        return 0.0


def aggressive_cleanup(verbose: bool = False) -> dict:
    """Limpeza agressiva de memória (estilo Xavante).

    Sequência:
      1. gc.collect gen 0, 1, 2 (3 gerações completas)
      2. torch.cuda.empty_cache() (se CUDA disponível)
      3. torch.cuda.synchronize() (se CUDA disponível)

    Args:
        verbose: logar memória antes/depois

    Returns:
        dict com rss_before_mb, rss_after_mb, freed_mb
    """
    rss_before = get_rss_mb()

    # 3 gerações de gc
    gc.collect(0)
    gc.collect(1)
    gc.collect(2)

    # CUDA cleanup (no-op se CPU-only)
    try:
        import torch
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
            torch.cuda.synchronize()
    except Exception:
        pass

    rss_after = get_rss_mb()
    freed = rss_before - rss_after

    if verbose:
        logger.info(
            "aggressive_cleanup: RSS %.1f → %.1f MB (freed %.1f MB)",
            rss_before, rss_after, freed,
        )
    return {
        "rss_before_mb": rss_before,
        "rss_after_mb": rss_after,
        "freed_mb": freed,
    }


@contextmanager
def production_cleanup(verbose: bool = False):
    """Context manager que garante cleanup agressivo ao sair (mesmo com exceção).

    Uso:
        with production_cleanup(verbose=True):
            # treino pesado aqui
            ...
        # cleanup automático ao sair do bloco
    """
    try:
        yield
    finally:
        aggressive_cleanup(verbose=verbose)


# ---------------------------------------------------------------------------
# TimeBudget — orçamento de tempo para treino/época/passos
# ---------------------------------------------------------------------------
@dataclass
class TimeBudget:
    """Orçamento de tempo para treino com timed steps.

    Permite:
      - max_total_s: tempo máximo total de treino
      - max_per_epoch_s: tempo máximo por época
      - is_over() / is_epoch_over(): verifica se orçamento estourou
      - remaining() / remaining_epoch(): tempo restante
      - reset_epoch(): reseta o contador de época (chamar no início de cada época)

    Uso:
        budget = TimeBudget(max_total_s=600, max_per_epoch_s=300)
        budget.reset_epoch()
        for epoch in range(2):
            for step in train_loop:
                if budget.is_over() or budget.is_epoch_over():
                    break
                ...
            budget.reset_epoch()
    """
    max_total_s: float = 600.0
    max_per_epoch_s: float = 300.0
    _start: float = field(default_factory=time.time, repr=False)
    _epoch_start: float = field(default_factory=time.time, repr=False)

    def reset_epoch(self):
        """Reseta o contador de época (chamar no início de cada época)."""
        self._epoch_start = time.time()

    def reset_total(self):
        """Reseta o contador total (chamar no início do treino)."""
        self._start = time.time()
        self._epoch_start = time.time()

    def elapsed(self) -> float:
        return time.time() - self._start

    def elapsed_epoch(self) -> float:
        return time.time() - self._epoch_start

    def remaining(self) -> float:
        return max(0.0, self.max_total_s - self.elapsed())

    def remaining_epoch(self) -> float:
        return max(0.0, self.max_per_epoch_s - self.elapsed_epoch())

    def is_over(self) -> bool:
        return self.elapsed() >= self.max_total_s

    def is_epoch_over(self) -> bool:
        return self.elapsed_epoch() >= self.max_per_epoch_s

    def should_save_partial(self) -> bool:
        """True se faltam < 10% do tempo (salvar parcial)."""
        return self.remaining() < (self.max_total_s * 0.1)

    def summary(self) -> dict:
        return {
            "elapsed_s": self.elapsed(),
            "remaining_s": self.remaining(),
            "epoch_elapsed_s": self.elapsed_epoch(),
            "epoch_remaining_s": self.remaining_epoch(),
            "is_over": self.is_over(),
            "is_epoch_over": self.is_epoch_over(),
        }


# ---------------------------------------------------------------------------
# StepTimer — mede tempo por passo com warning se lento
# ---------------------------------------------------------------------------
@dataclass
class StepTimer:
    """Mede tempo por passo e emite warnings se lento.

    Uso:
        timer = StepTimer(expected_s=0.5)
        for step in range(N):
            timer.start()
            # ... passo de treino ...
            timer.stop()  # loga se > 2x expected
    """
    expected_s: float = 0.5
    _start: float = 0.0
    _count: int = 0
    _total_s: float = 0.0
    _max_s: float = 0.0

    def start(self):
        self._start = time.time()

    def stop(self, step_label: str = "") -> float:
        elapsed = time.time() - self._start
        self._count += 1
        self._total_s += elapsed
        if elapsed > self._max_s:
            self._max_s = elapsed
        if elapsed > 2 * self.expected_s:
            logger.warning(
                "LENTO: step %s took %.2fs (expected ~%.2fs)",
                step_label or self._count, elapsed, self.expected_s,
            )
        elif elapsed < 0.5 * self.expected_s:
            logger.debug("RAPIDO: step %s took %.2fs", step_label, elapsed)
        return elapsed

    def avg(self) -> float:
        return self._total_s / max(1, self._count)

    def summary(self) -> dict:
        return {
            "count": self._count,
            "total_s": self._total_s,
            "avg_s": self.avg(),
            "max_s": self._max_s,
            "expected_s": self.expected_s,
        }


__all__ = [
    "get_rss_mb",
    "aggressive_cleanup",
    "production_cleanup",
    "TimeBudget",
    "StepTimer",
]