File size: 6,136 Bytes
087e4a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
from __future__ import annotations

import argparse
import json
import random
from pathlib import Path
from typing import Any


def read_records(path: Path) -> tuple[list[dict[str, Any]], str]:
    if path.suffix == ".jsonl":
        rows = [
            json.loads(line)
            for line in path.read_text(encoding="utf-8").splitlines()
            if line.strip()
        ]
        return [dict(row) for row in rows if isinstance(row, dict)], "jsonl"
    raw = json.loads(path.read_text(encoding="utf-8"))
    if isinstance(raw, list):
        return [dict(row) for row in raw if isinstance(row, dict)], "list"
    if isinstance(raw, dict):
        for key in ("data", "examples", "instances"):
            value = raw.get(key)
            if isinstance(value, list):
                return [dict(row) for row in value if isinstance(row, dict)], key
    raise ValueError(f"Unsupported dataset shape: {path}")


def write_records(path: Path, records: list[dict[str, Any]], shape: str) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    if shape == "jsonl":
        path.write_text(
            "".join(json.dumps(record, ensure_ascii=False) + "\n" for record in records),
            encoding="utf-8",
        )
        return
    payload: Any = records if shape == "list" else {shape: records}
    path.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")


def split_records(
    records: list[dict[str, Any]],
    *,
    val_ratio: float,
    min_val: int,
    max_val: int | None,
    seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
    if not records:
        return [], []
    shuffled = list(records)
    random.Random(seed).shuffle(shuffled)
    val_size = max(min_val, int(round(len(shuffled) * val_ratio)))
    if max_val is not None:
        val_size = min(val_size, max_val)
    val_size = max(1, min(val_size, len(shuffled) - 1)) if len(shuffled) > 1 else 1
    return shuffled[val_size:], shuffled[:val_size]


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        description="Split each task's train file into deterministic train and validation files."
    )
    parser.add_argument("--config", type=Path, required=True, help="Task-stream config to read.")
    parser.add_argument("--output-config", type=Path, help="Updated task-stream config to write.")
    parser.add_argument("--output-dir", type=Path, help="Common output directory for split files.")
    parser.add_argument("--val-ratio", type=float, default=0.2)
    parser.add_argument("--min-val", type=int, default=1)
    parser.add_argument("--max-val", type=int)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--overwrite", action="store_true")
    parser.add_argument("--dry-run", action="store_true")
    return parser


def main() -> int:
    args = build_parser().parse_args()
    if not 0 < args.val_ratio < 1:
        raise ValueError("--val-ratio must be between 0 and 1.")
    if args.min_val < 0 or (args.max_val is not None and args.max_val < 1):
        raise ValueError("Validation limits must be non-negative and --max-val must be positive.")

    config_path = args.config.resolve()
    config = json.loads(config_path.read_text(encoding="utf-8"))
    config_base = config_path.parent.parent if config_path.parent.name == "configs" else config_path.parent
    output_root = args.output_dir.resolve() if args.output_dir else None
    changed = False

    for task_index, task in enumerate(config.get("tasks", [])):
        task_name = str(task.get("task_name", "task"))
        splits = task.get("splits", {})
        if not isinstance(splits, dict) or not splits.get("train"):
            print(f"SKIP {task_name}: no train split")
            continue
        train_path = _resolve_manifest_path(str(splits["train"]), config_base)
        records, shape = read_records(train_path)
        train_records, val_records = split_records(
            records,
            val_ratio=args.val_ratio,
            min_val=args.min_val,
            max_val=args.max_val,
            seed=args.seed + task_index,
        )
        target_dir = output_root / task_name if output_root else train_path.parent
        suffix = ".jsonl" if shape == "jsonl" else ".json"
        split_train_path = target_dir / f"train_split{suffix}"
        split_val_path = target_dir / f"validation_split{suffix}"
        print(
            f"{task_name}: {len(records)} -> train={len(train_records)} "
            f"validation={len(val_records)} ({split_train_path}, {split_val_path})"
        )
        if args.dry_run:
            continue
        if not args.overwrite and (split_train_path.exists() or split_val_path.exists()):
            raise FileExistsError(
                f"Split file already exists for {task_name}; use --overwrite to replace."
            )
        write_records(split_train_path, train_records, shape)
        write_records(split_val_path, val_records, shape)
        splits["train"] = _portable_path(split_train_path, config_base)
        splits["validation"] = _portable_path(split_val_path, config_base)
        splits["val"] = splits["validation"]
        splits.setdefault("anchor", splits["validation"])
        changed = True

    if args.output_config and changed and not args.dry_run:
        args.output_config.parent.mkdir(parents=True, exist_ok=True)
        args.output_config.write_text(
            json.dumps(config, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
        )
    elif changed and not args.dry_run:
        print("Split files written. Pass --output-config to write an updated task-stream config.")
    return 0


def _resolve_manifest_path(raw_path: str, config_base: Path) -> Path:
    path = Path(raw_path).expanduser()
    return path.resolve() if path.is_absolute() else (config_base / path).resolve()


def _portable_path(path: Path, config_base: Path) -> str:
    try:
        return str(path.resolve().relative_to(config_base.resolve()))
    except ValueError:
        return str(path.resolve())


if __name__ == "__main__":
    raise SystemExit(main())