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