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