tiny-hinglish-turn-detector / scripts /analyze_failures.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
6.15 kB
#!/usr/bin/env python3
"""Create a privacy-safe failure queue from development predictions."""
from __future__ import annotations
import argparse
import hashlib
import json
from collections import defaultdict
from pathlib import Path
from typing import Any
ROOT = Path(__file__).resolve().parents[1]
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--metrics", required=True)
parser.add_argument("--predictions", help="default: <metrics stem>.predictions.jsonl")
parser.add_argument("--output", required=True)
parser.add_argument("--examples-per-type", type=int, default=20)
parser.add_argument("--min-slice-count", type=int, default=5)
return parser.parse_args()
def _resolve(value: str) -> Path:
path = Path(value)
return path if path.is_absolute() else ROOT / path
def _case_id(record_id: object) -> str:
return "case_" + hashlib.sha256(str(record_id).encode()).hexdigest()[:12]
def _load_predictions(path: Path) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
seen: set[str] = set()
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 or record_id in seen:
raise ValueError(f"{path}:{line_number}: missing or duplicate record_id")
seen.add(record_id)
rows.append(row)
if not rows:
raise ValueError("prediction file is empty")
return rows
def _review_case(row: dict[str, Any], prediction: int) -> dict[str, Any]:
return {
"case_id": _case_id(row["record_id"]),
"target": "END" if int(row["label"]) else "HOLD",
"predicted": "END" if prediction else "HOLD",
"p_end": float(row["probability"]),
"language": row.get("language"),
"dataset": row.get("dataset"),
"synthetic": row.get("synthetic"),
"filler_type": row.get("filler_type"),
"duration_bin": row.get("duration_bin"),
"review_note": "Listen under authorized local access; do not export audio or transcript.",
}
def main() -> int:
args = parse_args()
if args.examples_per_type < 1 or args.min_slice_count < 1:
raise SystemExit("example and slice counts must be positive")
metrics_path = _resolve(args.metrics)
try:
metrics = json.loads(metrics_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise SystemExit(f"invalid metrics JSON: {metrics_path}") from exc
threshold = metrics.get("threshold")
if not isinstance(threshold, int | float) or not 0.0 <= threshold <= 1.0:
raise SystemExit("metrics JSON has no valid threshold")
predictions_path = (
_resolve(args.predictions)
if args.predictions
else metrics_path.with_name(metrics_path.stem + ".predictions.jsonl")
)
rows = _load_predictions(predictions_path)
false_interruptions: list[dict[str, Any]] = []
missed_ends: list[dict[str, Any]] = []
slices: dict[str, dict[str, dict[str, int]]] = {
dimension: defaultdict(lambda: {"count": 0, "false_interruptions": 0, "missed_ends": 0})
for dimension in ("language", "dataset", "synthetic", "filler_type", "duration_bin")
}
for row in rows:
label = int(row["label"])
probability = float(row["probability"])
if label not in (0, 1) or not 0.0 <= probability <= 1.0:
raise SystemExit("predictions contain invalid labels or probabilities")
prediction = int(probability >= threshold)
is_false_interruption = prediction == 1 and label == 0
is_missed_end = prediction == 0 and label == 1
if is_false_interruption:
false_interruptions.append(_review_case(row, prediction))
elif is_missed_end:
missed_ends.append(_review_case(row, prediction))
for dimension, values in slices.items():
value = str(row.get(dimension, "<missing>"))
values[value]["count"] += 1
values[value]["false_interruptions"] += int(is_false_interruption)
values[value]["missed_ends"] += int(is_missed_end)
false_interruptions.sort(key=lambda row: float(row["p_end"]), reverse=True)
missed_ends.sort(key=lambda row: float(row["p_end"]))
filtered_slices = {
dimension: {
value: counts
for value, counts in sorted(values.items())
if counts["count"] >= args.min_slice_count
}
for dimension, values in slices.items()
}
report = {
"scope": metrics.get("data_scope"),
"split": metrics.get("split"),
"development_only": metrics.get("development_only"),
"threshold": threshold,
"privacy": (
"Case IDs are one-way hashes. No audio, transcript, raw record ID, or source path "
"is included. Hypotheses require authorized local listening."
),
"counts": {
"examples": len(rows),
"false_interruptions": len(false_interruptions),
"missed_ends": len(missed_ends),
},
"highest_confidence_false_interruptions": false_interruptions[: args.examples_per_type],
"highest_confidence_missed_ends": missed_ends[: args.examples_per_type],
"failure_counts_by_slice": filtered_slices,
}
output = _resolve(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(
json.dumps(report, indent=2, sort_keys=True, allow_nan=False) + "\n",
encoding="utf-8",
)
print(json.dumps(report["counts"], indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())