"""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)