168 lines
6.8 KiB
Python
Executable File
168 lines
6.8 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
"""Aggregate paired-video SSIM JSON files into a comparison report."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import math
|
|
import statistics
|
|
from collections import defaultdict
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
|
|
def percentile(values: list[float], q: float) -> float:
|
|
ordered = sorted(values)
|
|
position = (len(ordered) - 1) * q
|
|
lower, upper = math.floor(position), math.ceil(position)
|
|
if lower == upper:
|
|
return ordered[lower]
|
|
return ordered[lower] * (upper - position) + ordered[upper] * (position - lower)
|
|
|
|
|
|
def summarize(pairs: list[dict[str, Any]], threshold: float) -> dict[str, Any]:
|
|
values = [float(pair["ssim_all_mean"]) for pair in pairs]
|
|
return {
|
|
"videos": len(values),
|
|
"mean_video_ssim": statistics.fmean(values),
|
|
"median_video_ssim": statistics.median(values),
|
|
"p10_video_ssim": percentile(values, 0.10),
|
|
"min_video_ssim": min(values),
|
|
"videos_at_or_above_threshold": sum(value >= threshold for value in values),
|
|
"pass_rate": sum(value >= threshold for value in values) / len(values),
|
|
}
|
|
|
|
|
|
def parse_series(text: str) -> tuple[str, str, Path]:
|
|
try:
|
|
scheme, task, path = text.split("=", 2)
|
|
except ValueError as error:
|
|
raise argparse.ArgumentTypeError("series must be SCHEME=TASK=PATH") from error
|
|
return scheme, task, Path(path)
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--series", action="append", type=parse_series, required=True)
|
|
parser.add_argument("--output-dir", type=Path, required=True)
|
|
parser.add_argument("--threshold", type=float, default=0.90)
|
|
parser.add_argument("--reference", required=True)
|
|
args = parser.parse_args()
|
|
|
|
by_scheme: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
|
sources = []
|
|
for scheme, task, path in args.series:
|
|
payload = json.loads(path.read_text(encoding="utf-8"))
|
|
pairs = payload["pairs"]
|
|
if not pairs or {pair["task"] for pair in pairs} != {task}:
|
|
raise SystemExit(f"task mismatch for {path}: expected {task}")
|
|
by_scheme[scheme].extend(pairs)
|
|
sources.append({"scheme": scheme, "task": task, "path": str(path.resolve())})
|
|
|
|
video_counts = {scheme: len(pairs) for scheme, pairs in by_scheme.items()}
|
|
if len(set(video_counts.values())) != 1:
|
|
raise SystemExit(f"schemes have different video counts: {video_counts}")
|
|
base_videos = next(iter(video_counts.values()))
|
|
report: dict[str, Any] = {
|
|
"metric": "FFmpeg decoded YUV420 SSIM All; base self-comparison = 1.0",
|
|
"threshold": args.threshold,
|
|
"reference": args.reference,
|
|
"sources": sources,
|
|
"schemes": {
|
|
"base": {
|
|
"overall": {
|
|
"videos": base_videos,
|
|
"mean_video_ssim": 1.0,
|
|
"median_video_ssim": 1.0,
|
|
"p10_video_ssim": 1.0,
|
|
"min_video_ssim": 1.0,
|
|
"videos_at_or_above_threshold": base_videos,
|
|
"pass_rate": 1.0,
|
|
}
|
|
}
|
|
},
|
|
}
|
|
table_rows = []
|
|
for scheme, pairs in sorted(by_scheme.items()):
|
|
task_groups: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
|
resolution_groups: dict[int, list[dict[str, Any]]] = defaultdict(list)
|
|
for pair in pairs:
|
|
task_groups[str(pair["task"])].append(pair)
|
|
resolution_groups[int(pair["short_edge"])].append(pair)
|
|
scheme_result = {
|
|
"overall": summarize(pairs, args.threshold),
|
|
"by_task": {
|
|
task: summarize(group, args.threshold)
|
|
for task, group in sorted(task_groups.items())
|
|
},
|
|
"by_resolution": {
|
|
str(resolution): summarize(group, args.threshold)
|
|
for resolution, group in sorted(resolution_groups.items())
|
|
},
|
|
}
|
|
report["schemes"][scheme] = scheme_result
|
|
for group_type, groups in (
|
|
("overall", {"all": pairs}),
|
|
("task", task_groups),
|
|
("resolution", resolution_groups),
|
|
):
|
|
for group, members in groups.items():
|
|
table_rows.append(
|
|
{"scheme": scheme, "group_type": group_type, "group": group, **summarize(members, args.threshold)}
|
|
)
|
|
|
|
args.output_dir.mkdir(parents=True, exist_ok=True)
|
|
(args.output_dir / "comparison.json").write_text(
|
|
json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
|
)
|
|
columns = (
|
|
"scheme", "group_type", "group", "videos", "mean_video_ssim",
|
|
"median_video_ssim", "p10_video_ssim", "min_video_ssim",
|
|
"videos_at_or_above_threshold", "pass_rate",
|
|
)
|
|
with (args.output_dir / "comparison.tsv").open("w", encoding="utf-8") as handle:
|
|
handle.write("\t".join(columns) + "\n")
|
|
for row in table_rows:
|
|
handle.write("\t".join(str(row[column]) for column in columns) + "\n")
|
|
|
|
lines = [
|
|
"# MiniMax-H3 paired SSIM: base vs Cache-DiT vs Larry LoRA",
|
|
"",
|
|
f"Reference: `{args.reference}`",
|
|
"",
|
|
"Metric: FFmpeg-decoded YUV420 `SSIM All`, paired by request_id after exact prompt/seed/task/resolution/duration/aspect-ratio checks. Base self-comparison is 1.0.",
|
|
"",
|
|
"| Scheme | Videos | Mean | Median | Video P10 | Worst video | >= 0.90 |",
|
|
"|---|---:|---:|---:|---:|---:|---:|",
|
|
f"| base | {base_videos} | 1.000000 | 1.000000 | 1.000000 | 1.000000 | {base_videos}/{base_videos} |",
|
|
]
|
|
for scheme in sorted(by_scheme):
|
|
value = report["schemes"][scheme]["overall"]
|
|
lines.append(
|
|
f"| {scheme} | {value['videos']} | {value['mean_video_ssim']:.6f} | "
|
|
f"{value['median_video_ssim']:.6f} | {value['p10_video_ssim']:.6f} | "
|
|
f"{value['min_video_ssim']:.6f} | {value['videos_at_or_above_threshold']}/{value['videos']} |"
|
|
)
|
|
lines.extend(["", "## By task", ""])
|
|
for scheme in sorted(by_scheme):
|
|
for task, value in report["schemes"][scheme]["by_task"].items():
|
|
lines.append(
|
|
f"- {scheme} / {task}: mean={value['mean_video_ssim']:.6f}, "
|
|
f">=0.90={value['videos_at_or_above_threshold']}/{value['videos']}"
|
|
)
|
|
lines.extend(["", "## By resolution", ""])
|
|
for scheme in sorted(by_scheme):
|
|
values = report["schemes"][scheme]["by_resolution"]
|
|
rendered = ", ".join(
|
|
f"{resolution}p={value['mean_video_ssim']:.6f}"
|
|
for resolution, value in values.items()
|
|
)
|
|
lines.append(f"- {scheme}: {rendered}")
|
|
(args.output_dir / "README.md").write_text("\n".join(lines) + "\n", encoding="utf-8")
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|