File size: 5,479 Bytes
eafbe80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
"""
Long-Horizon Consistency metrics:
- Stable sequence length (frames until quality collapse)
- Frame-to-frame drift rate (fraction of consecutive pairs below similarity threshold)
Optional: CLIP-based drift when available.
"""
from __future__ import annotations

import argparse
import json
import os
from typing import Any

from .common import discover_evals_videos, load_video_frames

try:
    from skimage.metrics import peak_signal_noise_ratio as psnr_skimage
    from skimage.metrics import structural_similarity as ssim_skimage
    HAS_SKIMAGE = True
except ImportError:
    HAS_SKIMAGE = False


def _psnr_simple(img1: "np.ndarray", img2: "np.ndarray", data_range: float = 255.0) -> float:
    """PSNR from MSE (same size, float images)."""
    import numpy as np
    mse = np.mean((img1.astype(np.float64) - img2.astype(np.float64)) ** 2)
    if mse <= 0:
        return 100.0
    return float(10 * np.log10((data_range ** 2) / mse))


def _frame_similarity_psnr(f1: "np.ndarray", f2: "np.ndarray") -> float:
    """Consecutive frame PSNR; resize f2 to f1 if shape mismatch."""
    import numpy as np
    if f1.shape != f2.shape:
        from PIL import Image
        import cv2
        h, w = f1.shape[:2]
        f2 = cv2.resize(f2, (w, h), interpolation=cv2.INTER_LINEAR)
    if HAS_SKIMAGE:
        return float(psnr_skimage(f1, f2, data_range=255))
    return _psnr_simple(f1, f2)


def compute_stable_length_and_drift(
    frames: "np.ndarray",
    collapse_psnr_threshold: float = 15.0,
    drift_psnr_threshold: float = 18.0,
) -> tuple[int, float, list[float]]:
    """
    Compute stable sequence length (number of frames until first collapse) and drift rate.
    - collapse: first frame index i where PSNR(frames[i], frames[i-1]) < collapse_psnr_threshold; length = that i (or len(frames) if never).
    - drift_rate: fraction of consecutive pairs with PSNR < drift_psnr_threshold.
    Returns (stable_length, drift_rate, list of consecutive PSNRs).
    """
    import numpy as np
    n = frames.shape[0]
    if n <= 1:
        return n, 0.0, []

    psnrs = []
    stable_length = n
    for i in range(1, n):
        p = _frame_similarity_psnr(frames[i - 1], frames[i])
        psnrs.append(p)
        if p < collapse_psnr_threshold and stable_length == n:
            stable_length = i  # collapse at frame i (0-indexed: frame i is first "bad")
    pairs = max(1, n - 1)
    below = sum(1 for p in psnrs if p < drift_psnr_threshold)
    drift_rate = below / pairs
    return stable_length, drift_rate, psnrs


def run_long_horizon_consistency(
    evals_root: str,
    collapse_psnr_threshold: float = 15.0,
    drift_psnr_threshold: float = 18.0,
    video_paths: list[tuple[str, str]] | None = None,
) -> dict[str, Any]:
    """
    Run long-horizon consistency metrics on evals_ep0 outputs.
    Returns dict with per_video results and aggregate.
    """
    if video_paths is None:
        video_paths = discover_evals_videos(evals_root)

    per_video = []
    all_stable_lengths = []
    all_drift_rates = []

    for rel, absp in video_paths:
        if not os.path.isfile(absp):
            continue
        frames = load_video_frames(absp)
        if frames.size == 0:
            per_video.append({"rel": rel, "stable_length": 0, "drift_rate": 0.0, "num_frames": 0})
            continue
        n = frames.shape[0]
        stable_length, drift_rate, psnrs = compute_stable_length_and_drift(
            frames, collapse_psnr_threshold, drift_psnr_threshold
        )
        all_stable_lengths.append(stable_length)
        all_drift_rates.append(drift_rate)
        per_video.append({
            "rel": rel,
            "stable_length": stable_length,
            "drift_rate": drift_rate,
            "num_frames": n,
            "mean_consecutive_psnr": float(sum(psnrs) / len(psnrs)) if psnrs else 0.0,
        })

    agg = {}
    if all_stable_lengths:
        agg["mean_stable_length"] = float(sum(all_stable_lengths) / len(all_stable_lengths))
        agg["min_stable_length"] = int(min(all_stable_lengths))
        agg["max_stable_length"] = int(max(all_stable_lengths))
    if all_drift_rates:
        agg["mean_drift_rate"] = float(sum(all_drift_rates) / len(all_drift_rates))
        agg["max_drift_rate"] = float(max(all_drift_rates))

    return {
        "dimension": "long_horizon_consistency",
        "params": {
            "collapse_psnr_threshold": collapse_psnr_threshold,
            "drift_psnr_threshold": drift_psnr_threshold,
        },
        "per_video": per_video,
        "aggregate": agg,
        "num_videos": len(per_video),
    }


def main():
    p = argparse.ArgumentParser(description="Long-Horizon Consistency metrics")
    p.add_argument("--evals_root", type=str, required=True, help="evals_ep0 root (e.g. ckpt_dir/evals_ep0)")
    p.add_argument("--collapse_threshold", type=float, default=15.0, help="PSNR below this = collapse")
    p.add_argument("--drift_threshold", type=float, default=18.0, help="PSNR below this = drift pair")
    p.add_argument("--output", type=str, default=None, help="Write JSON here")
    args = p.parse_args()

    result = run_long_horizon_consistency(
        args.evals_root,
        collapse_psnr_threshold=args.collapse_threshold,
        drift_psnr_threshold=args.drift_threshold,
    )
    out = json.dumps(result, indent=2)
    print(out)
    if args.output:
        with open(args.output, "w") as f:
            f.write(out)


if __name__ == "__main__":
    main()