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