Self-Forcing / scripts /plot_frrf_chunk_metrics.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
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()