Download scripts/plot_frrf_chunk_metrics.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 5.22 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/plot_frrf_chunk_metrics.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/plot_frrf_chunk_metrics.py
-
curl -L -o plot_frrf_chunk_metrics.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/plot_frrf_chunk_metrics.py
5.22 kB
| #!/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() | |