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