Spaces:
Runtime error
Runtime error
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()
|