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