Download scripts/analyze_failures.py from suvradeepp/tiny-hinglish-turn-detector: direct link, hf CLI and curl.
- Browser
- Download file 6.15 kB
-
https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/scripts/analyze_failures.py
- Command line
-
hf download hf://suvradeepp/tiny-hinglish-turn-detector/scripts/analyze_failures.py
-
curl -L -o analyze_failures.py https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/scripts/analyze_failures.py
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()) | |