File size: 4,468 Bytes
35d483e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | #!/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())
|