| |
| """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() |
|
|