Downloads · 30 days
10
15% of all-time downloads
rokati/focal_gnn_one_graph
focal_gnn_one_graph is a machine learning model from rokati. Use it for the machine learning task on the model card, and read the license before you ship it in a product. It is set up for pytorch.
This is a Heterogeneous Graph Neural Network (GNN) model trained to predict Expected Goals (xG) in football/soccer using a single persistent graph with shot-indexed edges.
Downloads · 30 days
10
15% of all-time downloads
All-time downloads
65
Public
Repo size
1.9 MB
Likes
0
Public
Click a slice to open those files.
.pth1.9 MB · 99%
From the Hugging Face model README
This is a Heterogeneous Graph Neural Network (GNN) model trained to predict Expected Goals (xG) in football/soccer using a single persistent graph with shot-indexed edges.
The model uses the following contextual features for each shot:
All edges are indexed by shot_idx for efficient masking during prediction.
import torch
from torch_geometric.data import HeteroData
from huggingface_hub import hf_hub_download
import importlib.util
import json
# Download files
model_path = hf_hub_download(repo_id="rokati/focal_gnn_one_graph", filename="best_gnn_model.pth")
architecture_path = hf_hub_download(repo_id="rokati/focal_gnn_one_graph", filename="model_architecture.py")
config_path = hf_hub_download(repo_id="rokati/focal_gnn_one_graph", filename="config.json")
# Load configuration
with open(config_path, 'r') as f:
config = json.load(f)
# Load architecture
spec = importlib.util.spec_from_file_location("model_architecture", architecture_path)
model_module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(model_module)
# Create model instance
model = model_module.XGNet(
num_players=config['num_players'],
hid=config['hidden_dim'],
p=config['dropout_rate'],
heads=config['num_heads'],
num_layers=config['num_layers'],
use_norm=config['use_norm'],
num_global_features=config['num_global_features']
)
# Load weights
model.load_state_dict(torch.load(model_path, map_location=torch.device('cpu')))
model.eval()
# Prepare persistent graph (build once with all players/shots)
# Then predict by passing the graph and shot_idx:
with torch.no_grad():
xg_prediction = torch.sigmoid(model(graph, shot_idx)).item()
The model was trained with:
The heterogeneous GNN uses:
Unlike traditional approaches that create separate graphs for each shot, this model uses:
MIT