#!/usr/bin/env python3 """Sensitivity analysis for CLARA verification-feedback design choices.""" from __future__ import annotations import argparse import csv import json from pathlib import Path from typing import Any, Callable, Dict, List, Optional, Tuple import numpy as np from sklearn.metrics import accuracy_score, f1_score from transformers import CLIPProcessor, DebertaV2Tokenizer import sys PROJECT_ROOT = Path(__file__).resolve().parents[1] if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) from src.hfm_pipeline import ( HFMLoader, create_dataloaders as create_hfm_dataloaders, estimate_max_length as estimate_hfm_max_length, gather_logits_variant as gather_hfm_variant, load_checkpoint as load_hfm_checkpoint, resolve_device, ) from src.mvsa_multiple_pipeline import ( MVSALoader, create_dataloaders as create_mvsa_dataloaders, gather_logits_variant as gather_mvsa_variant, load_checkpoint as load_mvsa_checkpoint, ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Sensitivity analysis for verification-feedback module." ) parser.add_argument("--dataset", choices=["hfm", "mvsa_multiple"], required=True) parser.add_argument("--checkpoint", default=None) parser.add_argument("--data-root", default=None) parser.add_argument("--text-dir", default=None) parser.add_argument("--label-file", default=None) parser.add_argument("--output-dir", default=None) parser.add_argument("--batch-size", type=int, default=None) parser.add_argument("--max-length", type=int, default=None) parser.add_argument("--num-workers", type=int, default=None) parser.add_argument("--train-ratio", type=float, default=None) parser.add_argument("--val-ratio", type=float, default=None) parser.add_argument("--seed", type=int, default=None) parser.add_argument("--preprocessing-mode", choices=["paper", "strict"], default=None) parser.add_argument("--disable-paper-exact-counts", action="store_true") parser.add_argument("--device", default="auto") parser.add_argument( "--thresholds", type=float, nargs="*", default=[0.05, 0.10, 0.20], help="Thresholds used for disagreement gating.", ) parser.add_argument( "--confidence-gates", type=float, nargs="*", default=[0.70, 0.80], help="Confidence-gate values. Feedback applies only for predictions <= gate.", ) return parser.parse_args() def _load_hfm( args: argparse.Namespace, device: Any ) -> Tuple[Any, Dict[str, Any], Any, Callable[..., Tuple[np.ndarray, np.ndarray]], str]: checkpoint = args.checkpoint or "outputs/hfm/clara_hfm.pt" data_root = args.data_root or "data/HFM" text_dir = args.text_dir or str(Path(data_root) / "text") model, cfg, _ = load_hfm_checkpoint(checkpoint, device) cfg["image_root"] = data_root cfg["text_dir"] = text_dir if args.batch_size is not None: cfg["batch_size"] = args.batch_size if args.num_workers is not None: cfg["num_workers"] = args.num_workers if args.max_length is not None: cfg["max_length"] = args.max_length loader_obj = HFMLoader(cfg["text_dir"], cfg["image_root"]) all_samples = loader_obj.load() train_samples = loader_obj.get_split("train") val_samples = loader_obj.get_split("val") test_samples = loader_obj.get_split("test") clip_processor = CLIPProcessor.from_pretrained(cfg["vision_model_id"]) tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"]) max_length = cfg.get("max_length") if not max_length: max_length = int(estimate_hfm_max_length(all_samples)) pin_memory = bool(cfg.get("pin_memory", True) and device.type == "cuda") _, _, test_loader = create_hfm_dataloaders( train_samples=train_samples, val_samples=val_samples, test_samples=test_samples, clip_processor=clip_processor, tokenizer=tokenizer, batch_size=int(cfg.get("batch_size", 32)), max_length=int(max_length), num_workers=int(cfg.get("num_workers", 0)), pin_memory=pin_memory, weighted_train_sampler=False, ) return model, cfg, test_loader, gather_hfm_variant, "macro" def _load_mvsa_multiple( args: argparse.Namespace, device: Any ) -> Tuple[Any, Dict[str, Any], Any, Callable[..., Tuple[np.ndarray, np.ndarray]], str]: checkpoint = args.checkpoint or "outputs/mvsa_multiple/clara_mvsa_multiple.pt" 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") model, cfg, _ = load_mvsa_checkpoint(checkpoint, device) cfg["text_dir"] = text_dir cfg["label_file"] = label_file if args.batch_size is not None: cfg["batch_size"] = args.batch_size if args.num_workers is not None: cfg["num_workers"] = args.num_workers if args.max_length is not None: cfg["max_length"] = args.max_length if args.train_ratio is not None: cfg["train_ratio"] = args.train_ratio if args.val_ratio is not None: cfg["val_ratio"] = args.val_ratio if args.seed is not None: cfg["seed"] = int(args.seed) if args.preprocessing_mode is not None: cfg["preprocessing_mode"] = args.preprocessing_mode if args.disable_paper_exact_counts: cfg["paper_exact_counts"] = False loader_obj = MVSALoader(cfg["text_dir"], cfg["label_file"]) loader_obj.load( preprocessing_mode=str(cfg.get("preprocessing_mode", "paper")), require_unanimous=bool(cfg.get("require_unanimous", True)), require_cross_agree=bool(cfg.get("require_cross_agree", True)), paper_exact_counts=bool(cfg.get("paper_exact_counts", False)), ) train_samples, val_samples, test_samples = loader_obj.split( train_ratio=float(cfg.get("train_ratio", 0.8)), val_ratio=float(cfg.get("val_ratio", 0.1)), seed=int(cfg.get("seed", 42)), paper_811=bool(str(cfg.get("preprocessing_mode", "paper")).lower() == "paper"), ) clip_processor = CLIPProcessor.from_pretrained(cfg["vision_model_id"]) tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"]) pin_memory = bool(cfg.get("pin_memory", True) and device.type == "cuda") _, _, test_loader = create_mvsa_dataloaders( train_samples=train_samples, val_samples=val_samples, test_samples=test_samples, clip_processor=clip_processor, tokenizer=tokenizer, batch_size=int(cfg.get("batch_size", 48)), max_length=int(cfg.get("max_length", 128)), num_workers=int(cfg.get("num_workers", 0)), pin_memory=pin_memory, persistent_workers=bool(cfg.get("persistent_workers", True)), prefetch_factor=int(cfg.get("prefetch_factor", 2)), use_mixup_negative=False, mixup_alpha=float(cfg.get("mixup_alpha", 0.4)), negative_class_boost=float(cfg.get("negative_class_boost", 12.0)), min_ratio_negative=float(cfg.get("min_ratio_negative", 0.30)), weighted_train_sampler=False, ) return model, cfg, test_loader, gather_mvsa_variant, "weighted" def _evaluate_from_logits( logits: np.ndarray, labels: np.ndarray, f1_average: str, ) -> Tuple[float, float]: pred = logits.argmax(axis=-1) acc = float(accuracy_score(labels, pred)) f1 = float(f1_score(labels, pred, average=f1_average)) return acc, f1 def main() -> None: args = parse_args() device = resolve_device(args.device) if args.dataset == "hfm": model, cfg, test_loader, gather_variant, f1_average = _load_hfm(args, device) output_dir = Path(args.output_dir or "results/hfm") metric_name = "f1_macro" else: model, cfg, test_loader, gather_variant, f1_average = _load_mvsa_multiple(args, device) output_dir = Path(args.output_dir or "results/mvsa_multiple") metric_name = "f1_weighted" settings: List[Dict[str, Any]] = [ { "setting": "Full (baseline)", "variant": "full", "consensus_mode": None, "threshold": None, "confidence_gate": None, }, { "setting": "Alt consensus: abs_prob_diff", "variant": "vf_custom", "consensus_mode": "abs_prob_diff", "threshold": 0.0, "confidence_gate": None, }, { "setting": "Alt consensus: logit_diff", "variant": "vf_custom", "consensus_mode": "logit_diff", "threshold": 0.0, "confidence_gate": None, }, ] for threshold in args.thresholds: settings.append( { "setting": f"Threshold tau={float(threshold):.2f}", "variant": "vf_custom", "consensus_mode": "prob_diff", "threshold": float(threshold), "confidence_gate": None, } ) for gate in args.confidence_gates: settings.append( { "setting": f"Confidence gate={float(gate):.2f}", "variant": "vf_custom", "consensus_mode": "prob_diff", "threshold": 0.0, "confidence_gate": float(gate), } ) rows: List[Dict[str, Any]] = [] for item in settings: if item["variant"] == "full": logits, labels = gather_variant( model=model, loader=test_loader, device=device, variant="full", ) else: logits, labels = gather_variant( model=model, loader=test_loader, device=device, variant="vf_custom", consensus_mode=str(item["consensus_mode"]), threshold=float(item["threshold"]), confidence_gate=( float(item["confidence_gate"]) if item["confidence_gate"] is not None else None ), ) acc, f1 = _evaluate_from_logits(logits=logits, labels=labels, f1_average=f1_average) rows.append( { "setting": item["setting"], "variant": item["variant"], "consensus_mode": item["consensus_mode"], "threshold": item["threshold"], "confidence_gate": item["confidence_gate"], "accuracy": acc, metric_name: f1, } ) output_dir.mkdir(parents=True, exist_ok=True) csv_path = output_dir / "verification_feedback_sensitivity.csv" json_path = output_dir / "verification_feedback_sensitivity.json" with csv_path.open("w", encoding="utf-8", newline="") as f: fieldnames = [ "setting", "variant", "consensus_mode", "threshold", "confidence_gate", "accuracy", metric_name, ] writer = csv.DictWriter(f, fieldnames=fieldnames) writer.writeheader() for row in rows: writer.writerow(row) payload = { "dataset": args.dataset, "checkpoint": str(args.checkpoint) if args.checkpoint else None, "device": str(device), "metric_name": metric_name, "settings": rows, "config_used": cfg, } json_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") print("\nVerification-feedback sensitivity summary:") for row in rows: print( f"- {row['setting']:<34} | Acc={row['accuracy']:.4f} | " f"{metric_name}={row[metric_name]:.4f}" ) print(f"Saved CSV: {csv_path}") print(f"Saved JSON: {json_path}") if __name__ == "__main__": main()