#!/usr/bin/env python3 """Run verification-feedback sensitivity across multiple seeds and aggregate.""" from __future__ import annotations import argparse import csv import json import statistics import subprocess import sys from collections import defaultdict from pathlib import Path from typing import Dict, List SCRIPT_DIR = Path(__file__).resolve().parent def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Multi-seed sensitivity analysis runner.") parser.add_argument("--dataset", choices=["hfm", "mvsa_multiple"], required=True) parser.add_argument("--seeds", type=int, nargs="+", default=[1, 3, 5, 7, 11]) parser.add_argument("--data-root", default=None) parser.add_argument("--text-dir", default=None) parser.add_argument("--label-file", default=None) parser.add_argument("--outputs-root", default=None) parser.add_argument("--results-root", default=None) parser.add_argument("--checkpoint-name", default=None) parser.add_argument("--batch-size", type=int, default=None) parser.add_argument("--max-epochs", type=int, default=50) parser.add_argument("--learning-rate", type=float, default=None) parser.add_argument("--weight-decay", type=float, default=0.01) parser.add_argument("--warmup-ratio", type=float, default=None) parser.add_argument("--early-stopping-patience", type=int, default=None) parser.add_argument("--grad-accum-steps", type=int, default=None) parser.add_argument("--num-workers", type=int, default=6) parser.add_argument("--max-length", type=int, default=128) parser.add_argument("--train-ratio", type=float, default=0.8) parser.add_argument("--val-ratio", type=float, default=0.1) parser.add_argument("--preprocessing-mode", choices=["paper", "strict"], default="paper") parser.add_argument("--device", default="auto") parser.add_argument("--train-if-missing", action="store_true") parser.add_argument("--skip-existing", action="store_true") parser.add_argument("--dry-run", action="store_true") return parser.parse_args() def run_cmd(cmd: List[str], dry_run: bool = False) -> None: print(" ".join(cmd)) if dry_run: return subprocess.run(cmd, check=True) def read_rows(csv_path: Path) -> List[Dict[str, str]]: with csv_path.open("r", encoding="utf-8") as f: return list(csv.DictReader(f)) 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() py = sys.executable if args.dataset == "hfm": data_root = args.data_root or "data/HFM" text_dir = args.text_dir or str(Path(data_root) / "text") outputs_root = Path(args.outputs_root or "outputs/reviewer3_hfm_baseline") results_root = Path(args.results_root or "results/reviewer3_hfm_sensitivity") checkpoint_name = args.checkpoint_name or "clara_hfm.pt" else: data_root = args.data_root or "data/MVSA-Multiple" text_dir = args.text_dir or str(Path(data_root) / "data") label_file = args.label_file or str(Path(data_root) / "labelResultAll.txt") outputs_root = Path(args.outputs_root or "outputs/reviewer3_mvsa_multiple_baseline") results_root = Path(args.results_root or "results/reviewer3_mvsa_multiple_sensitivity") checkpoint_name = args.checkpoint_name or "clara_mvsa_multiple.pt" outputs_root.mkdir(parents=True, exist_ok=True) results_root.mkdir(parents=True, exist_ok=True) per_seed_files: List[Path] = [] for seed in args.seeds: seed_output = outputs_root / f"seed_{seed}" seed_result = results_root / f"seed_{seed}" seed_result.mkdir(parents=True, exist_ok=True) ckpt_path = seed_output / checkpoint_name sensitivity_csv = seed_result / "verification_feedback_sensitivity.csv" if args.skip_existing and sensitivity_csv.exists(): per_seed_files.append(sensitivity_csv) continue if not ckpt_path.exists(): if not args.train_if_missing: raise FileNotFoundError( f"Missing checkpoint for seed {seed}: {ckpt_path}. " "Use --train-if-missing to train baseline checkpoints first." ) if args.dataset == "hfm": train_cmd = [ py, str(SCRIPT_DIR / "train_hfm.py"), "--data-root", data_root, "--text-dir", text_dir, "--output-dir", str(seed_output), "--checkpoint-name", checkpoint_name, "--max-epochs", str(args.max_epochs), "--weight-decay", str(args.weight_decay), "--num-workers", str(args.num_workers), "--seed", str(seed), "--hidden-dim", "512", "--lora-rank", "8", "--device", args.device, ] if args.batch_size is not None: train_cmd.extend(["--batch-size", str(args.batch_size)]) if args.learning_rate is not None: train_cmd.extend(["--learning-rate", str(args.learning_rate)]) if args.warmup_ratio is not None: train_cmd.extend(["--warmup-ratio", str(args.warmup_ratio)]) if args.early_stopping_patience is not None: train_cmd.extend(["--early-stopping-patience", str(args.early_stopping_patience)]) if args.max_length is not None: train_cmd.extend(["--max-length", str(args.max_length)]) else: train_cmd = [ py, str(SCRIPT_DIR / "train_mvsa_multiple.py"), "--data-root", data_root, "--text-dir", text_dir, "--label-file", label_file, "--output-dir", str(seed_output), "--checkpoint-name", checkpoint_name, "--max-epochs", str(args.max_epochs), "--weight-decay", str(args.weight_decay), "--num-workers", str(args.num_workers), "--max-length", str(args.max_length), "--seed", str(seed), "--train-ratio", str(args.train_ratio), "--val-ratio", str(args.val_ratio), "--preprocessing-mode", args.preprocessing_mode, "--hidden-dim", "512", "--lora-rank", "8", "--text-unfreeze-mode", "unfreeze_all", "--device", args.device, ] if args.batch_size is not None: train_cmd.extend(["--batch-size", str(args.batch_size)]) if args.learning_rate is not None: train_cmd.extend(["--learning-rate", str(args.learning_rate)]) if args.warmup_ratio is not None: train_cmd.extend(["--warmup-ratio", str(args.warmup_ratio)]) if args.early_stopping_patience is not None: train_cmd.extend(["--early-stopping-patience", str(args.early_stopping_patience)]) if args.grad_accum_steps is not None: train_cmd.extend(["--grad-accum-steps", str(args.grad_accum_steps)]) run_cmd(train_cmd, dry_run=args.dry_run) sensitivity_cmd = [ py, str(SCRIPT_DIR / "sensitivity_verification_feedback.py"), "--dataset", args.dataset, "--checkpoint", str(ckpt_path), "--data-root", data_root, "--text-dir", text_dir, "--output-dir", str(seed_result), "--num-workers", str(args.num_workers), "--device", args.device, ] if args.batch_size is not None: sensitivity_cmd.extend(["--batch-size", str(args.batch_size)]) if args.max_length is not None: sensitivity_cmd.extend(["--max-length", str(args.max_length)]) if args.dataset == "mvsa_multiple": sensitivity_cmd.extend( [ "--label-file", label_file, "--train-ratio", str(args.train_ratio), "--val-ratio", str(args.val_ratio), "--seed", str(seed), "--preprocessing-mode", args.preprocessing_mode, ] ) run_cmd(sensitivity_cmd, dry_run=args.dry_run) if not args.dry_run: per_seed_files.append(sensitivity_csv) if args.dry_run: return grouped: Dict[str, Dict[str, List[float]]] = defaultdict(lambda: defaultdict(list)) metric_name = "" for csv_path in per_seed_files: rows = read_rows(csv_path) if not rows: continue for row in rows: setting = row["setting"] grouped[setting]["accuracy"].append(float(row["accuracy"])) other_metrics = [k for k in row.keys() if k.startswith("f1_")] if other_metrics: metric_name = other_metrics[0] grouped[setting][metric_name].append(float(row[metric_name])) out_rows: List[Dict[str, object]] = [] for setting, value_map in grouped.items(): acc_mean, acc_std = agg(value_map["accuracy"]) metric_mean, metric_std = agg(value_map[metric_name]) if metric_name else (float("nan"), float("nan")) out_rows.append( { "setting": setting, "accuracy_mean": acc_mean, "accuracy_std": acc_std, f"{metric_name}_mean": metric_mean, f"{metric_name}_std": metric_std, "n_seeds": len(value_map["accuracy"]), } ) out_rows.sort(key=lambda x: str(x["setting"])) agg_csv = results_root / "verification_feedback_sensitivity_aggregate.csv" with agg_csv.open("w", encoding="utf-8", newline="") as f: fieldnames = [ "setting", "accuracy_mean", "accuracy_std", f"{metric_name}_mean", f"{metric_name}_std", "n_seeds", ] writer = csv.DictWriter(f, fieldnames=fieldnames) writer.writeheader() for row in out_rows: writer.writerow(row) agg_json = results_root / "verification_feedback_sensitivity_aggregate.json" agg_json.write_text(json.dumps(out_rows, indent=2), encoding="utf-8") md_lines = [ f"| Setting | Accuracy (%) | {metric_name} (%) | n |", "|---|---:|---:|---:|", ] for row in out_rows: md_lines.append( "| {setting} | {acc_m:.2f} ± {acc_s:.2f} | {f1_m:.2f} ± {f1_s:.2f} | {n} |".format( setting=row["setting"], acc_m=float(row["accuracy_mean"]) * 100.0, acc_s=float(row["accuracy_std"]) * 100.0, f1_m=float(row[f"{metric_name}_mean"]) * 100.0, f1_s=float(row[f"{metric_name}_std"]) * 100.0, n=int(row["n_seeds"]), ) ) agg_md = results_root / "verification_feedback_sensitivity_aggregate.md" agg_md.write_text("\n".join(md_lines) + "\n", encoding="utf-8") print(f"Saved aggregate CSV: {agg_csv}") print(f"Saved aggregate JSON: {agg_json}") print(f"Saved aggregate MD: {agg_md}") if __name__ == "__main__": main()