File size: 3,104 Bytes
8fa3dd6
 
 
 
 
 
 
 
 
 
 
 
16bf4fc
 
 
 
 
 
 
 
8fa3dd6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16bf4fc
8fa3dd6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
import json
from pathlib import Path
import random
import re
from typing import List


def read_records(path: str) -> List[str]:
    text = Path(path).read_text(encoding="utf-8", errors="ignore")
    if "<|eos|>" in text:
        parts = text.split("<|eos|>")
    elif "<|record_end|>" in text:
        parts = text.split("<|record_end|>")
    elif "<|end|>" in text:
        # Backward-compatible fallback for older corpora. New chat-style corpora
        # should use <|eos|> as the record separator because <|end|> is also
        # used inside conversations to mark turn boundaries.
        parts = text.split("<|end|>")
    else:
        # Fallback for plain corpora: split on blank lines.
        parts = re.split(r"\n\s*\n", text)
    records = []
    for part in parts:
        part = part.strip()
        if part:
            records.append(part)
    return records


def write_records(path: str, records: List[str]) -> None:
    out = Path(path)
    out.parent.mkdir(parents=True, exist_ok=True)
    with out.open("w", encoding="utf-8") as f:
        for rec in records:
            f.write(rec.strip())
            f.write("\n<|eos|>\n")


def split_corpus(
    input_path: str,
    train_output: str,
    val_output: str,
    val_ratio: float = 0.02,
    min_val_records: int = 1,
    seed: int = 42,
) -> dict:
    if not 0.0 < val_ratio < 0.5:
        raise ValueError("val_ratio must be between 0 and 0.5")
    records = read_records(input_path)
    if len(records) < 2:
        raise ValueError("Need at least two records to create a validation split")
    rng = random.Random(seed)
    rng.shuffle(records)
    val_n = max(min_val_records, int(round(len(records) * val_ratio)))
    val_n = min(val_n, len(records) - 1)
    val_records = records[:val_n]
    train_records = records[val_n:]
    write_records(train_output, train_records)
    write_records(val_output, val_records)
    return {
        "input": input_path,
        "train_output": train_output,
        "val_output": val_output,
        "records_total": len(records),
        "records_train": len(train_records),
        "records_val": len(val_records),
        "val_ratio_actual": len(val_records) / len(records),
        "seed": seed,
    }


def main() -> None:
    parser = argparse.ArgumentParser(description="Split an Ares corpus into train/validation files.")
    parser.add_argument("--input", required=True)
    parser.add_argument("--train-output", required=True)
    parser.add_argument("--val-output", required=True)
    parser.add_argument("--val-ratio", type=float, default=0.02)
    parser.add_argument("--min-val-records", type=int, default=1)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()
    result = split_corpus(
        input_path=args.input,
        train_output=args.train_output,
        val_output=args.val_output,
        val_ratio=args.val_ratio,
        min_val_records=args.min_val_records,
        seed=args.seed,
    )
    print(json.dumps(result, indent=2))


if __name__ == "__main__":
    main()