Pafkun333's picture
Deploy World Cup Predictor V2
bc46574 verified
Raw
History Blame Contribute Delete
1.62 kB
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()