HAKO-v1 / hako /memory_manager.py
PowerMachine's picture
HAKO upload: hako/memory_manager.py
e4e6b61 verified
Raw History Blame Contribute Delete
4.72 kB
"""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)