#!/usr/bin/env python3 """Plot 10-prompt mean FRRF metrics versus the reused chunk. The three input ``per_prompt.csv`` files use slightly different identifiers (``prompt_id`` for Self/Causal-Forcing and ``case`` for WorldPlay), but share the metric columns. We deliberately aggregate from the per-prompt rows so that PSNR is also an arithmetic mean over the ten prompts. """ from __future__ import annotations import argparse import csv from collections import defaultdict from pathlib import Path import matplotlib.pyplot as plt DEFAULT_ROOTS = { "Self-Forcing": Path( "/data3/chenzhuo/workspace/Self-Forcing/outputs/" "single_chunk_frrf_14chunks_first10" ), "Causal-Forcing": Path( "/data3/chenzhuo/workspace/Causal-Forcing/outputs/" "single_chunk_frrf_14chunks_first10" ), "HY-WorldPlay": Path( "/data3/chenzhuo/workspace/HY-WorldPlay-DEV/outputs/" "moviebench_single_chunk_frrf_14chunks_first10" ), } METRICS = ("psnr", "ssim", "lpips") Y_LABELS = {"psnr": "PSNR (dB)", "ssim": "SSIM", "lpips": "LPIPS"} COLORS = { "Self-Forcing": "#1f77b4", "Causal-Forcing": "#d62728", "HY-WorldPlay": "#2ca02c", } def read_prompt_means(root: Path, num_chunks: int = 14) -> dict[str, list[float]]: """Return arithmetic means over prompts for every metric and chunk.""" csv_path = root / "per_prompt.csv" if not csv_path.exists(): raise FileNotFoundError(csv_path) # values[chunk][metric] -> list of prompt-level values values: dict[int, dict[str, list[float]]] = defaultdict( lambda: {metric: [] for metric in METRICS} ) with csv_path.open(newline="") as handle: reader = csv.DictReader(handle) required = {"reuse_chunk", *METRICS} missing = required.difference(reader.fieldnames or ()) if missing: raise ValueError(f"{csv_path} is missing columns: {sorted(missing)}") for row in reader: chunk = int(row["reuse_chunk"]) if not 0 <= chunk < num_chunks: raise ValueError(f"unexpected reuse_chunk={chunk} in {csv_path}") for metric in METRICS: values[chunk][metric].append(float(row[metric])) result: dict[str, list[float]] = {} for metric in METRICS: means = [] for chunk in range(num_chunks): prompt_values = values[chunk][metric] if len(prompt_values) != 10: raise ValueError( f"{csv_path}: chunk {chunk} has {len(prompt_values)} rows; " "expected 10 prompts" ) means.append(sum(prompt_values) / len(prompt_values)) result[metric] = means return result def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--output", type=Path, default=Path( "/data3/chenzhuo/workspace/Self-Forcing/outputs/plots/" "frrf_14chunks_metrics_10prompt_mean.png" ), help="PNG output path (a PDF with the same stem is written too).", ) parser.add_argument("--num-chunks", type=int, default=14) return parser.parse_args() def main() -> None: args = parse_args() data = { label: read_prompt_means(root, args.num_chunks) for label, root in DEFAULT_ROOTS.items() } plt.rcParams.update( { "font.size": 11, "axes.labelsize": 12, "axes.titlesize": 13, "legend.fontsize": 10.5, "xtick.labelsize": 10, "ytick.labelsize": 10, "savefig.bbox": "tight", } ) fig, axes = plt.subplots(1, 3, figsize=(15.2, 4.7), sharex=True) chunks = list(range(args.num_chunks)) for axis, metric in zip(axes, METRICS): for label, values in data.items(): axis.plot( chunks, values[metric], color=COLORS[label], marker="o", markersize=4.5, linewidth=2.0, label=label, ) axis.set_title(metric.upper()) axis.set_xlabel("Reuse chunk") axis.set_ylabel(Y_LABELS[metric]) axis.set_xticks(chunks) axis.grid(True, linestyle="--", linewidth=0.7, alpha=0.35) axis.set_axisbelow(True) axis.spines["top"].set_visible(False) axis.spines["right"].set_visible(False) # One shared legend for all three panels. handles, labels = axes[0].get_legend_handles_labels() fig.legend( handles, labels, loc="upper center", bbox_to_anchor=(0.5, 0.995), ncol=3, frameon=False, ) fig.suptitle( "FRRF reuse-chunk error", y=1.045, fontsize=14, fontweight="semibold", ) fig.tight_layout(rect=(0, 0, 1, 1.0), w_pad=2.0) args.output.parent.mkdir(parents=True, exist_ok=True) fig.savefig(args.output, dpi=300) fig.savefig(args.output.with_suffix(".pdf")) print(f"saved {args.output}") print(f"saved {args.output.with_suffix('.pdf')}") if __name__ == "__main__": main()