| |
| """Evaluate a checkpoint with product-correct metrics and subgroup slices.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import json |
| import os |
| import sys |
| 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)) |
|
|
| OFFICIAL_TEST_DATASET_ID = "pipecat-ai/smart-turn-data-v3.2-test" |
| OFFICIAL_TEST_REVISION = "0500378e8ed6d38e37b016e24d261e8e6c6a6859" |
| OFFICIAL_TEST_EXPECTED_ROWS = 31_527 |
| OFFICIAL_TEST_LOCAL_DIR = REPOSITORY_ROOT / "data/raw/smart-turn-data-v3.2-test" |
| OFFICIAL_TEST_HF_SOURCES = { |
| OFFICIAL_TEST_DATASET_ID, |
| f"hf://datasets/{OFFICIAL_TEST_DATASET_ID}", |
| } |
| OFFICIAL_TEST_SHARDS = { |
| "data/train-00000-of-00010.parquet": ( |
| 486_502_678, |
| "a87c75806b814ee7379998b6f9dc65a6433c01bfec2875e62c5d1ccd2b37257a", |
| ), |
| "data/train-00001-of-00010.parquet": ( |
| 489_429_742, |
| "c408bd3b31cc3cb907280fa5d3186f0f5ba08c6beb84532a3828d685b967b7d8", |
| ), |
| "data/train-00002-of-00010.parquet": ( |
| 479_920_042, |
| "2b50ff3346f8aecc6c4b0b706b593c5b3b174b73ce84c795b515b68e6abb3788", |
| ), |
| "data/train-00003-of-00010.parquet": ( |
| 486_565_988, |
| "0afd86b7d1cdf03ffb804a00278fb76477d29ef008b71ad1d7a746bb1b25850c", |
| ), |
| "data/train-00004-of-00010.parquet": ( |
| 477_831_330, |
| "4c600774512010880f72dd86cc8abe5d46df594615e472351bd270cc94e1ff66", |
| ), |
| "data/train-00005-of-00010.parquet": ( |
| 495_015_303, |
| "b4537a3b96498481b98c5d60b8d84ad05109ae0acbeafc18759b15ff1d0d9335", |
| ), |
| "data/train-00006-of-00010.parquet": ( |
| 481_054_312, |
| "ef0eb0085b55e05fc5594f035c8011afb78d039e1b3e50c579a96e9295d9acec", |
| ), |
| "data/train-00007-of-00010.parquet": ( |
| 478_266_840, |
| "eddd1db2f95fff2f08e18ca1fe73fe0c5d9eb8030bea97da2d109297a0158b67", |
| ), |
| "data/train-00008-of-00010.parquet": ( |
| 479_859_534, |
| "96314dc8bb77515a5d1d02e8cb7c1410da54169d33d6f73ca2f87738ec1269f0", |
| ), |
| "data/train-00009-of-00010.parquet": ( |
| 483_305_860, |
| "769283c79bea4ae6eebdfe7d09fe481f8a154c53b5e3c29fe2c71d7b55dfc862", |
| ), |
| } |
|
|
|
|
| def _sha256(path: Path) -> str: |
| digest = hashlib.sha256() |
| with path.open("rb") as handle: |
| for block in iter(lambda: handle.read(1024 * 1024), b""): |
| digest.update(block) |
| return digest.hexdigest() |
|
|
|
|
| def _portable_path(path: Path) -> str: |
| try: |
| return path.resolve().relative_to(REPOSITORY_ROOT).as_posix() |
| except ValueError: |
| return path.name |
|
|
|
|
| def _evaluation_source_evidence( |
| source: str | None, |
| *, |
| dataset_id: str | None, |
| revision: str | None, |
| ) -> dict[str, Any]: |
| local = _local_source_path(source) |
| if local is not None and local.is_file() and not local.is_symlink(): |
| return { |
| "kind": "file", |
| "path": _portable_path(local), |
| "bytes": local.stat().st_size, |
| "sha256": _sha256(local), |
| "dataset_id": dataset_id, |
| "revision": revision, |
| } |
| if local is not None: |
| return { |
| "kind": "directory", |
| "path": _portable_path(local), |
| "dataset_id": dataset_id, |
| "revision": revision, |
| } |
| return { |
| "kind": "huggingface", |
| "identifier": dataset_id or source, |
| "revision": revision, |
| } |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--checkpoint", required=True) |
| parser.add_argument("--source", help="manifest, local dataset, or HF dataset ID") |
| parser.add_argument( |
| "--dataset-id", |
| help=( |
| "source provenance identity; required and checked for a local official-test snapshot" |
| ), |
| ) |
| parser.add_argument("--revision", help="immutable dataset revision") |
| parser.add_argument("--split", default="test") |
| parser.add_argument("--source-root", help="base directory for manifest source_file entries") |
| parser.add_argument("--output", default="artifacts/evaluation/metrics.json") |
| parser.add_argument( |
| "--predictions-output", |
| help="optional JSONL path (default: beside metrics); keep raw record IDs out of reports/", |
| ) |
| parser.add_argument("--batch-size", type=int, default=32) |
| parser.add_argument("--num-workers", type=int, default=0) |
| parser.add_argument("--max-examples", type=int) |
| parser.add_argument("--threshold", type=float, help="default: calibrated checkpoint threshold") |
| parser.add_argument("--min-slice-count", type=int, default=25) |
| parser.add_argument("--bootstrap-samples", type=int, default=1_000) |
| parser.add_argument("--device", default="auto") |
| parser.add_argument("--smoke-test", action="store_true") |
| parser.add_argument( |
| "--allow-sealed-test", |
| action="store_true", |
| help="explicitly permit the official smart-turn test after model/threshold freeze", |
| ) |
| parser.add_argument( |
| "--frozen-manifest", |
| help="hash-bound manifest produced by scripts/freeze_candidate.py", |
| ) |
| return parser.parse_args() |
|
|
|
|
| def _device(torch: Any, requested: str) -> Any: |
| if requested != "auto": |
| return torch.device(requested) |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
|
|
| def _hf_dataset_id_from_source(source: str | None) -> str | None: |
| """Return an exact canonical HF identity; never infer one from a path basename.""" |
|
|
| if not source: |
| return None |
| normalized = source.rstrip("/") |
| if normalized in OFFICIAL_TEST_HF_SOURCES: |
| return OFFICIAL_TEST_DATASET_ID |
| return None |
|
|
|
|
| def _is_official_test_source( |
| source: str | None, |
| dataset_id: str | None = None, |
| ) -> bool: |
| inferred = _hf_dataset_id_from_source(source) |
| if dataset_id is not None and inferred is not None and dataset_id != inferred: |
| raise ValueError("--dataset-id conflicts with the exact Hugging Face source identity") |
| return (dataset_id or inferred) == OFFICIAL_TEST_DATASET_ID |
|
|
|
|
| def _local_source_path(source: str | None) -> Path | None: |
| if not source: |
| return None |
| candidate = Path(source).expanduser() |
| if not candidate.is_absolute(): |
| candidate = REPOSITORY_ROOT / candidate |
| return candidate.resolve() if candidate.exists() else None |
|
|
|
|
| def _validate_official_test_request(args: argparse.Namespace) -> bool: |
| """Fail closed when any official-test control is incomplete or contradictory.""" |
|
|
| try: |
| official = _is_official_test_source(args.source, args.dataset_id) |
| except ValueError as exc: |
| raise SystemExit(str(exc)) from exc |
| local_source = _local_source_path(args.source) |
| canonical_local_source = ( |
| local_source is not None and local_source == OFFICIAL_TEST_LOCAL_DIR.resolve() |
| ) |
| official_controls_used = bool(args.allow_sealed_test or args.frozen_manifest) |
| if canonical_local_source and args.dataset_id != OFFICIAL_TEST_DATASET_ID: |
| raise SystemExit( |
| f"the local official-test snapshot requires --dataset-id {OFFICIAL_TEST_DATASET_ID}" |
| ) |
| if official_controls_used and not official: |
| raise SystemExit( |
| "official-test controls require the exact dataset identity " |
| f"--dataset-id {OFFICIAL_TEST_DATASET_ID} (or that exact HF source)" |
| ) |
| if not official: |
| return False |
| if args.smoke_test: |
| raise SystemExit("official-test provenance cannot be combined with --smoke-test") |
| if args.revision != OFFICIAL_TEST_REVISION: |
| raise SystemExit( |
| "official test requires the exact pinned revision: " + OFFICIAL_TEST_REVISION |
| ) |
| if not args.allow_sealed_test or not args.frozen_manifest: |
| raise SystemExit("official test requires both --allow-sealed-test and --frozen-manifest") |
| if args.max_examples is not None: |
| raise SystemExit("--max-examples is forbidden for official-test evaluation") |
| if args.split not in {"test", "train"}: |
| raise SystemExit( |
| "official-test --split must be semantic 'test' or physical Hugging Face 'train'" |
| ) |
| return True |
|
|
|
|
| def _verify_local_official_snapshot(source: Path) -> int: |
| """Verify the immutable local snapshot before loading any official-test rows.""" |
|
|
| if not source.is_dir() or source.is_symlink(): |
| raise ValueError("local official-test source must be a non-symlink directory") |
| actual = { |
| path.relative_to(source).as_posix() for path in source.rglob("*.parquet") if path.is_file() |
| } |
| expected = set(OFFICIAL_TEST_SHARDS) |
| if actual != expected: |
| missing = sorted(expected - actual) |
| unexpected = sorted(actual - expected) |
| raise ValueError( |
| f"official-test shard inventory mismatch: missing={missing}, unexpected={unexpected}" |
| ) |
| try: |
| import pyarrow.parquet as parquet |
| except ImportError as exc: |
| raise ValueError("official-test snapshot verification requires pyarrow") from exc |
| rows = 0 |
| for relative, (expected_bytes, expected_sha256) in OFFICIAL_TEST_SHARDS.items(): |
| path = source / relative |
| if path.is_symlink() or not path.is_file(): |
| raise ValueError(f"official-test shard is not a regular file: {relative}") |
| if path.stat().st_size != expected_bytes: |
| raise ValueError(f"official-test shard size mismatch: {relative}") |
| if _sha256(path) != expected_sha256: |
| raise ValueError(f"official-test shard SHA-256 mismatch: {relative}") |
| try: |
| rows += int(parquet.ParquetFile(path).metadata.num_rows) |
| except Exception as exc: |
| raise ValueError(f"cannot read official-test Parquet metadata: {relative}") from exc |
| if rows != OFFICIAL_TEST_EXPECTED_ROWS: |
| raise ValueError( |
| "official-test row-count mismatch: " |
| f"expected {OFFICIAL_TEST_EXPECTED_ROWS}, found {rows}" |
| ) |
| return rows |
|
|
|
|
| def main() -> int: |
| args = parse_args() |
| official_test = _validate_official_test_request(args) |
| |
| |
| |
| source_split = "train" if official_test else args.split |
| report_split = "test" if official_test else args.split |
| try: |
| import torch |
| except ImportError as exc: |
| raise SystemExit("Evaluation requires PyTorch") from exc |
|
|
| from turn_detection.models import LogMelConfig, LogMelFrontend, load_model_checkpoint |
| from turn_detection.provenance import verify_freeze_manifest |
| from turn_detection.training.datasets import ( |
| AudioFeatureCollator, |
| SyntheticFeatureDataset, |
| build_record_dataloader, |
| ) |
| from turn_detection.training.metrics import ( |
| binary_classification_metrics, |
| grouped_bootstrap_interval, |
| metrics_at_fpr_budgets, |
| operational_metrics, |
| reliability_bins, |
| sliced_metrics, |
| ) |
|
|
| checkpoint_path = Path(args.checkpoint) |
| if not checkpoint_path.is_absolute(): |
| checkpoint_path = REPOSITORY_ROOT / checkpoint_path |
| device = _device(torch, args.device) |
| model, checkpoint = load_model_checkpoint(checkpoint_path, map_location=device) |
| model.to(device).eval() |
| metadata = dict(checkpoint.get("metadata", {})) |
| run_metadata = metadata.get("run_metadata", {}) |
| if not isinstance(run_metadata, dict): |
| run_metadata = {} |
| smoke_checkpoint = bool(metadata.get("smoke_test", False)) |
| training_status = str( |
| run_metadata.get("status", "smoke" if smoke_checkpoint else "development") |
| ) |
| development_only = smoke_checkpoint or training_status.lower() != "final" |
| feature_config = LogMelConfig.from_mapping(metadata.get("feature_config", {})) |
| frontend = LogMelFrontend(feature_config) |
| max_seconds = float(metadata.get("max_seconds", 8.0)) |
| threshold = float( |
| args.threshold if args.threshold is not None else checkpoint.get("threshold", 0.5) |
| ) |
| dataset_revision = args.revision |
| freeze_manifest_sha256: str | None = None |
| if official_test: |
| if not args.allow_sealed_test or not args.frozen_manifest: |
| raise SystemExit( |
| "official test requires both --allow-sealed-test and --frozen-manifest" |
| ) |
| frozen_path = Path(args.frozen_manifest) |
| if not frozen_path.is_absolute(): |
| frozen_path = REPOSITORY_ROOT / frozen_path |
| try: |
| frozen = verify_freeze_manifest( |
| frozen_path, |
| REPOSITORY_ROOT, |
| checkpoint_path=checkpoint_path, |
| threshold=threshold, |
| ) |
| except ValueError as exc: |
| raise SystemExit(f"official-test freeze validation failed: {exc}") from exc |
| frozen_dataset = frozen["official_test"] |
| if frozen_dataset.get("dataset_id") != OFFICIAL_TEST_DATASET_ID: |
| raise SystemExit("official-test freeze targets the wrong dataset identity") |
| if frozen_dataset.get("revision") != args.revision: |
| raise SystemExit("official-test revision differs from frozen manifest") |
| freeze_manifest_sha256 = _sha256(frozen_path) |
| local_source = _local_source_path(args.source) |
| if local_source is not None: |
| try: |
| _verify_local_official_snapshot(local_source) |
| except ValueError as exc: |
| raise SystemExit(f"official-test snapshot validation failed: {exc}") from exc |
| elif _hf_dataset_id_from_source(args.source) is None: |
| raise SystemExit("declared local official-test source does not exist") |
|
|
| if args.smoke_test: |
| from torch.utils.data import DataLoader |
|
|
| dataset = SyntheticFeatureDataset(32, feature_config.n_mels, 96, seed=18) |
| loader = DataLoader( |
| dataset, |
| batch_size=args.batch_size, |
| collate_fn=AudioFeatureCollator(frontend, max_seconds=1.0), |
| ) |
| else: |
| if not args.source: |
| raise SystemExit("--source is required unless --smoke-test is used") |
| loader = build_record_dataloader( |
| args.source, |
| split=source_split, |
| frontend=frontend, |
| batch_size=args.batch_size, |
| max_seconds=max_seconds, |
| shuffle=False, |
| num_workers=args.num_workers, |
| token=os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN"), |
| revision=dataset_revision, |
| max_examples=args.max_examples, |
| source_root=args.source_root, |
| ) |
| labels: list[int] = [] |
| probabilities: list[float] = [] |
| records: list[dict[str, Any]] = [] |
| languages: list[Any] = [] |
| datasets: list[Any] = [] |
| synthetic_values: list[Any] = [] |
| filler_types: list[str] = [] |
| duration_bins: list[str] = [] |
| turn_ids: list[str] = [] |
| turn_id_observed: list[bool] = [] |
| group_ids: list[str] = [] |
| durations: list[float] = [] |
|
|
| with torch.inference_mode(): |
| for batch in loader: |
| features = batch["log_mel"].to(device) |
| mask = batch["attention_mask"].to(device) |
| output = model(features, mask) |
| batch_probabilities = torch.sigmoid(output.endpoint_logits).cpu().tolist() |
| batch_labels = [int(value) for value in batch["endpoint"].tolist()] |
| for index, (target, probability) in enumerate( |
| zip(batch_labels, batch_probabilities, strict=True) |
| ): |
| midfiller = float(batch["midfiller"][index]) |
| endfiller = float(batch["endfiller"][index]) |
| if midfiller == 1.0: |
| filler_type = "midfiller" |
| elif endfiller == 1.0: |
| filler_type = "endfiller" |
| elif midfiller == 0.0 and endfiller == 0.0: |
| filler_type = "no_filler" |
| else: |
| filler_type = "unknown" |
| duration = batch["duration_seconds"][index] |
| if duration is None: |
| duration_bin = "unknown" |
| elif float(duration) < 2.0: |
| duration_bin = "<2s" |
| elif float(duration) < 4.0: |
| duration_bin = "2-4s" |
| elif float(duration) < 8.0: |
| duration_bin = "4-8s" |
| else: |
| duration_bin = ">=8s" |
| record = { |
| "record_id": batch["record_id"][index], |
| "turn_id": batch["turn_id"][index], |
| "group_id": batch["group_id"][index], |
| "label": target, |
| "probability": float(probability), |
| "prediction": int(probability >= threshold), |
| "language": batch["language"][index], |
| "dataset": batch["dataset"][index], |
| "synthetic": batch["synthetic"][index], |
| "filler_type": filler_type, |
| "duration_bin": duration_bin, |
| } |
| records.append(record) |
| filler_types.append(filler_type) |
| duration_bins.append(duration_bin) |
| labels.extend(batch_labels) |
| probabilities.extend(float(value) for value in batch_probabilities) |
| languages.extend(batch["language"]) |
| datasets.extend(batch["dataset"]) |
| synthetic_values.extend(batch["synthetic"]) |
| turn_ids.extend(batch["turn_id"]) |
| turn_id_observed.extend(bool(value) for value in batch["turn_id_observed"]) |
| group_ids.extend(batch["group_id"]) |
| durations.extend( |
| float(value) for value in batch["duration_seconds"] if value is not None |
| ) |
|
|
| if not labels: |
| raise SystemExit("evaluation source produced no examples") |
| if official_test and len(labels) != OFFICIAL_TEST_EXPECTED_ROWS: |
| raise SystemExit( |
| "official-test evaluation was incomplete: " |
| f"expected {OFFICIAL_TEST_EXPECTED_ROWS} examples, evaluated {len(labels)}" |
| ) |
| total_audio_seconds = sum(durations) if len(durations) == len(labels) else None |
| slice_report = sliced_metrics( |
| labels, |
| probabilities, |
| { |
| "language": languages, |
| "dataset": datasets, |
| "synthetic": synthetic_values, |
| "filler_type": filler_types, |
| "duration_bin": duration_bins, |
| }, |
| threshold=threshold, |
| min_count=args.min_slice_count, |
| ) |
|
|
| def worst_group(metric: str, maximize: bool) -> dict[str, Any] | None: |
| candidates: list[dict[str, Any]] = [] |
| for dimension, values in slice_report.items(): |
| for value, metrics in values.items(): |
| score = metrics.get(metric) |
| if score is not None: |
| candidates.append( |
| { |
| "dimension": dimension, |
| "value": value, |
| "count": metrics["count"], |
| metric: score, |
| } |
| ) |
| if not candidates: |
| return None |
| return sorted(candidates, key=lambda item: float(item[metric]), reverse=maximize)[0] |
|
|
| has_observed_turn_ids = bool(turn_id_observed) and all(turn_id_observed) |
| operational = operational_metrics( |
| labels, |
| probabilities, |
| threshold, |
| turn_ids=turn_ids if has_observed_turn_ids else None, |
| total_audio_seconds=total_audio_seconds, |
| ) |
| operational["sequence_metrics_available"] = has_observed_turn_ids |
| operational["scope"] = ( |
| "observed turn/conversation sequences" |
| if has_observed_turn_ids |
| else "independent labeled clips; per-hour rate is a clip-normalized proxy, not an " |
| "online conversation measurement" |
| ) |
| if not has_observed_turn_ids: |
| operational["sequence_metrics_unavailable_reason"] = ( |
| "source records contain no genuine turn_id or conversation_id" |
| ) |
|
|
| report = { |
| "checkpoint": _portable_path(checkpoint_path), |
| "checkpoint_sha256": _sha256(checkpoint_path), |
| "dataset_revision": dataset_revision, |
| "official_test": official_test, |
| "freeze_manifest_sha256": freeze_manifest_sha256, |
| "split": report_split, |
| "source_split": source_split, |
| "evaluation_source": _evaluation_source_evidence( |
| args.source, |
| dataset_id=args.dataset_id, |
| revision=args.revision or metadata.get("data_revision"), |
| ), |
| "development_only": development_only, |
| "training_status": training_status, |
| "data_scope": metadata.get("data_scope"), |
| "data_revision": metadata.get("data_revision"), |
| "threshold": threshold, |
| "metrics": binary_classification_metrics(labels, probabilities, threshold), |
| "operating_points": metrics_at_fpr_budgets(labels, probabilities), |
| "operating_points_note": ( |
| "Label-dependent threshold sweep for curve analysis only. The deployed result is " |
| "`metrics` at the checkpoint's frozen validation-selected threshold." |
| ), |
| "operational": operational, |
| "slices": slice_report, |
| "worst_groups": { |
| "highest_false_positive_rate": worst_group("false_positive_rate", True), |
| "lowest_recall": worst_group("recall", False), |
| }, |
| "grouped_bootstrap_95ci": { |
| "false_positive_rate": grouped_bootstrap_interval( |
| labels, |
| probabilities, |
| group_ids, |
| threshold, |
| metric="false_positive_rate", |
| samples=args.bootstrap_samples, |
| seed=17, |
| ), |
| "recall": grouped_bootstrap_interval( |
| labels, |
| probabilities, |
| group_ids, |
| threshold, |
| metric="recall", |
| samples=args.bootstrap_samples, |
| seed=17, |
| ), |
| }, |
| "reliability": reliability_bins(labels, probabilities), |
| } |
|
|
| output_path = Path(args.output) |
| if not output_path.is_absolute(): |
| output_path = REPOSITORY_ROOT / output_path |
| output_path.parent.mkdir(parents=True, exist_ok=True) |
| output_path.write_text( |
| json.dumps(report, indent=2, sort_keys=True, allow_nan=False), encoding="utf-8" |
| ) |
| prediction_path = ( |
| Path(args.predictions_output) |
| if args.predictions_output |
| else output_path.with_name(output_path.stem + ".predictions.jsonl") |
| ) |
| if not prediction_path.is_absolute(): |
| prediction_path = REPOSITORY_ROOT / prediction_path |
| prediction_path.parent.mkdir(parents=True, exist_ok=True) |
| with prediction_path.open("w", encoding="utf-8") as handle: |
| for record in records: |
| handle.write(json.dumps(record, ensure_ascii=False, allow_nan=False) + "\n") |
| print(json.dumps({"metrics": report["metrics"], "output": str(output_path)}, indent=2)) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|