tiny-hinglish-turn-detector / scripts /compare_models.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
10.1 kB
#!/usr/bin/env python3
"""Compare two aligned prediction files with paired group bootstrap intervals."""
from __future__ import annotations
import argparse
import hashlib
import json
import random
import sys
from collections import defaultdict
from pathlib import Path
from typing import Any
REPOSITORY_ROOT = Path(__file__).resolve().parents[1]
SOURCE_ROOT = REPOSITORY_ROOT / "src"
if str(SOURCE_ROOT) not in sys.path:
sys.path.insert(0, str(SOURCE_ROOT))
from turn_detection.training.metrics import ( # noqa: E402
average_precision,
binary_classification_metrics,
roc_auc,
threshold_at_max_fpr,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--baseline", required=True, help="baseline prediction JSONL")
parser.add_argument("--candidate", required=True, help="candidate prediction JSONL")
parser.add_argument("--output", required=True)
parser.add_argument("--baseline-name", default="baseline")
parser.add_argument("--candidate-name", default="candidate")
parser.add_argument("--fpr-budget", type=float, default=0.02)
parser.add_argument("--bootstrap-samples", type=int, default=2_000)
parser.add_argument("--seed", type=int, default=17)
return parser.parse_args()
def _load(path: Path) -> dict[str, dict[str, Any]]:
records: dict[str, dict[str, Any]] = {}
with path.open(encoding="utf-8") as handle:
for line_number, line in enumerate(handle, start=1):
if not line.strip():
continue
try:
row = json.loads(line)
except json.JSONDecodeError as exc:
raise ValueError(f"{path}:{line_number}: invalid JSON") from exc
if not isinstance(row, dict):
raise ValueError(f"{path}:{line_number}: expected an object")
record_id = str(row.get("record_id", ""))
if not record_id:
raise ValueError(f"{path}:{line_number}: record_id is missing")
if record_id in records:
raise ValueError(f"{path}:{line_number}: duplicate record_id {record_id!r}")
label = row.get("label")
probability = row.get("probability")
if label not in (0, 1, False, True):
raise ValueError(f"{path}:{line_number}: label must be binary")
if not isinstance(probability, int | float) or not 0.0 <= probability <= 1.0:
raise ValueError(f"{path}:{line_number}: probability must be in [0, 1]")
records[record_id] = row
if not records:
raise ValueError(f"{path}: no predictions")
return records
def _percentile(values: list[float], quantile: float) -> float:
ordered = sorted(values)
position = (len(ordered) - 1) * quantile
lower = int(position)
upper = min(lower + 1, len(ordered) - 1)
fraction = position - lower
return ordered[lower] * (1.0 - fraction) + ordered[upper] * fraction
def _interval(values: list[float], point: float) -> dict[str, float]:
return {
"estimate": point,
"lower": _percentile(values, 0.025),
"upper": _percentile(values, 0.975),
"bootstrap_fraction_delta_above_zero": sum(value > 0.0 for value in values) / len(values),
}
def _portable_path(path: Path) -> str:
try:
return path.resolve().relative_to(REPOSITORY_ROOT).as_posix()
except ValueError:
return path.name
def _file_evidence(path: Path) -> dict[str, str | int]:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return {
"path": _portable_path(path),
"bytes": path.stat().st_size,
"sha256": digest.hexdigest(),
}
def main() -> int:
args = parse_args()
if not 0.0 <= args.fpr_budget <= 1.0:
raise SystemExit("--fpr-budget must be in [0, 1]")
if args.bootstrap_samples < 1:
raise SystemExit("--bootstrap-samples must be positive")
baseline_path = Path(args.baseline)
candidate_path = Path(args.candidate)
if not baseline_path.is_absolute():
baseline_path = REPOSITORY_ROOT / baseline_path
if not candidate_path.is_absolute():
candidate_path = REPOSITORY_ROOT / candidate_path
baseline = _load(baseline_path)
candidate = _load(candidate_path)
if baseline.keys() != candidate.keys():
only_baseline = len(baseline.keys() - candidate.keys())
only_candidate = len(candidate.keys() - baseline.keys())
raise SystemExit(
"prediction sets are not aligned: "
f"{only_baseline} only in baseline, {only_candidate} only in candidate"
)
record_ids = sorted(baseline)
labels: list[int] = []
baseline_scores: list[float] = []
candidate_scores: list[float] = []
groups: list[str] = []
for record_id in record_ids:
baseline_row = baseline[record_id]
candidate_row = candidate[record_id]
if int(baseline_row["label"]) != int(candidate_row["label"]):
raise SystemExit(f"label mismatch for record {record_id}")
baseline_group = str(baseline_row.get("group_id") or record_id)
candidate_group = str(candidate_row.get("group_id") or record_id)
if baseline_group != candidate_group:
raise SystemExit(f"group mismatch for record {record_id}")
labels.append(int(baseline_row["label"]))
baseline_scores.append(float(baseline_row["probability"]))
candidate_scores.append(float(candidate_row["probability"]))
groups.append(baseline_group)
baseline_point = threshold_at_max_fpr(labels, baseline_scores, args.fpr_budget)
candidate_point = threshold_at_max_fpr(labels, candidate_scores, args.fpr_budget)
baseline_threshold = float(baseline_point["threshold"])
candidate_threshold = float(candidate_point["threshold"])
by_group: dict[str, list[int]] = defaultdict(list)
for index, group in enumerate(groups):
by_group[group].append(index)
group_names = sorted(by_group)
rng = random.Random(args.seed)
bootstrap_deltas: dict[str, list[float]] = defaultdict(list)
for _ in range(args.bootstrap_samples):
sampled_groups = [rng.choice(group_names) for _ in group_names]
indices = [index for group in sampled_groups for index in by_group[group]]
sample_labels = [labels[index] for index in indices]
baseline_sample = [baseline_scores[index] for index in indices]
candidate_sample = [candidate_scores[index] for index in indices]
baseline_auc = roc_auc(sample_labels, baseline_sample)
candidate_auc = roc_auc(sample_labels, candidate_sample)
baseline_ap = average_precision(sample_labels, baseline_sample)
candidate_ap = average_precision(sample_labels, candidate_sample)
if baseline_auc is not None and candidate_auc is not None:
bootstrap_deltas["roc_auc"].append(candidate_auc - baseline_auc)
if baseline_ap is not None and candidate_ap is not None:
bootstrap_deltas["average_precision"].append(candidate_ap - baseline_ap)
baseline_fixed = binary_classification_metrics(
sample_labels, baseline_sample, baseline_threshold
)
candidate_fixed = binary_classification_metrics(
sample_labels, candidate_sample, candidate_threshold
)
for metric in ("recall", "false_positive_rate"):
baseline_value = baseline_fixed[metric]
candidate_value = candidate_fixed[metric]
if baseline_value is not None and candidate_value is not None:
bootstrap_deltas[metric].append(candidate_value - baseline_value)
baseline_metrics = binary_classification_metrics(labels, baseline_scores, baseline_threshold)
candidate_metrics = binary_classification_metrics(labels, candidate_scores, candidate_threshold)
point_deltas = {
"roc_auc": float(candidate_metrics["roc_auc"] - baseline_metrics["roc_auc"]),
"average_precision": float(
candidate_metrics["average_precision"] - baseline_metrics["average_precision"]
),
"recall": float(candidate_metrics["recall"] - baseline_metrics["recall"]),
"false_positive_rate": float(
candidate_metrics["false_positive_rate"] - baseline_metrics["false_positive_rate"]
),
}
comparison = {
"scope": "paired development comparison; not independent-test evidence",
"count": len(labels),
"group_count": len(group_names),
"positive_count": sum(labels),
"negative_count": len(labels) - sum(labels),
"fpr_budget": args.fpr_budget,
"threshold_note": (
"Each threshold was selected on this same development set. Bootstrap samples use "
"those fixed thresholds; intervals do not remove model-selection or calibration bias."
),
"baseline": {
"name": args.baseline_name,
"predictions": _file_evidence(baseline_path),
"metrics": baseline_metrics,
},
"candidate": {
"name": args.candidate_name,
"predictions": _file_evidence(candidate_path),
"metrics": candidate_metrics,
},
"candidate_minus_baseline_95ci": {
metric: _interval(bootstrap_deltas[metric], point)
for metric, point in point_deltas.items()
},
"bootstrap": {
"unit": "audit leakage group",
"samples": args.bootstrap_samples,
"seed": args.seed,
},
}
output = Path(args.output)
if not output.is_absolute():
output = REPOSITORY_ROOT / output
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(
json.dumps(comparison, indent=2, sort_keys=True, allow_nan=False) + "\n",
encoding="utf-8",
)
print(json.dumps(comparison, indent=2, sort_keys=True, allow_nan=False))
return 0
if __name__ == "__main__":
raise SystemExit(main())