"""Compare grid search results across ITM and SiTo configs. Usage: python scripts/compare_grid_results.py python scripts/compare_grid_results.py --grid-root tmp_2/grid_search tmp_2/grid_search_sito """ import argparse import json from pathlib import Path def find_summaries(grid_root: Path): results = [] for summary_file in sorted(grid_root.rglob("all_summary.json")): rel = summary_file.relative_to(grid_root) config_name = rel.parts[0] with open(summary_file) as f: data = json.load(f) results.append((config_name, str(rel.parent), data)) return results def parse_timing_log(timing_log: Path): timings = {} current_config = None if not timing_log.exists(): return timings with open(timing_log) as f: for line in f: line = line.strip() if line.startswith("--- Config:"): current_config = line.split("Config:")[1].strip().rstrip(" ---") elif line.startswith("s/it:") and current_config: timings[current_config] = float(line.split(":")[1].strip()) elif line.startswith("Total time:") and current_config: timings.setdefault(current_config, {}) if isinstance(timings[current_config], dict): timings[current_config] = None return timings def main(): parser = argparse.ArgumentParser() parser.add_argument( "--grid-root", nargs="+", default=["tmp_2/grid_search", "tmp_2/grid_search_sito"], ) args = parser.parse_args() workspace = Path(__file__).resolve().parent.parent all_results = [] all_timings = {} for grid_dir in args.grid_root: grid_path = workspace / grid_dir if not grid_path.exists(): continue results = find_summaries(grid_path) all_results.extend([(grid_dir, *r) for r in results]) timing_log = grid_path / "timing_log.txt" if timing_log.exists(): with open(timing_log) as f: content = f.read() current_config = None for line in content.split("\n"): line = line.strip() if line.startswith("--- Config:"): current_config = line.split("Config:")[1].strip().rstrip(" ---") elif "s/it:" in line and current_config: try: s_per_it = float(line.split("s/it:")[1].strip()) all_timings[f"{grid_dir}/{current_config}"] = s_per_it except ValueError: pass if not all_results: print("No results found. Run grid search scripts first.") return print("=" * 90) print(f"{'Config':<35} {'PSNR↑':>8} {'SSIM↑':>8} {'LPIPS↓':>8} {'s/it':>8} {'Speedup':>8}") print("=" * 90) baseline_time = None rows = [] for grid_dir, config_name, rel_path, data in all_results: psnr = data.get("psnr_mean", data.get("psnr", "N/A")) ssim = data.get("ssim_mean", data.get("ssim", "N/A")) lpips = data.get("lpips_mean", data.get("lpips", "N/A")) key = f"{grid_dir}/{config_name}" s_it = all_timings.get(key) if "dense_baseline" in config_name and s_it is not None: baseline_time = s_it rows.append((config_name, psnr, ssim, lpips, s_it, grid_dir)) for config_name, psnr, ssim, lpips, s_it, grid_dir in rows: psnr_str = f"{psnr:.2f}" if isinstance(psnr, (int, float)) else str(psnr) ssim_str = f"{ssim:.4f}" if isinstance(ssim, (int, float)) else str(ssim) lpips_str = f"{lpips:.4f}" if isinstance(lpips, (int, float)) else str(lpips) s_it_str = f"{s_it:.2f}s" if s_it else "N/A" speedup_str = "" if s_it and baseline_time and baseline_time > 0: speedup = baseline_time / s_it speedup_str = f"{speedup:.2f}x" label = config_name if "sito" in grid_dir: label = f"[SiTo] {config_name}" elif "grid_search" in grid_dir and "sito" not in grid_dir: label = f"[ITM] {config_name}" if "itm" in config_name else f" {config_name}" print(f"{label:<35} {psnr_str:>8} {ssim_str:>8} {lpips_str:>8} {s_it_str:>8} {speedup_str:>8}") print("=" * 90) print("\n↑ = higher is better, ↓ = lower is better") if baseline_time: print(f"Baseline s/it: {baseline_time:.2f}s") if __name__ == "__main__": main()