| |
| """Audit raw smart-turn data and emit a leakage-grouped JSONL manifest.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import os |
| import sys |
| from collections.abc import Iterable, Iterator, Mapping |
| from pathlib import Path |
| from typing import Any |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| SRC_ROOT = PROJECT_ROOT / "src" |
| if str(SRC_ROOT) not in sys.path: |
| sys.path.insert(0, str(SRC_ROOT)) |
|
|
| from turn_detection.data import ( |
| DatasetReadError, |
| GroupingConfig, |
| OptionalDependencyError, |
| audit_records, |
| iter_records, |
| write_json, |
| write_manifest, |
| ) |
|
|
| DEFAULT_DATASET = PROJECT_ROOT / "data" / "raw" / "smart-turn-data-v3.2-train" |
| DEFAULT_MANIFEST = PROJECT_ROOT / "data" / "processed" / "manifest.jsonl" |
| DEFAULT_REPORT = PROJECT_ROOT / "artifacts" / "data_audit.json" |
|
|
|
|
| def _parser() -> argparse.ArgumentParser: |
| parser = argparse.ArgumentParser( |
| description=( |
| "Stream local Parquet/JSON/CSV or a Hugging Face dataset, validate every row, " |
| "hash exact audio, and construct transitive leakage groups." |
| ) |
| ) |
| parser.add_argument( |
| "inputs", |
| nargs="*", |
| help=( |
| "Local file/directory paths, or one Hugging Face dataset id. " |
| f"Defaults to {DEFAULT_DATASET.relative_to(PROJECT_ROOT)}." |
| ), |
| ) |
| parser.add_argument( |
| "--output", type=Path, default=DEFAULT_MANIFEST, help="Output JSONL manifest" |
| ) |
| parser.add_argument( |
| "--report", type=Path, default=DEFAULT_REPORT, help="Output audit report JSON" |
| ) |
| parser.add_argument("--hf-split", default="train", help="Hugging Face split name") |
| parser.add_argument("--revision", help="Pinned Hugging Face dataset revision") |
| parser.add_argument( |
| "--token-env", default="HF_TOKEN", help="Environment variable containing the HF token" |
| ) |
| parser.add_argument( |
| "--env-file", |
| type=Path, |
| default=PROJECT_ROOT / ".env", |
| help="Optional KEY=VALUE file used only when --token-env is unset", |
| ) |
| parser.add_argument( |
| "--batch-size", type=int, default=256, help="Parquet rows per streamed Arrow batch" |
| ) |
| parser.add_argument("--limit", type=int, help="Audit at most N records (useful for smoke runs)") |
| parser.add_argument( |
| "--progress-every", |
| type=int, |
| default=1000, |
| help="Print progress every N input rows; set 0 to disable", |
| ) |
| parser.add_argument( |
| "--no-text-grouping", |
| action="store_true", |
| help="Do not link repeated normalized transcripts/prompts", |
| ) |
| parser.add_argument( |
| "--fail-on-error", |
| action="store_true", |
| help="Return a failure status when any invalid record is found (outputs are still written)", |
| ) |
| return parser |
|
|
|
|
| def _token_from_env(name: str, env_file: Path | None) -> str | None: |
| token = os.environ.get(name) |
| if token or env_file is None or not env_file.is_file(): |
| return token |
| |
| for raw_line in env_file.read_text(encoding="utf-8").splitlines(): |
| line = raw_line.strip() |
| if not line or line.startswith("#") or "=" not in line: |
| continue |
| key, value = line.split("=", 1) |
| if key.strip() in (name, name.casefold()): |
| value = value.strip() |
| if len(value) >= 2 and value[0] == value[-1] and value[0] in ("'", '"'): |
| value = value[1:-1] |
| return value or None |
| return None |
|
|
|
|
| def _progress( |
| records: Iterable[Mapping[str, Any]], |
| *, |
| every: int, |
| ) -> Iterator[Mapping[str, Any]]: |
| for count, record in enumerate(records, start=1): |
| if every > 0 and count % every == 0: |
| print(f"audited input rows: {count:,}", file=sys.stderr, flush=True) |
| yield record |
|
|
|
|
| def main(argv: list[str] | None = None) -> int: |
| args = _parser().parse_args(argv) |
| if args.batch_size <= 0: |
| raise SystemExit("--batch-size must be positive") |
| if args.limit is not None and args.limit < 0: |
| raise SystemExit("--limit cannot be negative") |
| if args.progress_every < 0: |
| raise SystemExit("--progress-every cannot be negative") |
|
|
| inputs = args.inputs or [str(DEFAULT_DATASET)] |
| source: str | Path | list[str | Path] |
| source = inputs[0] if len(inputs) == 1 else inputs |
| token = _token_from_env(args.token_env, args.env_file) |
| try: |
| records = iter_records( |
| source, |
| split=args.hf_split, |
| revision=args.revision, |
| token=token, |
| batch_size=args.batch_size, |
| limit=args.limit, |
| ) |
| manifest, report = audit_records( |
| _progress(records, every=args.progress_every), |
| grouping_config=GroupingConfig(include_text=not args.no_text_grouping), |
| ) |
| write_manifest(args.output, manifest) |
| write_json(args.report, report) |
| except (FileNotFoundError, DatasetReadError, OptionalDependencyError, ValueError) as exc: |
| print(f"audit failed: {exc}", file=sys.stderr) |
| return 2 |
|
|
| invalid = int(report["records"]["invalid"]) |
| print( |
| f"wrote {len(manifest):,} rows to {args.output} " |
| f"({invalid:,} invalid); report: {args.report}", |
| file=sys.stderr, |
| ) |
| return 1 if args.fail_on_error and invalid else 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|