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()