ravel / scripts /aggregate_ablation_multiseed.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
4.48 kB
#!/usr/bin/env python3
"""Aggregate ablation summaries across multiple seeds."""
from __future__ import annotations
import argparse
import csv
import json
import statistics
from pathlib import Path
from typing import Dict, List
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Aggregate ablation CSV files across seeds.")
parser.add_argument("--results-pattern", required=True, help="Pattern with {seed} placeholder")
parser.add_argument("--seeds", type=int, nargs="+", default=[1, 3, 5, 7, 11])
parser.add_argument("--output-dir", required=True)
parser.add_argument(
"--variant-order",
nargs="*",
default=["Full", "w/o Verification", "w/o Feedback", "w/o Co-Attention", "Text-only", "Vision-only", "w/o Text", "w/o Image"],
)
return parser.parse_args()
def agg(values: List[float]) -> tuple[float, float]:
if not values:
return float("nan"), float("nan")
if len(values) == 1:
return values[0], 0.0
return float(statistics.mean(values)), float(statistics.stdev(values))
def main() -> None:
args = parse_args()
metric_map: Dict[str, Dict[str, List[float]]] = {}
f1_key = ""
for seed in args.seeds:
path = Path(args.results_pattern.format(seed=seed))
if not path.exists():
raise FileNotFoundError(f"Missing ablation CSV for seed {seed}: {path}")
with path.open("r", encoding="utf-8") as f:
rows = list(csv.DictReader(f))
if not rows:
continue
if not f1_key:
f1_candidates = [k for k in rows[0].keys() if k.startswith("f1_")]
if not f1_candidates:
raise RuntimeError(f"Could not find F1 column in {path}")
f1_key = f1_candidates[0]
for row in rows:
variant = row["variant"]
metric_map.setdefault(variant, {"accuracy": [], f1_key: []})
metric_map[variant]["accuracy"].append(float(row["accuracy"]))
metric_map[variant][f1_key].append(float(row[f1_key]))
ordered_variants: List[str] = []
for variant in args.variant_order:
if variant in metric_map:
ordered_variants.append(variant)
for variant in metric_map.keys():
if variant not in ordered_variants:
ordered_variants.append(variant)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
rows_out: List[Dict[str, object]] = []
for variant in ordered_variants:
acc_mean, acc_std = agg(metric_map[variant]["accuracy"])
f1_mean, f1_std = agg(metric_map[variant][f1_key])
rows_out.append(
{
"variant": variant,
"accuracy_mean": acc_mean,
"accuracy_std": acc_std,
f"{f1_key}_mean": f1_mean,
f"{f1_key}_std": f1_std,
"n_seeds": len(metric_map[variant]["accuracy"]),
"accuracy_pm": f"{acc_mean * 100.0:.2f} ± {acc_std * 100.0:.2f}",
"f1_pm": f"{f1_mean * 100.0:.2f} ± {f1_std * 100.0:.2f}",
}
)
csv_path = output_dir / "ablation_multiseed_aggregate.csv"
with csv_path.open("w", encoding="utf-8", newline="") as f:
fieldnames = [
"variant",
"accuracy_mean",
"accuracy_std",
f"{f1_key}_mean",
f"{f1_key}_std",
"n_seeds",
"accuracy_pm",
"f1_pm",
]
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
for row in rows_out:
writer.writerow(row)
md_lines = [
f"| Variant | Accuracy (%) | {f1_key} (%) | n |",
"|---|---:|---:|---:|",
]
for row in rows_out:
md_lines.append(
"| {variant} | {acc} | {f1} | {n} |".format(
variant=row["variant"],
acc=row["accuracy_pm"],
f1=row["f1_pm"],
n=int(row["n_seeds"]),
)
)
md_path = output_dir / "ablation_multiseed_aggregate.md"
md_path.write_text("\n".join(md_lines) + "\n", encoding="utf-8")
json_path = output_dir / "ablation_multiseed_aggregate.json"
json_path.write_text(json.dumps(rows_out, indent=2), encoding="utf-8")
print(f"Saved CSV: {csv_path}")
print(f"Saved MD: {md_path}")
print(f"Saved JSON:{json_path}")
if __name__ == "__main__":
main()