PBC / pbc3_quick_rd.py
EgeEken's picture
pbc3: restore PIL resampling and fix benchmarks
08edfc9
Raw
History Blame Contribute Delete
9.65 kB
import json
import os
import subprocess
import sys
import threading
import time
import numpy as np
import optuna
from PIL import Image, ImageOps
import pbc3_sweep
from PBC3 import PBC3, preload_numba
from pbc3_types import PBC3Config
RESULTS_PATH = os.path.join(pbc3_sweep.PROJECT_DIR, "pbc3_quick_rd.json")
OLD_PBC_DIR = os.path.join(pbc3_sweep.PROJECT_DIR, "old_pbc")
PRESETS = ("compression", "balanced", "quality", "high_quality")
_LOCK = threading.Lock()
_STOP = threading.Event()
_THREAD = None
_STATE = {"running": False, "done": 0, "total": None, "current": None, "started": None, "error": None, "log": []}
def _log(msg):
with _LOCK:
_STATE["log"] = (_STATE["log"] + [f"{time.strftime('%H:%M:%S')} {msg}"])[-5000:]
print(f"[quick-rd] {msg}", flush=True)
def _new_bar():
with _LOCK:
_STATE["log"] = (_STATE["log"] + [""])[-5000:]
def _bar_tick():
with _LOCK:
if _STATE["log"]:
_STATE["log"][-1] += "|"
def status():
with _LOCK:
s = dict(_STATE)
s["log"] = list(_STATE["log"])
return s
def stop():
_STOP.set()
return {"stopping": True}
def start(n_trials=1, run_old=True):
global _THREAD
with _LOCK:
if _STATE["running"]:
return {"error": "A quick RD benchmark is already running."}
if not os.path.isdir(pbc3_sweep.DATA_DIR):
return {"error": "Dataset folder missing: hpt_data/"}
n_trials = max(1, min(10, int(n_trials or 1)))
run_old = bool(run_old)
if run_old and not os.path.exists(os.path.join(OLD_PBC_DIR, "PBC3.py")):
return {"error": "old_pbc/PBC3.py not found. Add the old runtime files or disable live old comparison."}
_STOP.clear()
_THREAD = threading.Thread(target=_run, args=(n_trials, run_old), daemon=True)
_THREAD.start()
return {"started": True, "n_trials": n_trials, "run_old": run_old}
def _mse(a, b):
return float(np.mean((a.astype(np.float32) - b.astype(np.float32)) ** 2))
def _avg(rows):
if not rows:
return {"speed": None, "bpp": None, "mse": None}
return {
"speed": float(np.mean([r["seconds"] for r in rows])),
"bpp": float(np.mean([r["bpp"] for r in rows])),
"mse": float(np.mean([r["mse"] for r in rows])),
}
def _eval_current_preset(name, images, trial):
rows = []
cfg = getattr(PBC3Config, name)()
for im in images:
if _STOP.is_set():
raise InterruptedError("quick RD stopped")
arr = im["arr"]
pixels = arr.shape[0] * arr.shape[1]
res = PBC3.compress(Image.fromarray(arr), config=cfg)
recon = np.asarray(res.image.convert("RGB").resize((arr.shape[1], arr.shape[0])))
rows.append({
"name": im["name"], "mp": im["mp"], "trial": trial,
"seconds": float(res.encode_seconds), "bpp": float(res.total_bits / pixels), "mse": _mse(arr, recon),
})
_bar_tick()
return rows
def _eval_old_preset_subprocess(name, images, trial):
payload = [{
"name": im["name"],
"path": im.get("path") or os.path.join(pbc3_sweep.DATA_DIR, im["name"]),
"mp": im["mp"],
} for im in images]
code = r'''
import json, os, sys, time
import numpy as np
from PIL import Image, ImageOps
sys.path.insert(0, sys.argv[1])
from PBC3 import PBC3
try:
from pbc3_types import PBC3Config
except Exception:
from pbc_types import PBC3Config
from PBC3 import preload_numba
preset = sys.argv[2]
trial = int(sys.argv[3])
images = json.loads(sys.stdin.read())
preload_numba()
out = []
for im in images:
img = ImageOps.exif_transpose(Image.open(im["path"])).convert("RGB")
arr = np.asarray(img)
pixels = arr.shape[0] * arr.shape[1]
cfg = getattr(PBC3Config, preset)()
res = PBC3.compress(img, config=cfg)
recon = np.asarray(res.image.convert("RGB").resize((arr.shape[1], arr.shape[0])))
mse = float(np.mean((arr.astype(np.float32) - recon.astype(np.float32)) ** 2))
out.append({"name": im["name"], "mp": im["mp"], "trial": trial, "seconds": float(res.encode_seconds), "bpp": float(res.total_bits / pixels), "mse": mse})
print(json.dumps(out))
'''
env = os.environ.copy()
env["PYTHONPATH"] = OLD_PBC_DIR + os.pathsep + env.get("PYTHONPATH", "")
p = subprocess.run(
[sys.executable, "-c", code, OLD_PBC_DIR, name, str(trial)],
input=json.dumps(payload), text=True, capture_output=True, cwd=pbc3_sweep.PROJECT_DIR, env=env,
)
if p.returncode != 0:
raise RuntimeError(f"old_pbc {name} failed: {p.stderr[-1200:]}")
for _ in images:
_bar_tick()
return json.loads(p.stdout)
def _run(n_trials, run_old):
total = len(PRESETS) * n_trials * (2 if run_old else 1)
with _LOCK:
_STATE.update(running=True, done=0, total=total, current=None, started=time.time(), error=None, log=[])
done = 0
try:
images = pbc3_sweep.dataset()
if not images:
raise RuntimeError("No images found in hpt_data/")
preload_numba(os.path.join(pbc3_sweep.PROJECT_DIR, "patch_policy.npz"))
_log(f"loaded {len(images)} images · {n_trials} trial(s)" + (" · live old enabled" if run_old else ""))
out = {"created": time.time(), "n_trials": n_trials, "presets": []}
for trial in range(1, n_trials + 1):
for source, evaluator in (("old_live", _eval_old_preset_subprocess), ("new", _eval_current_preset)):
if source == "old_live" and not run_old:
continue
for name in PRESETS:
if _STOP.is_set():
raise InterruptedError("quick RD stopped")
with _LOCK:
_STATE["current"] = {"source": source, "preset": name, "trial": trial}
_log(f"trial {trial}/{n_trials} {source} {name} started")
_new_bar()
rows = evaluator(name, images, trial)
values = _avg(rows)
out["presets"].append({"source": source, "kind": "preset", "preset": name, "trial": trial, "rows": rows, **values})
done += 1
with _LOCK:
_STATE["done"] = done
_log(f"trial {trial}/{n_trials} {source} {name} done | MSE {values['mse']:.3f} | bpp {values['bpp']:.5f} | speed {values['speed']:.3f}s")
with open(RESULTS_PATH, "w", encoding="utf-8") as f:
json.dump(out, f)
except Exception as e:
with _LOCK:
_STATE["error"] = str(e)
_log(f"ERROR: {e}")
finally:
with _LOCK:
_STATE["running"] = False
_STATE["current"] = None
_log("quick RD stopped")
def _trial_row(t, mp_min=None, mp_max=None):
ua = t.user_attrs
cfg = ua.get("config") or {}
rows = ua.get("per_image") or []
if mp_min is not None:
rows = [r for r in rows if float(r.get("mp", 0)) >= float(mp_min)]
if mp_max is not None:
rows = [r for r in rows if float(r.get("mp", 0)) <= float(mp_max)]
if rows:
vals = _avg(rows)
elif t.values and len(t.values) == 3:
vals = {"speed": float(t.values[0]), "bpp": float(t.values[1]), "mse": float(t.values[2])}
else:
return None
kind = ua.get("kind")
if ua.get("preset"):
return {"source": "old", "kind": "preset", "preset": ua.get("preset"), **vals}
if cfg.get("codec"):
return {"source": "codec", "kind": "codec", "codec": cfg.get("codec"), "q": cfg.get("q"), **vals}
if kind == "baseline" and ua.get("baseline"):
return {"source": "codec", "kind": "codec", "codec": str(ua.get("baseline")).split("_")[0].upper(), "q": cfg.get("q"), **vals}
return None
def _load_old_db(mp_min=None, mp_max=None):
if not os.path.exists(pbc3_sweep.DB_PATH):
return {"exists": False, "rows": []}
try:
study = optuna.load_study(study_name=pbc3_sweep.STUDY, storage=pbc3_sweep.STORAGE)
except Exception:
return {"exists": False, "rows": []}
rows = []
for t in study.trials:
if t.state != optuna.trial.TrialState.COMPLETE:
continue
ua = t.user_attrs
baseline = ua.get("baseline") or ""
cfg = ua.get("config") or {}
if ua.get("preset") or baseline.startswith("preset_") or cfg.get("codec"):
r = _trial_row(t, mp_min, mp_max)
if r:
rows.append(r)
return {"exists": True, "rows": rows}
def _load_live(mp_min=None, mp_max=None):
if not os.path.exists(RESULTS_PATH):
return []
with open(RESULTS_PATH, "r", encoding="utf-8") as f:
data = json.load(f)
grouped = {}
for p in data.get("presets", []):
rows = p.get("rows") or []
if mp_min is not None:
rows = [r for r in rows if float(r.get("mp", 0)) >= float(mp_min)]
if mp_max is not None:
rows = [r for r in rows if float(r.get("mp", 0)) <= float(mp_max)]
key = (p.get("source"), p.get("preset"))
grouped.setdefault(key, []).extend(rows)
out = []
for (source, preset), rows in grouped.items():
out.append({"source": source, "kind": "preset", "preset": preset, **_avg(rows)})
return out
def results(mp_min=None, mp_max=None):
old = _load_old_db(mp_min, mp_max)
live = _load_live(mp_min, mp_max)
return {
"old_db_exists": old["exists"],
"new_exists": os.path.exists(RESULTS_PATH),
"old_pbc_exists": os.path.exists(os.path.join(OLD_PBC_DIR, "PBC3.py")),
"rows": old["rows"] + live,
}