File size: 4,719 Bytes
e4e6b61
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Aggressive memory & storage management for HAKO.

Proposition P-MEM (bounded-resource operation). If every phase satisfies
    (i)  RSS(t) <= ram_soft_cap for all t (enforced by watchdog degrade),
    (ii) disk blobs not referenced by decomposed artifacts are deleted,
    (iii) telemetry/checkpoints are size-capped and compressed,
then the pipeline operates indefinitely within a fixed resource envelope.
Proof sketch: (i) bounds the state space of the resident set by construction
(degrade-or-abort policy with monotone batch reduction); (ii) makes total disk
usage a sum of compressed artifacts (monotone, bounded by artifact budget);
(iii) bounds every log by its cap. Induction over phases completes the bound.
"""
from __future__ import annotations

import gc
import logging
import os
import shutil
import time
from pathlib import Path
from typing import Iterable

import psutil

log = logging.getLogger("hako.mem")


class MemoryManager:
    def __init__(self, cfg) -> None:
        self.cfg = cfg
        vm = psutil.virtual_memory()
        self.ram_soft_cap = int(cfg.ram_soft_cap_ratio * vm.total)
        self._last_report = 0.0

    # ------------------------------------------------------------------ ram
    def rss(self) -> int:
        return psutil.Process(os.getpid()).memory_info().rss

    def pressure(self) -> float:
        """RSS pressure in [0, 1+] relative to the soft cap."""
        return self.rss() / max(1, self.ram_soft_cap)

    def check(self, degrade_hooks: Iterable | None = None) -> str:
        """Watchdog: returns 'ok' | 'degraded'. Calls hooks to shrink state."""
        p = self.pressure()
        if p < 0.85:
            return "ok"
        gc.collect()
        if p < 1.0:
            return "ok"
        action = "degraded"
        for hook in degrade_hooks or []:
            try:
                hook()
            except Exception as exc:  # pragma: no cover - defensive
                log.warning("degrade hook failed: %s", exc)
        gc.collect()
        if self.pressure() > 1.05:
            raise MemoryError(
                f"RSS {self.rss()/2**30:.2f}GiB exceeded soft cap "
                f"{self.ram_soft_cap/2**30:.2f}GiB after degradation")
        return action

    def report(self, tag: str = "") -> None:
        now = time.time()
        if now - self._last_report < 15:
            return
        self._last_report = now
        vm = psutil.virtual_memory()
        log.info("mem[%s] rss=%.2fGiB cap=%.2fGiB sys_avail=%.2fGiB",
                 tag, self.rss() / 2**30, self.ram_soft_cap / 2**30,
                 vm.available / 2**30)

    # ----------------------------------------------------------------- disk
    def disk_free(self) -> int:
        return psutil.disk_usage("/").free

    def ensure_disk(self, need_bytes: int) -> None:
        if self.disk_free() >= max(need_bytes, self.cfg.disk_min_free_bytes):
            return
        freed = self.sweep_caches()
        log.info("disk sweep freed %.2f MiB", freed / 2**20)
        if self.disk_free() < max(need_bytes, self.cfg.disk_min_free_bytes):
            raise OSError("disk exhausted even after cache sweep")

    def sweep_caches(self) -> int:
        """Delete HF hub blobs and work leftovers that have decomposed twins."""
        freed = 0
        hf_blob_dirs = [
            Path(self.cfg.hf_cache) / "hub",
        ]
        for base in hf_blob_dirs:
            if not base.exists():
                continue
            for model_dir in base.iterdir():
                marker = Path(self.cfg.artifacts_dir) / (
                    model_dir.name.replace("models--", "") + ".decomposed.json")
                if marker.exists():
                    # decomposed copy exists -> raw weights are expendable
                    size = sum(f.stat().st_size for f in model_dir.rglob("*")
                               if f.is_file())
                    shutil.rmtree(model_dir, ignore_errors=True)
                    freed += size
        gc.collect()
        return freed

    @staticmethod
    def cap_file_size(path: Path, max_bytes: int, trim_head: bool = True) -> None:
        """Keep a log under max_bytes by trimming the oldest half (JSONL-safe)."""
        p = Path(path)
        if not p.exists() or p.stat().st_size <= max_bytes:
            return
        lines = p.read_text(encoding="utf-8", errors="ignore").splitlines(keepends=True)
        keep = lines[len(lines) // 2:] if trim_head else lines[:len(lines) // 2]
        tmp = p.with_suffix(p.suffix + ".tmp")
        tmp.write_text("".join(keep), encoding="utf-8")
        tmp.replace(p)

    @staticmethod
    def rm_tree_quiet(path: Path) -> None:
        shutil.rmtree(path, ignore_errors=True)