#!/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: .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, "")) 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())