File size: 4,481 Bytes
ea8bfa1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
#!/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()