FPLpredictor / train_model.py
bennymach's picture
Upload 7 files
314bef8 verified
Raw History Blame Contribute Delete
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()