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, }