#!/usr/bin/env python3 """Create and validate deterministic leakage-aware train/development splits.""" from __future__ import annotations import argparse import sys from pathlib import Path 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 ( # noqa: E402 DEFAULT_SPLIT_RATIOS, assign_splits, build_split_report, parse_split_ratios, read_manifest, write_json, write_manifest, ) DEFAULT_INPUT = PROJECT_ROOT / "data" / "processed" / "manifest.jsonl" DEFAULT_OUTPUT = PROJECT_ROOT / "data" / "processed" / "splits.jsonl" DEFAULT_REPORT = PROJECT_ROOT / "artifacts" / "split_report.json" def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description=( "Assign whole transitive leakage groups to deterministic approximately stratified splits. " "This command reads JSONL only and does not import PyArrow." ) ) parser.add_argument("--input", type=Path, default=DEFAULT_INPUT, help="Audit manifest JSONL") parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT, help="Split manifest JSONL") parser.add_argument("--report", type=Path, default=DEFAULT_REPORT, help="Split summary JSON") parser.add_argument( "--split", dest="split_specs", action="append", metavar="NAME=WEIGHT", help="Repeat for each split; defaults to train=0.9 and validation=0.1", ) parser.add_argument("--seed", type=int, default=42, help="Deterministic assignment seed") parser.add_argument( "--stratify", default="endpoint,language,dataset", help="Comma-separated marginal fields to balance", ) parser.add_argument( "--holdout-field", action="append", default=[], help="Link equal field values into one split (repeat for source/speaker-held-out splits)", ) parser.add_argument("--limit", type=int, help="Read at most N manifest rows (smoke runs)") parser.add_argument( "--include-invalid", action="store_true", help="Include rows with validation errors; by default they are excluded from model splits", ) return parser def main(argv: list[str] | None = None) -> int: args = _parser().parse_args(argv) if args.limit is not None and args.limit < 0: raise SystemExit("--limit cannot be negative") try: ratios = parse_split_ratios(args.split_specs) if args.split_specs else DEFAULT_SPLIT_RATIOS stratify_fields = tuple( field.strip() for field in args.stratify.split(",") if field.strip() ) input_rows = list(read_manifest(args.input, limit=args.limit)) if args.include_invalid: eligible = input_rows else: eligible = [row for row in input_rows if not row.get("validation_errors")] split_rows = assign_splits( eligible, ratios=ratios, seed=args.seed, stratify_fields=stratify_fields, holdout_fields=tuple(args.holdout_field), ) report = build_split_report( split_rows, stratify_fields=stratify_fields, holdout_fields=tuple(args.holdout_field), ) report["configuration"] = { "ratios": dict(ratios), "seed": args.seed, "stratify_fields": list(stratify_fields), "holdout_fields": list(args.holdout_field), "include_invalid": bool(args.include_invalid), } report["input_records"] = len(input_rows) report["excluded_invalid_records"] = len(input_rows) - len(eligible) if not report["leakage"]["is_valid"]: print("split validation failed: leakage was detected", file=sys.stderr) return 2 write_manifest(args.output, split_rows) write_json(args.report, report) except (FileNotFoundError, OSError, ValueError) as exc: print(f"split preparation failed: {exc}", file=sys.stderr) return 2 counts = ", ".join(f"{name}={count:,}" for name, count in report["split_counts"].items()) print( f"wrote {len(split_rows):,} rows to {args.output} ({counts}); report: {args.report}", file=sys.stderr, ) return 0 if __name__ == "__main__": raise SystemExit(main())