| """
|
| bert/calibrate.py
|
| Threshold calibration for the entailment rejection gate.
|
|
|
| A raw softmax score is NOT a true probability.
|
| After training, run a hard-negative validation set (entity/date swapped claims)
|
| and plot the Precision-Recall curve to find the threshold T where:
|
| FPR_entailment == 0 (we never pass a contradiction as entailment)
|
|
|
| In a verification engine: precision > recall.
|
| Better to drop a true claim than to cite a hallucinated one.
|
| """
|
|
|
| from __future__ import annotations
|
|
|
| import argparse
|
| import json
|
| from pathlib import Path
|
| from typing import List, Tuple
|
|
|
| import numpy as np
|
| import torch
|
| from sklearn.metrics import precision_recall_curve, average_precision_score
|
| from tqdm import tqdm
|
|
|
| from bert.dataset import (
|
| CrossEncoderDataset,
|
| EntailmentExample,
|
| build_combined_dataset,
|
| generate_hard_negatives,
|
| make_dataloader,
|
| )
|
| from bert.model import BertCrossEncoderVerifier, CrossEncoderConfig, build_model, load_tokenizer
|
|
|
|
|
| @torch.no_grad()
|
| def collect_scores(
|
| model: BertCrossEncoderVerifier,
|
| loader: torch.utils.data.DataLoader,
|
| device: torch.device,
|
| ) -> Tuple[np.ndarray, np.ndarray]:
|
| """Return (entailment_scores, true_labels) over the full dataset."""
|
| model.eval()
|
| all_scores: List[float] = []
|
| all_labels: List[int] = []
|
|
|
| for batch in tqdm(loader, desc="calibrate"):
|
| input_ids = batch["input_ids"].to(device)
|
| attention_mask = batch["attention_mask"].to(device)
|
| token_type_ids = batch.get("token_type_ids")
|
| if token_type_ids is not None:
|
| token_type_ids = token_type_ids.to(device)
|
| labels = batch["label"]
|
|
|
| scores = model.predict_entailment_score(input_ids, attention_mask, token_type_ids)
|
| all_scores.extend(scores.cpu().tolist())
|
| all_labels.extend(labels.tolist())
|
|
|
| return np.array(all_scores), np.array(all_labels)
|
|
|
|
|
| def calibrate_threshold(
|
| scores: np.ndarray,
|
| labels: np.ndarray,
|
| target_fpr: float = 0.0,
|
| ) -> Tuple[float, dict]:
|
| """
|
| Find the minimum threshold T such that no Contradiction (label=0)
|
| is classified as Entailment (score >= T).
|
|
|
| target_fpr=0.0 means zero false positive rate for Entailment.
|
| Returns (threshold, metrics_at_threshold).
|
| """
|
|
|
| binary_labels = (labels == 2).astype(int)
|
|
|
| precision, recall, thresholds = precision_recall_curve(binary_labels, scores)
|
| ap = average_precision_score(binary_labels, scores)
|
|
|
|
|
|
|
|
|
|
|
| best_threshold = 1.0
|
| non_entailment = scores[labels != 2]
|
| for t in sorted(set(thresholds)):
|
| fpr = (non_entailment >= t).mean() if len(non_entailment) > 0 else 0.0
|
| if fpr <= target_fpr:
|
| best_threshold = float(t)
|
| break
|
|
|
|
|
| preds = (scores >= best_threshold).astype(int)
|
| tp = ((preds == 1) & (binary_labels == 1)).sum()
|
| fp = ((preds == 1) & (binary_labels == 0)).sum()
|
| fn = ((preds == 0) & (binary_labels == 1)).sum()
|
| tn = ((preds == 0) & (binary_labels == 0)).sum()
|
|
|
| precision_at_t = tp / (tp + fp) if (tp + fp) > 0 else 0.0
|
| recall_at_t = tp / (tp + fn) if (tp + fn) > 0 else 0.0
|
| f1_at_t = (
|
| 2 * precision_at_t * recall_at_t / (precision_at_t + recall_at_t)
|
| if (precision_at_t + recall_at_t) > 0 else 0.0
|
| )
|
| actual_fpr = fp / (fp + tn) if (fp + tn) > 0 else 0.0
|
|
|
| metrics = {
|
| "threshold": best_threshold,
|
| "average_precision": float(ap),
|
| "precision": float(precision_at_t),
|
| "recall": float(recall_at_t),
|
| "f1": float(f1_at_t),
|
| "fpr": float(actual_fpr),
|
| "tp": int(tp), "fp": int(fp), "fn": int(fn), "tn": int(tn),
|
| }
|
| return best_threshold, metrics
|
|
|
|
|
| def run_calibration(
|
| checkpoint_path: Path,
|
| data_dir: Path,
|
| output_path: Path,
|
| backbone: str = "microsoft/deberta-v3-base",
|
| max_length: int = 512,
|
| batch_size: int = 64,
|
| n_hard_negatives: int = 2000,
|
| seed: int = 42,
|
| ) -> None:
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| tokenizer = load_tokenizer(backbone)
|
| config = CrossEncoderConfig(backbone=backbone, max_length=max_length)
|
| model = build_model(config).to(device)
|
| state = torch.load(checkpoint_path, map_location=device)
|
| if "model_state_dict" in state:
|
| state = state["model_state_dict"]
|
| model.load_state_dict(state)
|
| print(f"[calibrate] loaded {checkpoint_path}")
|
|
|
|
|
| val_ds = build_combined_dataset(data_dir, tokenizer, "dev", max_length, seed)
|
| hard_negs = generate_hard_negatives(val_ds.examples, n=n_hard_negatives, seed=seed)
|
| combined = CrossEncoderDataset(
|
| val_ds.examples + hard_negs, tokenizer, max_length
|
| )
|
| loader = make_dataloader(combined, batch_size, shuffle=False, num_workers=2)
|
| print(f"[calibrate] {len(combined)} examples ({n_hard_negatives} hard negatives)")
|
|
|
| scores, labels = collect_scores(model, loader, device)
|
|
|
| threshold, metrics = calibrate_threshold(scores, labels, target_fpr=0.0)
|
| print(f"\n[calibrate] Results (target FPR=0.0):")
|
| for k, v in metrics.items():
|
| print(f" {k}: {v}")
|
|
|
| output_path.parent.mkdir(parents=True, exist_ok=True)
|
| with output_path.open("w") as f:
|
| json.dump(metrics, f, indent=2)
|
| print(f"\n[calibrate] threshold config saved → {output_path}")
|
| print(f" Use threshold={threshold:.4f} in the inference daemon.")
|
|
|
|
|
| if __name__ == "__main__":
|
| parser = argparse.ArgumentParser()
|
| parser.add_argument("--checkpoint", type=Path, required=True)
|
| parser.add_argument("--data_dir", type=Path, required=True)
|
| parser.add_argument("--output", type=Path, default=Path("config/threshold.json"))
|
| parser.add_argument("--backbone", type=str, default="microsoft/deberta-v3-base")
|
| parser.add_argument("--max_length", type=int, default=512)
|
| parser.add_argument("--batch_size", type=int, default=64)
|
| args = parser.parse_args()
|
| run_calibration(**vars(args))
|
|
|