sskj-h3/throughput/common/summarize_ssim_comparison.py
2026-08-31 15:57:13 +08:00

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