File size: 1,619 Bytes
bc46574
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
from pathlib import Path

import pandas as pd

from worldcup_predictor.features.match_features import chronological_train_validation_test_split


def split_dataset(
    input_path: Path,
    output_dir: Path,
    *,
    validation_fraction: float = 0.15,
    test_fraction: float = 0.2,
) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
    data = pd.read_csv(input_path)
    train, validation, test = chronological_train_validation_test_split(
        data,
        validation_fraction=validation_fraction,
        test_fraction=test_fraction,
    )
    output_dir.mkdir(parents=True, exist_ok=True)
    train.to_csv(output_dir / "train.csv", index=False)
    validation.to_csv(output_dir / "validation.csv", index=False)
    test.to_csv(output_dir / "test.csv", index=False)
    return train, validation, test


def main() -> None:
    parser = argparse.ArgumentParser(description="Create chronological train/validation/test splits.")
    parser.add_argument("input", type=Path)
    parser.add_argument("output_dir", type=Path)
    parser.add_argument("--validation-fraction", type=float, default=0.15)
    parser.add_argument("--test-fraction", type=float, default=0.2)
    args = parser.parse_args()

    train, validation, test = split_dataset(
        args.input,
        args.output_dir,
        validation_fraction=args.validation_fraction,
        test_fraction=args.test_fraction,
    )
    print(f"train_rows: {len(train)}")
    print(f"validation_rows: {len(validation)}")
    print(f"test_rows: {len(test)}")


if __name__ == "__main__":
    main()