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()