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_pv_interaction_gnn_report.html" | |
| JSON_PATH = MODEL_DIR / "attack_pv_interaction_gnn_metrics.json" | |
| MODEL_PATH = MODEL_DIR / "attack_pv_interaction_gnn_bundle.pt" | |
| ATTACK_PRED_PATH = MODEL_DIR / "attack_interaction_gnn_test_predictions.parquet" | |
| PV_PRED_PATH = MODEL_DIR / "pv_interaction_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", | |
| } | |
| ATTACK_PREV_JSON = MODEL_DIR / "attack_distribution_gnn_metrics.json" | |
| PV_PREV_JSON = MODEL_DIR / "pv_distribution_gnn_metrics.json" | |
| class TaskData: | |
| df: pd.DataFrame | |
| global_x: np.ndarray | |
| attack_node_x: np.ndarray | |
| defense_node_x: np.ndarray | |
| y_dist: np.ndarray | |
| baseline_long: np.ndarray | |
| baseline_short: np.ndarray | |
| train_idx: np.ndarray | |
| val_idx: np.ndarray | |
| test_idx: np.ndarray | |
| val_start_date: str | |
| task_name: str | |
| class SplitData: | |
| global_x: np.ndarray | |
| attack_node_x: np.ndarray | |
| defense_node_x: np.ndarray | |
| y_dist: 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_graphs() -> tuple[torch.Tensor, 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) | |
| zone_to_idx = {z: i for i, z in enumerate(ZONE_ORDER)} | |
| adj = np.zeros((n, n), dtype=np.float32) | |
| cross = np.eye(n, dtype=np.float32) | |
| edge_pairs: list[tuple[int, int]] = [] | |
| for z, neighs in neighbors.items(): | |
| i = zone_to_idx[z] | |
| for neigh in neighs: | |
| j = zone_to_idx[neigh] | |
| adj[i, j] = 1.0 | |
| cross[i, j] = 1.0 | |
| edge_pairs.append((i, j)) | |
| deg = np.where(adj.sum(axis=1, keepdims=True) > 0, adj.sum(axis=1, keepdims=True), 1.0) | |
| adj = adj / deg | |
| cross_deg = np.where(cross.sum(axis=1, keepdims=True) > 0, cross.sum(axis=1, keepdims=True), 1.0) | |
| cross = cross / cross_deg | |
| return torch.tensor(adj, dtype=torch.float32), torch.tensor(cross, dtype=torch.float32), torch.tensor(edge_pairs, dtype=torch.long) | |
| ADJ_MATRIX, CROSS_MATRIX, EDGE_INDEX = _normalized_graphs() | |
| def _standardize_global(global_features: pd.DataFrame, train_idx: np.ndarray) -> 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) | |
| return scaled, { | |
| "fill_values": fill_values.to_dict(), | |
| "means": means.to_dict(), | |
| "stds": stds.to_dict(), | |
| "global_feature_columns": list(global_features.columns), | |
| } | |
| def _standardize_nodes(node_tensor: np.ndarray, train_idx: 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, :, :] | |
| return scaled.astype(np.float32), { | |
| "fill_values": fill.tolist(), | |
| "means": means.tolist(), | |
| "stds": stds.tolist(), | |
| } | |
| def _make_split(task: TaskData, idx: np.ndarray) -> SplitData: | |
| return SplitData( | |
| global_x=task.global_x[idx].astype(np.float32), | |
| attack_node_x=task.attack_node_x[idx].astype(np.float32), | |
| defense_node_x=task.defense_node_x[idx].astype(np.float32), | |
| y_dist=task.y_dist[idx].astype(np.float32), | |
| baseline_long=task.baseline_long[idx].astype(np.float32), | |
| baseline_short=task.baseline_short[idx].astype(np.float32), | |
| metadata=task.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.attack_node_x), | |
| torch.from_numpy(split.defense_node_x), | |
| torch.from_numpy(split.y_dist), | |
| torch.from_numpy(split.baseline_long), | |
| torch.from_numpy(split.baseline_short), | |
| ) | |
| return DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=shuffle) | |
| class IntraGraphBlock(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 CrossGraphBlock(nn.Module): | |
| def __init__(self, src_dim: int, dst_dim: int, out_dim: int, dropout: float) -> None: | |
| super().__init__() | |
| self.dst_lin = nn.Linear(dst_dim, out_dim) | |
| self.src_lin = nn.Linear(src_dim, out_dim) | |
| self.norm = nn.LayerNorm(out_dim) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, dst: torch.Tensor, src: torch.Tensor, cross_adj: torch.Tensor) -> torch.Tensor: | |
| src_msg = torch.einsum("ij,bjf->bif", cross_adj, src) | |
| h = self.dst_lin(dst) + self.src_lin(src_msg) | |
| h = self.norm(h) | |
| h = torch.relu(h) | |
| return self.dropout(h) | |
| class InteractionDistributionGNN(nn.Module): | |
| def __init__(self, attack_node_dim: int, defense_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.attack_encoder = nn.Sequential(nn.Linear(attack_node_dim + global_hidden, hidden_dim), nn.ReLU()) | |
| self.defense_encoder = nn.Sequential(nn.Linear(defense_node_dim + global_hidden, hidden_dim), nn.ReLU()) | |
| self.attack_blocks = nn.ModuleList([IntraGraphBlock(hidden_dim, hidden_dim, 0.06) for _ in range(num_layers)]) | |
| self.defense_blocks = nn.ModuleList([IntraGraphBlock(hidden_dim, hidden_dim, 0.06) for _ in range(num_layers)]) | |
| self.cross_to_attack = nn.ModuleList([CrossGraphBlock(hidden_dim, hidden_dim, hidden_dim, 0.05) for _ in range(num_layers)]) | |
| self.cross_to_defense = nn.ModuleList([CrossGraphBlock(hidden_dim, hidden_dim, hidden_dim, 0.05) for _ in range(num_layers)]) | |
| self.gate_head = nn.Linear(hidden_dim * 2, 1) | |
| self.delta_head = nn.Linear(hidden_dim * 2, 1) | |
| def forward( | |
| self, | |
| attack_node_x: torch.Tensor, | |
| defense_node_x: torch.Tensor, | |
| global_x: torch.Tensor, | |
| baseline_long: torch.Tensor, | |
| baseline_short: torch.Tensor, | |
| adj: torch.Tensor, | |
| cross: torch.Tensor, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| g = self.global_encoder(global_x) | |
| g_rep = g.unsqueeze(1).expand(-1, attack_node_x.size(1), -1) | |
| att = self.attack_encoder(torch.cat([attack_node_x, g_rep], dim=-1)) | |
| dfn = self.defense_encoder(torch.cat([defense_node_x, g_rep], dim=-1)) | |
| for att_block, dfn_block, cross_att, cross_dfn in zip(self.attack_blocks, self.defense_blocks, self.cross_to_attack, self.cross_to_defense): | |
| att = att + att_block(att, adj) + cross_att(att, dfn, cross) | |
| dfn = dfn + dfn_block(dfn, adj) + cross_dfn(dfn, att, cross) | |
| pair = torch.cat([att, dfn], dim=-1) | |
| gate = torch.sigmoid(self.gate_head(pair)).squeeze(-1) | |
| mixed = gate * baseline_short + (1.0 - gate) * baseline_long | |
| delta = self.delta_head(pair).squeeze(-1) | |
| logits = torch.log(torch.clamp(mixed, min=1e-6)) + delta | |
| pred = torch.softmax(logits, dim=1) | |
| return pred, delta, gate, mixed | |
| def _distribution_loss(pred: torch.Tensor, target: torch.Tensor, delta: torch.Tensor, gate: torch.Tensor) -> tuple[torch.Tensor, dict[str, float]]: | |
| eps = 1e-8 | |
| kl = torch.nn.functional.kl_div(torch.log(torch.clamp(pred, min=eps)), target, reduction="batchmean") | |
| smooth = torch.mean((delta[:, EDGE_INDEX[:, 0]] - delta[:, EDGE_INDEX[:, 1]]) ** 2) | |
| 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) | |
| return loss, {"kl": float(kl.item()), "smooth": float(smooth.item()), "gate_entropy": float(gate_entropy.item())} | |
| def _train_model(task: TaskData) -> tuple[InteractionDistributionGNN, list[dict[str, float]]]: | |
| train = _make_split(task, task.train_idx) | |
| val = _make_split(task, task.val_idx) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model = InteractionDistributionGNN( | |
| attack_node_dim=train.attack_node_x.shape[2], | |
| defense_node_dim=train.defense_node_x.shape[2], | |
| global_dim=train.global_x.shape[1], | |
| ).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) | |
| cross = CROSS_MATRIX.to(device) | |
| for epoch in range(1, MAX_EPOCHS + 1): | |
| model.train() | |
| running = 0.0 | |
| n_batches = 0 | |
| for global_x, attack_node_x, defense_node_x, y_dist, baseline_long, baseline_short in train_loader: | |
| global_x = global_x.to(device) | |
| attack_node_x = attack_node_x.to(device) | |
| defense_node_x = defense_node_x.to(device) | |
| y_dist = y_dist.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| optimizer.zero_grad() | |
| pred, delta, gate, mixed = model(attack_node_x, defense_node_x, global_x, baseline_long, baseline_short, adj, cross) | |
| loss, parts = _distribution_loss(pred, y_dist, 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 _evaluate_loader(model: InteractionDistributionGNN, 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) | |
| cross = CROSS_MATRIX.to(device) | |
| with torch.no_grad(): | |
| for global_x, attack_node_x, defense_node_x, y_dist, baseline_long, baseline_short in loader: | |
| global_x = global_x.to(device) | |
| attack_node_x = attack_node_x.to(device) | |
| defense_node_x = defense_node_x.to(device) | |
| y_dist = y_dist.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| pred, delta, gate, mixed = model(attack_node_x, defense_node_x, global_x, baseline_long, baseline_short, adj, cross) | |
| loss, parts = _distribution_loss(pred, y_dist, 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 _predict(model: InteractionDistributionGNN, 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) | |
| cross = CROSS_MATRIX.to(device) | |
| with torch.no_grad(): | |
| for global_x, attack_node_x, defense_node_x, y_dist, baseline_long, baseline_short in loader: | |
| global_x = global_x.to(device) | |
| attack_node_x = attack_node_x.to(device) | |
| defense_node_x = defense_node_x.to(device) | |
| baseline_long = baseline_long.to(device) | |
| baseline_short = baseline_short.to(device) | |
| pred, delta, gate, mixed = model(attack_node_x, defense_node_x, global_x, baseline_long, baseline_short, adj, cross) | |
| 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(y_true: np.ndarray, pred_model: np.ndarray, pred_long: np.ndarray, pred_short: np.ndarray) -> dict[str, float]: | |
| def kl_proxy(y, p): | |
| return float(np.mean(np.sum(y * (np.log(np.clip(y, 1e-8, 1.0)) - np.log(np.clip(p, 1e-8, 1.0))), axis=1))) | |
| 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": kl_proxy(y_true, pred_model), | |
| "season_baseline_kl_proxy": kl_proxy(y_true, pred_long), | |
| "short8_baseline_kl_proxy": kl_proxy(y_true, pred_short), | |
| } | |
| def _load_prev_metrics(path: Path) -> dict[str, float]: | |
| return json.loads(path.read_text(encoding="utf-8"))["test_metrics"] | |
| def _history_plot(history: list[dict[str, float]], title: str) -> str: | |
| hist = pd.DataFrame(history) | |
| fig, ax = plt.subplots(figsize=(7.5, 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(title, 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 _task_match_quad(row: pd.Series, prefix_target: 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"{prefix_target}__{zone}"]) for zone in ZONE_ORDER} | |
| model = {zone: float(row[f"pred_model__{prefix_target}__{zone}"]) for zone in ZONE_ORDER} | |
| season = {zone: float(row[f"pred_season__{prefix_target}__{zone}"]) for zone in ZONE_ORDER} | |
| short8 = {zone: float(row[f"pred_short8__{prefix_target}__{zone}"]) for zone in ZONE_ORDER} | |
| _draw_pitch_distribution(axes[0], real, "Real") | |
| _draw_pitch_distribution(axes[1], model, "Interaccion") | |
| _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 _render_compare_table(title: str, current: dict[str, float], previous: dict[str, float], season: dict[str, float], short8: dict[str, float]) -> str: | |
| return f""" | |
| <section class="card"> | |
| <h3>{html.escape(title)}</h3> | |
| <table> | |
| <tr><th>Modelo</th><th>MAE</th><th>JSD</th><th>KL</th></tr> | |
| <tr><td>Interaccion GNN</td><td>{current['model_mae']:.4f}</td><td>{current['model_jsd']:.4f}</td><td>{current['model_kl_proxy']:.4f}</td></tr> | |
| <tr><td>GNN anterior</td><td>{previous['model_mae']:.4f}</td><td>{previous['model_jsd']:.4f}</td><td>{previous['model_kl_proxy']:.4f}</td></tr> | |
| <tr><td>Baseline temporada</td><td>{season['mae']:.4f}</td><td>{season['jsd']:.4f}</td><td>{season['kl']:.4f}</td></tr> | |
| <tr><td>Baseline ultimos 8</td><td>{short8['mae']:.4f}</td><td>{short8['jsd']:.4f}</td><td>{short8['kl']:.4f}</td></tr> | |
| </table> | |
| </section> | |
| """ | |
| def _build_report(summary: dict, attack_hist_img: str, pv_hist_img: str, attack_racing_sections: list[str], pv_racing_sections: list[str]) -> str: | |
| attack_test = summary["attack"]["test_metrics"] | |
| pv_test = summary["pv"]["test_metrics"] | |
| return f"""<!DOCTYPE html> | |
| <html lang="es"> | |
| <head> | |
| <meta charset="utf-8" /> | |
| <title>Interaccion ataque-defensa GNN</title> | |
| <style> | |
| body {{ margin: 0; background: #f1f4ef; color: #14342B; font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif; }} | |
| .wrap {{ max-width: 1400px; 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; }} | |
| .grid2 {{ display: grid; grid-template-columns: repeat(2, minmax(0, 1fr)); gap: 16px; margin-bottom: 18px; }} | |
| 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; }} | |
| .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 de interaccion ataque-defensa</h1> | |
| <p class="lead">Cada zona tiene dos representaciones: una ofensiva propia y una defensiva del rival. Hay message passing dentro de cada grafo y tambien entre ambos grafos, para modelar explicitamente el matchup zona a zona. Se compara contra los dos baselines y contra la GNN anterior.</p> | |
| <section class="hero"> | |
| <div class="stat"><h3>Attack Test MAE</h3><p>{attack_test['model_mae']:.4f}</p></div> | |
| <div class="stat"><h3>PV Test MAE</h3><p>{pv_test['model_mae']:.4f}</p></div> | |
| <div class="stat"><h3>Attack Mejora vs Prev</h3><p>{summary['attack']['previous_test_metrics']['model_mae'] - attack_test['model_mae']:+.4f}</p></div> | |
| <div class="stat"><h3>PV Mejora vs Prev</h3><p>{summary['pv']['previous_test_metrics']['model_mae'] - pv_test['model_mae']:+.4f}</p></div> | |
| </section> | |
| <div class="grid2"> | |
| {_render_compare_table( | |
| "Ataque - test", | |
| summary["attack"]["test_metrics"], | |
| summary["attack"]["previous_test_metrics"], | |
| {"mae": summary["attack"]["test_metrics"]["season_baseline_mae"], "jsd": summary["attack"]["test_metrics"]["season_baseline_jsd"], "kl": summary["attack"]["test_metrics"]["season_baseline_kl_proxy"]}, | |
| {"mae": summary["attack"]["test_metrics"]["short8_baseline_mae"], "jsd": summary["attack"]["test_metrics"]["short8_baseline_jsd"], "kl": summary["attack"]["test_metrics"]["short8_baseline_kl_proxy"]}, | |
| )} | |
| {_render_compare_table( | |
| "PV - test", | |
| summary["pv"]["test_metrics"], | |
| summary["pv"]["previous_test_metrics"], | |
| {"mae": summary["pv"]["test_metrics"]["season_baseline_mae"], "jsd": summary["pv"]["test_metrics"]["season_baseline_jsd"], "kl": summary["pv"]["test_metrics"]["season_baseline_kl_proxy"]}, | |
| {"mae": summary["pv"]["test_metrics"]["short8_baseline_mae"], "jsd": summary["pv"]["test_metrics"]["short8_baseline_jsd"], "kl": summary["pv"]["test_metrics"]["short8_baseline_kl_proxy"]}, | |
| )} | |
| </div> | |
| <div class="grid2"> | |
| <section class="card"> | |
| <h3>Ataque - entrenamiento</h3> | |
| <img src="data:image/png;base64,{attack_hist_img}" alt="Historial ataque" /> | |
| </section> | |
| <section class="card"> | |
| <h3>PV - entrenamiento</h3> | |
| <img src="data:image/png;base64,{pv_hist_img}" alt="Historial PV" /> | |
| </section> | |
| </div> | |
| <section> | |
| <h2>Racing - Ataque</h2> | |
| {''.join(attack_racing_sections)} | |
| </section> | |
| <section> | |
| <h2>Racing - PV</h2> | |
| {''.join(pv_racing_sections)} | |
| </section> | |
| </div> | |
| </body> | |
| </html>""" | |
| def _build_attack_task(df_base: pd.DataFrame) -> TaskData: | |
| df = df_base[df_base["usable_for_model"]].copy().reset_index(drop=True) | |
| attack_targets = [f"target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| y_dist = df[attack_targets].to_numpy(dtype=np.float32) | |
| 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) | |
| attack_node_cols = [ | |
| [f"short_mean__actual_attack_share__{zone}", f"long_mean__actual_attack_share__{zone}", f"short_mean__actual_pv_share__{zone}", f"long_mean__actual_pv_share__{zone}"] | |
| for zone in ZONE_ORDER | |
| ] | |
| defense_node_cols = [ | |
| [f"opp__short_mean__actual_conceded_attack_share__{zone}", f"opp__long_mean__actual_conceded_attack_share__{zone}", f"opp__short_mean__actual_conceded_pv_share__{zone}", f"opp__long_mean__actual_conceded_pv_share__{zone}"] | |
| for zone in ZONE_ORDER | |
| ] | |
| attack_node_x = np.stack([df[cols].apply(pd.to_numeric, errors="coerce").to_numpy(dtype=float) for cols in attack_node_cols], axis=1) | |
| defense_node_x = np.stack([df[cols].apply(pd.to_numeric, errors="coerce").to_numpy(dtype=float) for cols in defense_node_cols], axis=1) | |
| wide_features, _, _, _, _ = base._feature_matrix(df) | |
| node_cols = {c for cols in attack_node_cols + defense_node_cols for c in cols} | set(attack_targets) | |
| global_features = wide_features[[c for c in wide_features.columns if c not in node_cols]].copy() | |
| train_idx_pd, val_idx_pd, val_start_date = base._build_temporal_validation(df[df["split"] == "train"].copy()) | |
| train_idx = train_idx_pd.to_numpy() | |
| val_idx = val_idx_pd.to_numpy() | |
| test_idx = df.index[df["split"] == "test"].to_numpy() | |
| global_x, _ = _standardize_global(global_features, train_idx) | |
| attack_node_x, _ = _standardize_nodes(attack_node_x, train_idx) | |
| defense_node_x, _ = _standardize_nodes(defense_node_x, train_idx) | |
| return TaskData(df, global_x, attack_node_x, defense_node_x, y_dist, baseline_long, baseline_short, train_idx, val_idx, test_idx, val_start_date, "attack") | |
| def _clip_and_normalize(mat: np.ndarray) -> np.ndarray: | |
| clipped = np.clip(mat, 0.0, None) | |
| total = clipped.sum(axis=1, keepdims=True) | |
| out = np.zeros_like(clipped, dtype=np.float32) | |
| mask = total.squeeze(-1) > 0 | |
| if mask.any(): | |
| out[mask] = (clipped[mask] / total[mask]).astype(np.float32) | |
| return out | |
| def _build_pv_task(df_base: pd.DataFrame) -> TaskData: | |
| df = df_base.copy() | |
| zone_cols = [f"zone_pvAdded__{zone}" for zone in ZONE_ORDER] | |
| df["pv_positive_total"] = np.clip(df[zone_cols].to_numpy(dtype=float), 0, None).sum(axis=1) | |
| df = df[df["usable_for_model"] & (df["pv_positive_total"] > 0)].copy().reset_index(drop=True) | |
| y_dist = _clip_and_normalize(df[[f"zone_pvAdded__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float)) | |
| baseline_long = _clip_and_normalize(df[[f"long_mean__zone_pvAdded__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float)) | |
| baseline_short = _clip_and_normalize(df[[f"short_mean__zone_pvAdded__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float)) | |
| attack_node_cols = [ | |
| [f"short_mean__actual_pv_share__{zone}", f"long_mean__actual_pv_share__{zone}", f"short_mean__actual_attack_share__{zone}", f"long_mean__actual_attack_share__{zone}"] | |
| for zone in ZONE_ORDER | |
| ] | |
| defense_node_cols = [ | |
| [f"opp__short_mean__actual_conceded_pv_share__{zone}", f"opp__long_mean__actual_conceded_pv_share__{zone}", f"opp__short_mean__actual_conceded_attack_share__{zone}", f"opp__long_mean__actual_conceded_attack_share__{zone}"] | |
| for zone in ZONE_ORDER | |
| ] | |
| attack_node_x = np.stack([df[cols].apply(pd.to_numeric, errors="coerce").to_numpy(dtype=float) for cols in attack_node_cols], axis=1) | |
| defense_node_x = np.stack([df[cols].apply(pd.to_numeric, errors="coerce").to_numpy(dtype=float) for cols in defense_node_cols], axis=1) | |
| wide_features, _, _, _, _ = base._feature_matrix(df) | |
| node_cols = {c for cols in attack_node_cols + defense_node_cols for c in cols} | |
| global_features = wide_features[[c for c in wide_features.columns if c not in node_cols]].copy() | |
| train_idx_pd, val_idx_pd, val_start_date = base._build_temporal_validation(df[df["split"] == "train"].copy()) | |
| train_idx = train_idx_pd.to_numpy() | |
| val_idx = val_idx_pd.to_numpy() | |
| test_idx = df.index[df["split"] == "test"].to_numpy() | |
| global_x, _ = _standardize_global(global_features, train_idx) | |
| attack_node_x, _ = _standardize_nodes(attack_node_x, train_idx) | |
| defense_node_x, _ = _standardize_nodes(defense_node_x, train_idx) | |
| return TaskData(df, global_x, attack_node_x, defense_node_x, y_dist, baseline_long, baseline_short, train_idx, val_idx, test_idx, val_start_date, "pv") | |
| def _run_task(task: TaskData) -> dict: | |
| model, history = _train_model(task) | |
| test_split = _make_split(task, task.test_idx) | |
| val_split = _make_split(task, task.val_idx) | |
| test_pred, test_gates, test_mixed = _predict(model, test_split) | |
| val_pred, val_gates, val_mixed = _predict(model, val_split) | |
| test_metrics = _metrics(test_split.y_dist.astype(float), test_pred, test_split.baseline_long.astype(float), test_split.baseline_short.astype(float)) | |
| val_metrics = _metrics(val_split.y_dist.astype(float), val_pred, val_split.baseline_long.astype(float), val_split.baseline_short.astype(float)) | |
| return { | |
| "model_state_dict": {k: v.detach().cpu() for k, v in model.state_dict().items()}, | |
| "history": history, | |
| "test_split": test_split, | |
| "val_split": val_split, | |
| "test_pred": test_pred, | |
| "val_pred": val_pred, | |
| "test_gates": test_gates, | |
| "test_mixed": test_mixed, | |
| "test_metrics": test_metrics, | |
| "val_metrics": val_metrics, | |
| } | |
| def _attach_predictions(split: SplitData, pred: np.ndarray, gates: np.ndarray, mixed: np.ndarray, target_prefix: str) -> pd.DataFrame: | |
| out = split.metadata.copy().reset_index(drop=True) | |
| for i, zone in enumerate(ZONE_ORDER): | |
| out[f"{target_prefix}__{zone}"] = split.y_dist[:, i] | |
| out[f"pred_model__{target_prefix}__{zone}"] = pred[:, i] | |
| out[f"pred_season__{target_prefix}__{zone}"] = split.baseline_long[:, i] | |
| out[f"pred_short8__{target_prefix}__{zone}"] = split.baseline_short[:, i] | |
| out[f"gate__{zone}"] = gates[:, i] | |
| out[f"mixed_base__{zone}"] = mixed[:, i] | |
| return out | |
| def _racing_sections(df: pd.DataFrame, target_prefix: str) -> list[str]: | |
| racing = df[df["teamId"] == base.RACING_TEAM_ID].sort_values(["fecha", "matchId"]).tail(3) | |
| sections = [] | |
| for _, row in racing.iterrows(): | |
| real = row[[f"{target_prefix}__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| model = row[[f"pred_model__{target_prefix}__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| season = row[[f"pred_season__{target_prefix}__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| short8 = row[[f"pred_short8__{target_prefix}__{zone}" for zone in ZONE_ORDER]].to_numpy(dtype=float) | |
| 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 interaccion {base._mean_abs_error(real[None, :], model[None, :]):.4f} | baseline temporada {base._mean_abs_error(real[None, :], season[None, :]):.4f} | baseline ultimos 8 {base._mean_abs_error(real[None, :], short8[None, :]):.4f}</p> | |
| <img src="data:image/png;base64,{_task_match_quad(row, target_prefix)}" alt="Distribuciones" /> | |
| </section> | |
| """ | |
| ) | |
| return sections | |
| def main() -> None: | |
| _set_seed() | |
| MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| REPORTS_DIR.mkdir(parents=True, exist_ok=True) | |
| df_base = base._load_dataset() | |
| attack_task = _build_attack_task(df_base) | |
| pv_task = _build_pv_task(df_base) | |
| attack_result = _run_task(attack_task) | |
| pv_result = _run_task(pv_task) | |
| attack_prev = _load_prev_metrics(ATTACK_PREV_JSON) | |
| pv_prev = _load_prev_metrics(PV_PREV_JSON) | |
| attack_pred_df = _attach_predictions(attack_result["test_split"], attack_result["test_pred"], attack_result["test_gates"], attack_result["test_mixed"], "target_attack_share") | |
| pv_pred_df = _attach_predictions(pv_result["test_split"], pv_result["test_pred"], pv_result["test_gates"], pv_result["test_mixed"], "target_pv_dist") | |
| attack_keep = ["matchId", "fecha", "league", "season", "teamId", "team_name", "opponent_name", "is_home", "goals_for", "goals_against", "n_prior_matches", "opp_n_prior_matches"] | |
| attack_keep += [f"target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| attack_keep += [f"pred_model__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| attack_keep += [f"pred_season__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| attack_keep += [f"pred_short8__target_attack_share__{zone}" for zone in ZONE_ORDER] | |
| attack_keep += [f"gate__{zone}" for zone in ZONE_ORDER] | |
| attack_keep += [f"mixed_base__{zone}" for zone in ZONE_ORDER] | |
| attack_pred_df[attack_keep].to_parquet(ATTACK_PRED_PATH, index=False) | |
| pv_keep = ["matchId", "fecha", "league", "season", "teamId", "team_name", "opponent_name", "is_home", "goals_for", "goals_against", "n_prior_matches", "opp_n_prior_matches"] | |
| pv_keep += [f"target_pv_dist__{zone}" for zone in ZONE_ORDER] | |
| pv_keep += [f"pred_model__target_pv_dist__{zone}" for zone in ZONE_ORDER] | |
| pv_keep += [f"pred_season__target_pv_dist__{zone}" for zone in ZONE_ORDER] | |
| pv_keep += [f"pred_short8__target_pv_dist__{zone}" for zone in ZONE_ORDER] | |
| pv_keep += [f"gate__{zone}" for zone in ZONE_ORDER] | |
| pv_keep += [f"mixed_base__{zone}" for zone in ZONE_ORDER] | |
| pv_pred_df[pv_keep].to_parquet(PV_PRED_PATH, index=False) | |
| summary = { | |
| "attack": { | |
| "train_rows": int(len(attack_task.train_idx)), | |
| "val_rows": int(len(attack_task.val_idx)), | |
| "test_rows": int(len(attack_task.test_idx)), | |
| "val_start_date": attack_task.val_start_date, | |
| "test_metrics": attack_result["test_metrics"], | |
| "val_metrics": attack_result["val_metrics"], | |
| "previous_test_metrics": attack_prev, | |
| "mean_gate_test": {zone: float(attack_result["test_gates"][:, i].mean()) for i, zone in enumerate(ZONE_ORDER)}, | |
| }, | |
| "pv": { | |
| "train_rows": int(len(pv_task.train_idx)), | |
| "val_rows": int(len(pv_task.val_idx)), | |
| "test_rows": int(len(pv_task.test_idx)), | |
| "val_start_date": pv_task.val_start_date, | |
| "test_metrics": pv_result["test_metrics"], | |
| "val_metrics": pv_result["val_metrics"], | |
| "previous_test_metrics": pv_prev, | |
| "mean_gate_test": {zone: float(pv_result["test_gates"][:, i].mean()) for i, zone in enumerate(ZONE_ORDER)}, | |
| }, | |
| } | |
| torch.save( | |
| { | |
| "attack_model_state_dict": attack_result["model_state_dict"], | |
| "pv_model_state_dict": pv_result["model_state_dict"], | |
| "zone_order": ZONE_ORDER, | |
| "summary": summary, | |
| }, | |
| MODEL_PATH, | |
| ) | |
| JSON_PATH.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8") | |
| report_html = _build_report( | |
| summary, | |
| _history_plot(attack_result["history"], "Ataque - interaccion"), | |
| _history_plot(pv_result["history"], "PV - interaccion"), | |
| _racing_sections(attack_pred_df, "target_attack_share"), | |
| _racing_sections(pv_pred_df, "target_pv_dist"), | |
| ) | |
| REPORT_PATH.write_text(report_html, encoding="utf-8") | |
| print(f"Modelo guardado en: {MODEL_PATH}") | |
| print(f"Predicciones ataque guardadas en: {ATTACK_PRED_PATH}") | |
| print(f"Predicciones PV guardadas en: {PV_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() | |