Spaces:
Running
Running
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from io import BytesIO | |
| from pathlib import Path | |
| import base64 | |
| import html | |
| import json | |
| import random | |
| import sys | |
| import matplotlib.pyplot as plt | |
| from matplotlib import colors | |
| from matplotlib.patches import Rectangle | |
| from mplsoccer import Pitch | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| from torch import nn | |
| from torch.utils.data import DataLoader, TensorDataset | |
| SCRIPT_DIR = Path(__file__).resolve().parent | |
| PROJECT_ROOT = SCRIPT_DIR.parent | |
| MODEL_DIR = PROJECT_ROOT / "data" / "modeling" | |
| REPORTS_DIR = PROJECT_ROOT / "reports" | |
| REPORT_PATH = REPORTS_DIR / "attack_distribution_gnn_report.html" | |
| JSON_PATH = MODEL_DIR / "attack_distribution_gnn_metrics.json" | |
| MODEL_PATH = MODEL_DIR / "attack_distribution_gnn_bundle.pt" | |
| PRED_PATH = MODEL_DIR / "attack_distribution_gnn_test_predictions.parquet" | |
| if str(SCRIPT_DIR) not in sys.path: | |
| sys.path.insert(0, str(SCRIPT_DIR)) | |
| import train_attack_prediction_ffn as base # noqa: E402 | |
| RANDOM_SEED = 42 | |
| BATCH_SIZE = 256 | |
| MAX_EPOCHS = 260 | |
| PATIENCE = 32 | |
| LEARNING_RATE = 4e-4 | |
| WEIGHT_DECAY = 1e-5 | |
| SMOOTH_LAMBDA = 0.005 | |
| GATE_ENTROPY_LAMBDA = 0.001 | |
| ZONE_ORDER = [ | |
| "Deep_Cross__Der_", | |
| "Half_Space__Der_", | |
| "Creativity_Zone", | |
| "Half_Space__Izq_", | |
| "Deep_Cross__Izq_", | |
| "Cross__Der_", | |
| "Cut_Back__Der_", | |
| "Scoring_Zone", | |
| "Cut_Back__Izq_", | |
| "Cross__Izq_", | |
| ] | |
| PRETTY_ZONE = { | |
| "Scoring_Zone": "Scoring Zone", | |
| "Creativity_Zone": "Creativity Zone", | |
| "Half_Space__Izq_": "Half-Space Izq", | |
| "Half_Space__Der_": "Half-Space Der", | |
| "Cut_Back__Izq_": "Cut-Back Izq", | |
| "Cut_Back__Der_": "Cut-Back Der", | |
| "Cross__Izq_": "Cross Izq", | |
| "Cross__Der_": "Cross Der", | |
| "Deep_Cross__Izq_": "Deep Cross Izq", | |
| "Deep_Cross__Der_": "Deep Cross Der", | |
| } | |
| NODE_FEATURE_PREFIXES = [ | |
| "short_mean__actual_attack_share__", | |
| "long_mean__actual_attack_share__", | |
| "short_mean__actual_pv_share__", | |
| "long_mean__actual_pv_share__", | |
| "short_mean__actual_conceded_attack_share__", | |
| "long_mean__actual_conceded_attack_share__", | |
| "short_mean__actual_conceded_pv_share__", | |
| "long_mean__actual_conceded_pv_share__", | |
| "opp__short_mean__actual_attack_share__", | |
| "opp__long_mean__actual_attack_share__", | |
| "opp__short_mean__actual_pv_share__", | |
| "opp__long_mean__actual_pv_share__", | |
| "opp__short_mean__actual_conceded_attack_share__", | |
| "opp__long_mean__actual_conceded_attack_share__", | |
| "opp__short_mean__actual_conceded_pv_share__", | |
| "opp__long_mean__actual_conceded_pv_share__", | |
| ] | |
| class SplitData: | |
| global_x: np.ndarray | |
| node_x: np.ndarray | |
| y_attack: np.ndarray | |
| baseline_long: np.ndarray | |
| baseline_short: np.ndarray | |
| metadata: pd.DataFrame | |
| def _set_seed(seed: int = RANDOM_SEED) -> None: | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| def _attack_zone_rectangles() -> dict[str, list[tuple[float, float, float, float]]]: | |
| zones: dict[str, list[tuple[float, float, float, float]]] = {} | |
| def add(z: str, x0: float, x1: float, y0: float, y1: float) -> None: | |
| zones.setdefault(z, []).append((x0, y0, x1 - x0, y1 - y0)) | |
| add("Scoring_Zone", 83, 100, 37, 63) | |
| add("Cut_Back__Izq_", 83, 100, 63, 79) | |
| add("Cross__Izq_", 83, 100, 79, 100) | |
| add("Cut_Back__Der_", 83, 100, 21, 37) | |
| add("Cross__Der_", 83, 100, 0, 21) | |
| add("Creativity_Zone", 60, 83, 37, 63) | |
| add("Half_Space__Izq_", 60, 83, 63, 79) | |
| add("Deep_Cross__Izq_", 60, 83, 79, 100) | |
| add("Half_Space__Der_", 60, 83, 21, 37) | |
| add("Deep_Cross__Der_", 60, 83, 0, 21) | |
| return zones | |
| ZONES_RECTS = _attack_zone_rectangles() | |
| def _img_to_base64(fig: plt.Figure) -> str: | |
| buf = BytesIO() | |
| fig.savefig(buf, format="png", dpi=180, bbox_inches="tight", facecolor=fig.get_facecolor()) | |
| plt.close(fig) | |
| return base64.b64encode(buf.getvalue()).decode("ascii") | |
| def _normalized_adjacency() -> tuple[torch.Tensor, torch.Tensor]: | |
| neighbors = { | |
| "Deep_Cross__Der_": ["Half_Space__Der_", "Cross__Der_", "Deep_Cross__Izq_"], | |
| "Half_Space__Der_": ["Deep_Cross__Der_", "Creativity_Zone", "Cut_Back__Der_", "Half_Space__Izq_"], | |
| "Creativity_Zone": ["Half_Space__Der_", "Half_Space__Izq_", "Scoring_Zone", "Cut_Back__Der_", "Cut_Back__Izq_"], | |
| "Half_Space__Izq_": ["Creativity_Zone", "Deep_Cross__Izq_", "Cut_Back__Izq_", "Half_Space__Der_"], | |
| "Deep_Cross__Izq_": ["Half_Space__Izq_", "Cross__Izq_", "Deep_Cross__Der_"], | |
| "Cross__Der_": ["Deep_Cross__Der_", "Cut_Back__Der_", "Cross__Izq_"], | |
| "Cut_Back__Der_": ["Cross__Der_", "Scoring_Zone", "Half_Space__Der_", "Cut_Back__Izq_", "Creativity_Zone"], | |
| "Scoring_Zone": ["Cut_Back__Der_", "Cut_Back__Izq_", "Creativity_Zone"], | |
| "Cut_Back__Izq_": ["Cross__Izq_", "Scoring_Zone", "Half_Space__Izq_", "Cut_Back__Der_", "Creativity_Zone"], | |
| "Cross__Izq_": ["Deep_Cross__Izq_", "Cut_Back__Izq_", "Cross__Der_"], | |
| } | |
| n = len(ZONE_ORDER) | |
| adj = np.zeros((n, n), dtype=np.float32) | |
| edge_pairs: list[tuple[int, int]] = [] | |
| zone_to_idx = {zone: i for i, zone in enumerate(ZONE_ORDER)} | |
| for zone, neighs in neighbors.items(): | |
| i = zone_to_idx[zone] | |
| for neigh in neighs: | |
| j = zone_to_idx[neigh] | |
| adj[i, j] = 1.0 | |
| edge_pairs.append((i, j)) | |
| deg = adj.sum(axis=1, keepdims=True) | |
| deg = np.where(deg > 0, deg, 1.0) | |
| adj = adj / deg | |
| edge_index = torch.tensor(edge_pairs, dtype=torch.long) | |
| return torch.tensor(adj, dtype=torch.float32), edge_index | |
| ADJ_MATRIX, EDGE_INDEX = _normalized_adjacency() | |
| class GraphBlock(nn.Module): | |
| def __init__(self, in_dim: int, out_dim: int, dropout: float) -> None: | |
| super().__init__() | |
| self.self_lin = nn.Linear(in_dim, out_dim) | |
| self.neigh_lin = nn.Linear(in_dim, out_dim) | |
| self.norm = nn.LayerNorm(out_dim) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor, adj: torch.Tensor) -> torch.Tensor: | |
| neigh = torch.einsum("ij,bjf->bif", adj, x) | |
| h = self.self_lin(x) + self.neigh_lin(neigh) | |
| h = self.norm(h) | |
| h = torch.relu(h) | |
| return self.dropout(h) | |
| class ResidualAttackGNN(nn.Module): | |
| def __init__(self, node_dim: int, global_dim: int, hidden_dim: int = 96, global_hidden: int = 96, num_layers: int = 3) -> None: | |
| super().__init__() | |
| self.global_encoder = nn.Sequential( | |
| nn.Linear(global_dim, 192), | |
| nn.ReLU(), | |
| nn.Dropout(0.10), | |
| nn.Linear(192, global_hidden), | |
| nn.ReLU(), | |
| ) | |
| self.node_encoder = nn.Sequential( | |
| nn.Linear(node_dim + global_hidden, hidden_dim), | |
| nn.ReLU(), | |
| ) | |
| self.graph_blocks = nn.ModuleList([GraphBlock(hidden_dim, hidden_dim, dropout=0.06) for _ in range(num_layers)]) | |
| self.gate_head = nn.Linear(hidden_dim, 1) | |
| self.delta_head = nn.Linear(hidden_dim, 1) | |
| def forward( | |
| self, | |
| node_x: torch.Tensor, | |
| global_x: torch.Tensor, | |
| baseline_long: torch.Tensor, | |
| baseline_short: torch.Tensor, | |
| adj: torch.Tensor, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| g = self.global_encoder(global_x) | |
| g_rep = g.unsqueeze(1).expand(-1, node_x.size(1), -1) | |
| h = self.node_encoder(torch.cat([node_x, g_rep], dim=-1)) | |
| gate = torch.sigmoid(self.gate_head(h)).squeeze(-1) | |
| mixed = gate * baseline_short + (1.0 - gate) * baseline_long | |
| for block in self.graph_blocks: | |
| h = h + block(h, adj) | |
| delta = self.delta_head(h).squeeze(-1) | |
| logits = torch.log(torch.clamp(mixed, min=1e-6)) + delta | |
| pred = torch.softmax(logits, dim=1) | |
| return pred, delta, gate, mixed | |
| def _load_data() -> tuple[pd.DataFrame, list[str]]: | |
| df = base._load_dataset() | |
| df = df[df["usable_for_model"]].copy().reset_index(drop=True) | |
| attack_targets = [f"target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| missing = [c for c in attack_targets if c not in df.columns] | |
| if missing: | |
| raise ValueError(f"Faltan targets esperados: {missing}") | |
| return df, attack_targets | |
| def _build_feature_matrices(df: pd.DataFrame, attack_targets: list[str]) -> tuple[pd.DataFrame, np.ndarray, np.ndarray, list[str], list[str]]: | |
| wide_features, _, _, _, dummy_cols = base._feature_matrix(df) | |
| node_cols_by_zone: dict[str, list[str]] = {zone: [] for zone in ZONE_ORDER} | |
| used_node_cols: set[str] = set() | |
| for prefix in NODE_FEATURE_PREFIXES: | |
| for zone in ZONE_ORDER: | |
| col = f"{prefix}{zone}" | |
| if col in df.columns: | |
| node_cols_by_zone[zone].append(col) | |
| used_node_cols.add(col) | |
| node_feat_names = [f"{prefix}{zone}" for prefix in NODE_FEATURE_PREFIXES for zone in ZONE_ORDER if f"{prefix}{zone}" in df.columns] | |
| node_tensor = np.stack( | |
| [df[node_cols_by_zone[zone]].apply(pd.to_numeric, errors="coerce").to_numpy(dtype=float) for zone in ZONE_ORDER], | |
| axis=1, | |
| ) | |
| global_exclude_cols = used_node_cols | set(attack_targets) | |
| global_cols = [c for c in wide_features.columns if c not in global_exclude_cols] | |
| global_features = wide_features[global_cols].copy() | |
| return global_features, node_tensor, np.array(ZONE_ORDER), global_cols, node_feat_names | |
| def _standardize_global(train_idx: np.ndarray, global_features: pd.DataFrame) -> tuple[np.ndarray, dict]: | |
| fill_values = global_features.iloc[train_idx].median(numeric_only=False) | |
| filled = global_features.fillna(fill_values) | |
| means = filled.iloc[train_idx].mean(axis=0) | |
| stds = filled.iloc[train_idx].std(axis=0, ddof=0).replace(0, 1.0) | |
| scaled = ((filled - means) / stds).to_numpy(dtype=np.float32) | |
| bundle = { | |
| "fill_values": fill_values.to_dict(), | |
| "means": means.to_dict(), | |
| "stds": stds.to_dict(), | |
| "global_feature_columns": list(global_features.columns), | |
| } | |
| return scaled, bundle | |
| def _standardize_node(train_idx: np.ndarray, node_tensor: np.ndarray) -> tuple[np.ndarray, dict]: | |
| train = node_tensor[train_idx] | |
| fill = np.nanmedian(train, axis=0) | |
| filled = np.where(np.isnan(node_tensor), fill[None, :, :], node_tensor) | |
| means = filled[train_idx].mean(axis=0) | |
| stds = filled[train_idx].std(axis=0, ddof=0) | |
| stds = np.where(stds > 0, stds, 1.0) | |
| scaled = (filled - means[None, :, :]) / stds[None, :, :] | |
| bundle = { | |
| "means": means.tolist(), | |
| "stds": stds.tolist(), | |
| "fill_values": fill.tolist(), | |
| } | |
| return scaled.astype(np.float32), bundle | |
| def _train_val_test_indices(df: pd.DataFrame) -> tuple[np.ndarray, np.ndarray, np.ndarray, str]: | |
| train_df = df[df["split"] == "train"].copy() | |
| train_idx_pd, val_idx_pd, val_start_date = base._build_temporal_validation(train_df) | |
| test_idx = df.index[df["split"] == "test"].to_numpy() | |
| return train_idx_pd.to_numpy(), val_idx_pd.to_numpy(), test_idx, val_start_date | |
| def _make_split( | |
| df: pd.DataFrame, | |
| idx: np.ndarray, | |
| global_x: np.ndarray, | |
| node_x: np.ndarray, | |
| y_attack: np.ndarray, | |
| baseline_long: np.ndarray, | |
| baseline_short: np.ndarray, | |
| ) -> SplitData: | |
| return SplitData( | |
| global_x=global_x[idx].astype(np.float32), | |
| node_x=node_x[idx].astype(np.float32), | |
| y_attack=y_attack[idx].astype(np.float32), | |
| baseline_long=baseline_long[idx].astype(np.float32), | |
| baseline_short=baseline_short[idx].astype(np.float32), | |
| metadata=df.iloc[idx].copy().reset_index(drop=True), | |
| ) | |
| def _make_loader(split: SplitData, shuffle: bool) -> DataLoader: | |
| dataset = TensorDataset( | |
| torch.from_numpy(split.global_x), | |
| torch.from_numpy(split.node_x), | |
| torch.from_numpy(split.y_attack), | |
| torch.from_numpy(split.baseline_long), | |
| torch.from_numpy(split.baseline_short), | |
| ) | |
| return DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=shuffle) | |
| def _distribution_loss( | |
| pred: torch.Tensor, | |
| target: torch.Tensor, | |
| delta: torch.Tensor, | |
| gate: torch.Tensor, | |
| ) -> tuple[torch.Tensor, dict[str, float]]: | |
| eps = 1e-8 | |
| log_pred = torch.log(torch.clamp(pred, min=eps)) | |
| kl = torch.nn.functional.kl_div(log_pred, target, reduction="batchmean") | |
| if EDGE_INDEX.numel(): | |
| smooth = torch.mean((delta[:, EDGE_INDEX[:, 0]] - delta[:, EDGE_INDEX[:, 1]]) ** 2) | |
| else: | |
| smooth = torch.zeros((), device=pred.device) | |
| gate_entropy = -torch.mean(gate * torch.log(torch.clamp(gate, min=eps)) + (1 - gate) * torch.log(torch.clamp(1 - gate, min=eps))) | |
| loss = kl + (SMOOTH_LAMBDA * smooth) + (GATE_ENTROPY_LAMBDA * gate_entropy) | |
| parts = { | |
| "kl": float(kl.item()), | |
| "smooth": float(smooth.item()), | |
| "gate_entropy": float(gate_entropy.item()), | |
| } | |
| return loss, parts | |
| def _evaluate_loader(model: ResidualAttackGNN, loader: DataLoader, device: torch.device) -> dict[str, float]: | |
| model.eval() | |
| total = 0.0 | |
| total_kl = 0.0 | |
| n_batches = 0 | |
| adj = ADJ_MATRIX.to(device) | |
| with torch.no_grad(): | |
| for global_x, node_x, y_attack, baseline_long, baseline_short in loader: | |
| global_x = global_x.to(device) | |
| node_x = node_x.to(device) | |
| y_attack = y_attack.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| pred, delta, gate, mixed = model(node_x, global_x, baseline_long, baseline_short, adj) | |
| loss, parts = _distribution_loss(pred, y_attack, delta, gate) | |
| total += float(loss.item()) | |
| total_kl += parts["kl"] | |
| n_batches += 1 | |
| return {"loss": total / max(n_batches, 1), "kl": total_kl / max(n_batches, 1)} | |
| def _train_model(train: SplitData, val: SplitData) -> tuple[ResidualAttackGNN, list[dict[str, float]]]: | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model = ResidualAttackGNN( | |
| node_dim=train.node_x.shape[2], | |
| global_dim=train.global_x.shape[1], | |
| hidden_dim=96, | |
| global_hidden=96, | |
| num_layers=3, | |
| ).to(device) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY) | |
| train_loader = _make_loader(train, shuffle=True) | |
| val_loader = _make_loader(val, shuffle=False) | |
| best_state = None | |
| best_val = float("inf") | |
| patience_left = PATIENCE | |
| history: list[dict[str, float]] = [] | |
| adj = ADJ_MATRIX.to(device) | |
| for epoch in range(1, MAX_EPOCHS + 1): | |
| model.train() | |
| running = 0.0 | |
| n_batches = 0 | |
| for global_x, node_x, y_attack, baseline_long, baseline_short in train_loader: | |
| global_x = global_x.to(device) | |
| node_x = node_x.to(device) | |
| y_attack = y_attack.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| optimizer.zero_grad() | |
| pred, delta, gate, mixed = model(node_x, global_x, baseline_long, baseline_short, adj) | |
| loss, parts = _distribution_loss(pred, y_attack, delta, gate) | |
| loss.backward() | |
| optimizer.step() | |
| running += float(loss.item()) | |
| n_batches += 1 | |
| val_metrics = _evaluate_loader(model, val_loader, device) | |
| history.append( | |
| { | |
| "epoch": epoch, | |
| "train_loss": running / max(n_batches, 1), | |
| "val_loss": val_metrics["loss"], | |
| "val_kl": val_metrics["kl"], | |
| } | |
| ) | |
| if val_metrics["loss"] < best_val - 1e-6: | |
| best_val = val_metrics["loss"] | |
| best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()} | |
| patience_left = PATIENCE | |
| else: | |
| patience_left -= 1 | |
| if patience_left <= 0: | |
| break | |
| if best_state is None: | |
| best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()} | |
| model.load_state_dict(best_state) | |
| return model, history | |
| def _predict(model: ResidualAttackGNN, split: SplitData) -> tuple[np.ndarray, np.ndarray, np.ndarray]: | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model.eval() | |
| preds: list[np.ndarray] = [] | |
| gates: list[np.ndarray] = [] | |
| mixeds: list[np.ndarray] = [] | |
| loader = _make_loader(split, shuffle=False) | |
| adj = ADJ_MATRIX.to(device) | |
| with torch.no_grad(): | |
| for global_x, node_x, y_attack, baseline_long, baseline_short in loader: | |
| global_x = global_x.to(device) | |
| node_x = node_x.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| pred, delta, gate, mixed = model(node_x, global_x, baseline_long, baseline_short, adj) | |
| preds.append(pred.cpu().numpy()) | |
| gates.append(gate.cpu().numpy()) | |
| mixeds.append(mixed.cpu().numpy()) | |
| return np.vstack(preds), np.vstack(gates), np.vstack(mixeds) | |
| def _metrics_against_baselines(y_true: np.ndarray, pred_model: np.ndarray, pred_long: np.ndarray, pred_short: np.ndarray) -> dict[str, float]: | |
| return { | |
| "model_mae": base._mean_abs_error(y_true, pred_model), | |
| "season_baseline_mae": base._mean_abs_error(y_true, pred_long), | |
| "short8_baseline_mae": base._mean_abs_error(y_true, pred_short), | |
| "model_jsd": base._jsd_mean(y_true, pred_model), | |
| "season_baseline_jsd": base._jsd_mean(y_true, pred_long), | |
| "short8_baseline_jsd": base._jsd_mean(y_true, pred_short), | |
| "model_kl_proxy": float(np.mean(np.sum(y_true * (np.log(np.clip(y_true, 1e-8, 1.0)) - np.log(np.clip(pred_model, 1e-8, 1.0))), axis=1))), | |
| "season_baseline_kl_proxy": float(np.mean(np.sum(y_true * (np.log(np.clip(y_true, 1e-8, 1.0)) - np.log(np.clip(pred_long, 1e-8, 1.0))), axis=1))), | |
| "short8_baseline_kl_proxy": float(np.mean(np.sum(y_true * (np.log(np.clip(y_true, 1e-8, 1.0)) - np.log(np.clip(pred_short, 1e-8, 1.0))), axis=1))), | |
| } | |
| def _history_plot(history: list[dict[str, float]]) -> str: | |
| hist = pd.DataFrame(history) | |
| fig, ax = plt.subplots(figsize=(8, 4.2), facecolor="#F6F7F4") | |
| ax.plot(hist["epoch"], hist["train_loss"], label="Train", color="#2B7A5A", linewidth=2) | |
| ax.plot(hist["epoch"], hist["val_loss"], label="Validacion", color="#D1495B", linewidth=2) | |
| ax.set_xlabel("Epoch") | |
| ax.set_ylabel("Loss") | |
| ax.set_title("Curva de entrenamiento GNN", fontsize=13, fontweight="bold", color="#14342B") | |
| ax.grid(alpha=0.2) | |
| ax.legend(frameon=False) | |
| return _img_to_base64(fig) | |
| def _draw_pitch_distribution(ax: plt.Axes, values: dict[str, float], title: str) -> None: | |
| pitch = Pitch(pitch_type="opta", pitch_length=100, pitch_width=100, line_color="#D9E0DA", linewidth=1.2) | |
| pitch.draw(ax=ax) | |
| ax.set_facecolor("#F6F7F4") | |
| vmax = max(values.values()) if values else 1.0 | |
| norm = colors.Normalize(vmin=0.0, vmax=max(vmax, 1e-6)) | |
| cmap = plt.cm.Greens | |
| for zone, rects in ZONES_RECTS.items(): | |
| value = values.get(zone, 0.0) | |
| for x, y, w, h in rects: | |
| ax.add_patch( | |
| Rectangle((x, y), w, h, facecolor=cmap(norm(value)), edgecolor="#FFFFFF", linewidth=1.5, alpha=0.84, zorder=1) | |
| ) | |
| ax.text(x + w / 2, y + h / 2, f"{value * 100:.1f}%", ha="center", va="center", fontsize=8.5, fontweight="bold", color="#16352C", zorder=3) | |
| ax.set_title(title, fontsize=12, fontweight="bold", color="#14342B", pad=10) | |
| def _match_quad_image(row: pd.Series, attack_targets: list[str]) -> str: | |
| fig, axes = plt.subplots(1, 4, figsize=(16, 4.6), facecolor="#F6F7F4") | |
| fig.subplots_adjust(wspace=0.08) | |
| real = {zone: float(row[f"target_attack_share__{zone}"]) for zone in ZONE_ORDER} | |
| model = {zone: float(row[f"pred_model__target_attack_share__{zone}"]) for zone in ZONE_ORDER} | |
| season = {zone: float(row[f"pred_season__target_attack_share__{zone}"]) for zone in ZONE_ORDER} | |
| short8 = {zone: float(row[f"pred_short8__target_attack_share__{zone}"]) for zone in ZONE_ORDER} | |
| _draw_pitch_distribution(axes[0], real, "Real") | |
| _draw_pitch_distribution(axes[1], model, "Modelo GNN") | |
| _draw_pitch_distribution(axes[2], season, "Baseline temporada") | |
| _draw_pitch_distribution(axes[3], short8, "Baseline ultimos 8") | |
| fig.suptitle( | |
| f"{row['fecha'].strftime('%Y-%m-%d')} | {row.get('team_name', 'Equipo')} vs {row.get('opponent_name', 'Rival')}", | |
| fontsize=15, | |
| fontweight="bold", | |
| color="#14342B", | |
| y=1.02, | |
| ) | |
| return _img_to_base64(fig) | |
| def _report_html(summary: dict, test_metrics: dict[str, float], val_metrics: dict[str, float], history_img: str, racing_sections: list[str]) -> str: | |
| return f"""<!DOCTYPE html> | |
| <html lang="es"> | |
| <head> | |
| <meta charset="utf-8" /> | |
| <title>GNN residual para distribucion de ataque</title> | |
| <style> | |
| body {{ margin: 0; background: #f1f4ef; color: #14342B; font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif; }} | |
| .wrap {{ max-width: 1380px; margin: 0 auto; padding: 30px 24px 48px; }} | |
| h1 {{ margin: 0 0 10px; font-size: 40px; }} | |
| .lead {{ margin: 0 0 22px; font-size: 18px; color: #35574D; }} | |
| .hero {{ display: grid; grid-template-columns: repeat(4, minmax(0, 1fr)); gap: 14px; margin-bottom: 20px; }} | |
| .stat, .card, .match {{ background: white; border-radius: 20px; padding: 18px 20px; box-shadow: 0 8px 24px rgba(12, 36, 28, 0.08); }} | |
| .stat h3 {{ margin: 0 0 8px; font-size: 12px; text-transform: uppercase; letter-spacing: .08em; color: #587468; }} | |
| .stat p {{ margin: 0; font-size: 28px; font-weight: 800; }} | |
| table {{ width: 100%; border-collapse: collapse; font-size: 14px; }} | |
| th, td {{ padding: 10px 8px; border-bottom: 1px solid #E5ECE6; text-align: left; }} | |
| th {{ color: #587468; text-transform: uppercase; font-size: 12px; letter-spacing: .06em; }} | |
| img {{ width: 100%; border-radius: 16px; display: block; }} | |
| .grid2 {{ display: grid; grid-template-columns: repeat(2, minmax(0, 1fr)); gap: 16px; margin-bottom: 18px; }} | |
| .card {{ margin-bottom: 18px; }} | |
| .match {{ margin-bottom: 18px; }} | |
| @media (max-width: 980px) {{ .hero, .grid2 {{ grid-template-columns: 1fr; }} }} | |
| </style> | |
| </head> | |
| <body> | |
| <div class="wrap"> | |
| <h1>GNN residual para distribucion de ataque</h1> | |
| <p class="lead">Modelo de grafos con 10 nodos-zona. Cada nodo recibe contexto propio y se comunica con zonas adyacentes y espejo. La salida final es una distribucion valida de ataque, construida como correccion residual sobre una mezcla aprendida entre baseline de temporada y baseline de ultimos 8.</p> | |
| <section class="hero"> | |
| <div class="stat"><h3>Train</h3><p>{summary['train_rows']}</p></div> | |
| <div class="stat"><h3>Val</h3><p>{summary['val_rows']}</p></div> | |
| <div class="stat"><h3>Test</h3><p>{summary['test_rows']}</p></div> | |
| <div class="stat"><h3>Val Start</h3><p>{html.escape(summary['val_start_date'])}</p></div> | |
| </section> | |
| <div class="grid2"> | |
| <section class="card"> | |
| <h2>Validacion</h2> | |
| <table> | |
| <tr><th>Metrica</th><th>Modelo</th><th>Temporada</th><th>Ultimos 8</th></tr> | |
| <tr><td>MAE</td><td>{val_metrics['model_mae']:.4f}</td><td>{val_metrics['season_baseline_mae']:.4f}</td><td>{val_metrics['short8_baseline_mae']:.4f}</td></tr> | |
| <tr><td>JSD</td><td>{val_metrics['model_jsd']:.4f}</td><td>{val_metrics['season_baseline_jsd']:.4f}</td><td>{val_metrics['short8_baseline_jsd']:.4f}</td></tr> | |
| <tr><td>KL</td><td>{val_metrics['model_kl_proxy']:.4f}</td><td>{val_metrics['season_baseline_kl_proxy']:.4f}</td><td>{val_metrics['short8_baseline_kl_proxy']:.4f}</td></tr> | |
| </table> | |
| </section> | |
| <section class="card"> | |
| <h2>Test</h2> | |
| <table> | |
| <tr><th>Metrica</th><th>Modelo</th><th>Temporada</th><th>Ultimos 8</th></tr> | |
| <tr><td>MAE</td><td>{test_metrics['model_mae']:.4f}</td><td>{test_metrics['season_baseline_mae']:.4f}</td><td>{test_metrics['short8_baseline_mae']:.4f}</td></tr> | |
| <tr><td>JSD</td><td>{test_metrics['model_jsd']:.4f}</td><td>{test_metrics['season_baseline_jsd']:.4f}</td><td>{test_metrics['short8_baseline_jsd']:.4f}</td></tr> | |
| <tr><td>KL</td><td>{test_metrics['model_kl_proxy']:.4f}</td><td>{test_metrics['season_baseline_kl_proxy']:.4f}</td><td>{test_metrics['short8_baseline_kl_proxy']:.4f}</td></tr> | |
| </table> | |
| </section> | |
| </div> | |
| <section class="card"> | |
| <h2>Entrenamiento</h2> | |
| <img src="data:image/png;base64,{history_img}" alt="Curva de entrenamiento" /> | |
| </section> | |
| <section> | |
| <h2>Ultimos 3 partidos de Racing en test</h2> | |
| {''.join(racing_sections)} | |
| </section> | |
| </div> | |
| </body> | |
| </html>""" | |
| def main() -> None: | |
| _set_seed() | |
| MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| REPORTS_DIR.mkdir(parents=True, exist_ok=True) | |
| df, attack_targets = _load_data() | |
| y_attack = df[attack_targets].to_numpy(dtype=np.float32) | |
| train_idx, val_idx, test_idx, val_start_date = _train_val_test_indices(df) | |
| global_features, node_tensor, zone_order_array, global_cols, node_feat_names = _build_feature_matrices(df, attack_targets) | |
| global_x, global_bundle = _standardize_global(train_idx, global_features) | |
| node_x, node_bundle = _standardize_node(train_idx, node_tensor) | |
| baseline_long = base._normalize_rows(df[[f"long_mean__actual_attack_share__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float)).astype(np.float32) | |
| baseline_short = base._normalize_rows(df[[f"short_mean__actual_attack_share__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float)).astype(np.float32) | |
| train_split = _make_split(df, train_idx, global_x, node_x, y_attack, baseline_long, baseline_short) | |
| val_split = _make_split(df, val_idx, global_x, node_x, y_attack, baseline_long, baseline_short) | |
| test_split = _make_split(df, test_idx, global_x, node_x, y_attack, baseline_long, baseline_short) | |
| model, history = _train_model(train_split, val_split) | |
| val_pred, val_gates, val_mixed = _predict(model, val_split) | |
| test_pred, test_gates, test_mixed = _predict(model, test_split) | |
| val_true = val_split.metadata[attack_targets].to_numpy(dtype=float) | |
| test_true = test_split.metadata[attack_targets].to_numpy(dtype=float) | |
| val_long = val_split.baseline_long.astype(float) | |
| val_short = val_split.baseline_short.astype(float) | |
| test_long = test_split.baseline_long.astype(float) | |
| test_short = test_split.baseline_short.astype(float) | |
| val_metrics = _metrics_against_baselines(val_true, val_pred, val_long, val_short) | |
| test_metrics = _metrics_against_baselines(test_true, test_pred, test_long, test_short) | |
| pred_df = test_split.metadata.copy().reset_index(drop=True) | |
| for i, zone in enumerate(ZONE_ORDER): | |
| target_col = f"target_attack_share__{zone}" | |
| pred_df[f"pred_model__{target_col}"] = test_pred[:, i] | |
| pred_df[f"pred_season__{target_col}"] = test_long[:, i] | |
| pred_df[f"pred_short8__{target_col}"] = test_short[:, i] | |
| pred_df[f"gate__{zone}"] = test_gates[:, i] | |
| pred_df[f"mixed_base__{zone}"] = test_mixed[:, i] | |
| keep_cols = [ | |
| "matchId", "fecha", "league", "season", "teamId", "team_name", "opponent_name", "is_home", | |
| "goals_for", "goals_against", "n_prior_matches", "opp_n_prior_matches", | |
| ] + attack_targets | |
| keep_cols += [f"pred_model__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| keep_cols += [f"pred_season__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| keep_cols += [f"pred_short8__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| keep_cols += [f"gate__{zone}" for zone in ZONE_ORDER] | |
| keep_cols += [f"mixed_base__{zone}" for zone in ZONE_ORDER] | |
| pred_df[keep_cols].to_parquet(PRED_PATH, index=False) | |
| history_img = _history_plot(history) | |
| racing_last3 = pred_df[pred_df["teamId"] == base.RACING_TEAM_ID].sort_values(["fecha", "matchId"]).tail(3) | |
| racing_sections = [] | |
| for _, row in racing_last3.iterrows(): | |
| att_true = row[attack_targets].to_numpy(dtype=float) | |
| att_model = row[[f"pred_model__target_attack_share__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| att_season = row[[f"pred_season__target_attack_share__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| att_short = row[[f"pred_short8__target_attack_share__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| img = _match_quad_image(row, attack_targets) | |
| racing_sections.append( | |
| f""" | |
| <section class="match"> | |
| <h3>{html.escape(row['fecha'].strftime('%Y-%m-%d'))} | {html.escape(str(row.get('team_name', 'Equipo')))} vs {html.escape(str(row.get('opponent_name', 'Rival')))}</h3> | |
| <p>MAE modelo {base._mean_abs_error(att_true[None, :], att_model[None, :]):.4f} | baseline temporada {base._mean_abs_error(att_true[None, :], att_season[None, :]):.4f} | baseline ultimos 8 {base._mean_abs_error(att_true[None, :], att_short[None, :]):.4f}</p> | |
| <img src="data:image/png;base64,{img}" alt="Distribuciones de ataque" /> | |
| </section> | |
| """ | |
| ) | |
| summary = { | |
| "model": "residual_attack_gnn", | |
| "train_rows": int(len(train_split.metadata)), | |
| "val_rows": int(len(val_split.metadata)), | |
| "test_rows": int(len(test_split.metadata)), | |
| "val_start_date": val_start_date, | |
| "global_feature_count": int(train_split.global_x.shape[1]), | |
| "node_feature_count": int(train_split.node_x.shape[2]), | |
| "zone_order": ZONE_ORDER, | |
| "val_metrics": val_metrics, | |
| "test_metrics": test_metrics, | |
| "mean_gate_test": {zone: float(test_gates[:, i].mean()) for i, zone in enumerate(ZONE_ORDER)}, | |
| "racing_last3_test_matches": [ | |
| {"matchId": row["matchId"], "fecha": row["fecha"].strftime("%Y-%m-%d"), "opponent_name": row.get("opponent_name")} | |
| for _, row in racing_last3.iterrows() | |
| ], | |
| } | |
| torch.save( | |
| { | |
| "model_state_dict": model.state_dict(), | |
| "global_bundle": global_bundle, | |
| "node_bundle": node_bundle, | |
| "zone_order": ZONE_ORDER, | |
| "attack_targets": attack_targets, | |
| "summary": summary, | |
| }, | |
| MODEL_PATH, | |
| ) | |
| JSON_PATH.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8") | |
| REPORT_PATH.write_text(_report_html(summary, test_metrics, val_metrics, history_img, racing_sections), encoding="utf-8") | |
| print(f"Modelo guardado en: {MODEL_PATH}") | |
| print(f"Predicciones guardadas en: {PRED_PATH}") | |
| print(f"Metricas guardadas en: {JSON_PATH}") | |
| print(f"Reporte guardado en: {REPORT_PATH}") | |
| print(json.dumps(summary, ensure_ascii=False, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |