Download train_model.py from bennymach/FPLpredictor: direct link, hf CLI and curl.
- Browser
- Download file 3.32 kB
-
https://huggingface.co/bennymach/FPLpredictor/resolve/main/train_model.py
- Command line
-
hf download hf://bennymach/FPLpredictor/train_model.py
-
curl -L -o train_model.py https://huggingface.co/bennymach/FPLpredictor/resolve/main/train_model.py
3.32 kB
| from __future__ import annotations | |
| import json | |
| from datetime import date | |
| import os | |
| from pathlib import Path | |
| import joblib | |
| import pandas as pd | |
| from huggingface_hub import HfApi | |
| from sklearn.metrics import mean_absolute_error | |
| from xgboost import XGBRegressor | |
| from data_pipeline import FEATURES, build_training_frame, cache_histories, fetch_fixtures, fetch_players | |
| ROOT = Path(__file__).resolve().parent | |
| MODEL_PATH = ROOT / "xgb_model.joblib" | |
| METADATA_PATH = ROOT / "model_metadata.json" | |
| CACHE_PATH = ROOT / ".cache" / "player_histories" | |
| MODEL_REPO_ID = os.getenv("MODEL_REPO_ID", "").strip() | |
| def new_model() -> XGBRegressor: | |
| return XGBRegressor( | |
| n_estimators=200, | |
| learning_rate=0.05, | |
| max_depth=3, | |
| objective="reg:squarederror", | |
| random_state=42, | |
| ) | |
| def evaluate_by_gameweek(frame: pd.DataFrame): | |
| """Use whole future gameweeks as validation sets, avoiding row leakage.""" | |
| gameweeks = sorted(frame["gameweek"].unique()) | |
| if len(gameweeks) < 3: | |
| raise ValueError("At least three completed gameweeks are required for validation.") | |
| split_at = max(1, int(len(gameweeks) * 0.8)) | |
| train_weeks, test_weeks = gameweeks[:split_at], gameweeks[split_at:] | |
| train = frame[frame["gameweek"].isin(train_weeks)] | |
| test = frame[frame["gameweek"].isin(test_weeks)] | |
| validation_model = new_model().fit(train[FEATURES], train["points"]) | |
| predictions = validation_model.predict(test[FEATURES]) | |
| return mean_absolute_error(test["points"], predictions), train_weeks, test_weeks | |
| def main() -> None: | |
| print("Fetching FPL data and player histories...") | |
| players = fetch_players() | |
| fixtures = fetch_fixtures() | |
| histories = cache_histories(players, CACHE_PATH) | |
| frame = build_training_frame(players, histories, fixtures) | |
| if frame.empty: | |
| raise RuntimeError("No valid training rows were created.") | |
| mae, train_weeks, test_weeks = evaluate_by_gameweek(frame) | |
| model = new_model().fit(frame[FEATURES], frame["points"]) | |
| joblib.dump(model, MODEL_PATH) | |
| metadata = { | |
| "trained_on": date.today().isoformat(), | |
| "features": FEATURES, | |
| "training_rows": len(frame), | |
| "gameweeks": sorted(int(week) for week in frame["gameweek"].unique()), | |
| "validation_mae": round(float(mae), 4), | |
| "validation_training_gameweeks": [int(week) for week in train_weeks], | |
| "validation_test_gameweeks": [int(week) for week in test_weeks], | |
| } | |
| METADATA_PATH.write_text(json.dumps(metadata, indent=2), encoding="utf-8") | |
| print(f"Validation MAE: {mae:.2f}") | |
| print(f"Saved model to {MODEL_PATH}") | |
| print(f"Saved metadata to {METADATA_PATH}") | |
| if MODEL_REPO_ID: | |
| api = HfApi() | |
| api.create_repo(repo_id=MODEL_REPO_ID, repo_type="model", exist_ok=True) | |
| for artifact in (MODEL_PATH, METADATA_PATH): | |
| api.upload_file( | |
| path_or_fileobj=str(artifact), | |
| path_in_repo=artifact.name, | |
| repo_id=MODEL_REPO_ID, | |
| repo_type="model", | |
| commit_message=f"Retrain model on {metadata['trained_on']}", | |
| ) | |
| print(f"Uploaded model artifacts to {MODEL_REPO_ID}") | |
| if __name__ == "__main__": | |
| main() |