File size: 4,541 Bytes
ec0a9aa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 | """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()
|