| |
| """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 ( |
| 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()) |
|
|