ravel / scripts /run_vf_sensitivity_multiseed.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
12.5 kB
#!/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()