from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import html
import json
import random
import sys
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 / "zone_transformer_report.html"
JSON_PATH = MODEL_DIR / "zone_transformer_metrics.json"
MODEL_PATH = MODEL_DIR / "zone_transformer_bundle.pt"
ATTACK_PRED_PATH = MODEL_DIR / "attack_zone_transformer_test_predictions.parquet"
PV_PRED_PATH = MODEL_DIR / "pv_zone_transformer_test_predictions.parquet"
if str(SCRIPT_DIR) not in sys.path:
sys.path.insert(0, str(SCRIPT_DIR))
import experiment_attack_distribution_gnn as attack_base # noqa: E402
import experiment_pv_distribution_gnn as pv_base # noqa: E402
import train_attack_prediction_ffn as base # noqa: E402
RANDOM_SEED = 42
BATCH_SIZE = 256
MAX_EPOCHS = 260
PATIENCE = 32
LEARNING_RATE = 3e-4
WEIGHT_DECAY = 1e-5
@dataclass
class SplitData:
global_x: np.ndarray
node_x: np.ndarray
y: 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)
class ZoneTransformer(nn.Module):
def __init__(
self,
node_dim: int,
global_dim: int,
num_zones: int,
hidden_dim: int = 96,
global_hidden: int = 96,
num_layers: int = 3,
num_heads: int = 4,
dropout: float = 0.10,
) -> None:
super().__init__()
self.global_encoder = nn.Sequential(
nn.Linear(global_dim, 192),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(192, global_hidden),
nn.ReLU(),
)
self.node_encoder = nn.Sequential(
nn.Linear(node_dim + global_hidden, hidden_dim),
nn.ReLU(),
)
self.pos_embedding = nn.Parameter(torch.zeros(1, num_zones, hidden_dim))
encoder_layer = nn.TransformerEncoderLayer(
d_model=hidden_dim,
nhead=num_heads,
dim_feedforward=hidden_dim * 4,
dropout=dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
self.norm = nn.LayerNorm(hidden_dim)
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,
) -> 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))
h = h + self.pos_embedding
h = self.transformer(h)
h = self.norm(h)
gate = torch.sigmoid(self.gate_head(h)).squeeze(-1)
mixed = gate * baseline_short + (1.0 - gate) * baseline_long
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 _make_split(
df: pd.DataFrame,
idx: np.ndarray,
global_x: np.ndarray,
node_x: np.ndarray,
y: 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=y[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),
torch.from_numpy(split.baseline_long),
torch.from_numpy(split.baseline_short),
)
return DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=shuffle)
def _evaluate_attack_loader(model: ZoneTransformer, loader: DataLoader, device: torch.device) -> dict[str, float]:
model.eval()
total = 0.0
total_kl = 0.0
n_batches = 0
with torch.no_grad():
for global_x, node_x, y, baseline_long, baseline_short in loader:
pred, delta, gate, _mixed = model(
node_x.to(device),
global_x.to(device),
baseline_long.to(device),
baseline_short.to(device),
)
loss, parts = attack_base._distribution_loss(pred, y.to(device), 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 _evaluate_pv_loader(model: ZoneTransformer, loader: DataLoader, device: torch.device) -> dict[str, float]:
model.eval()
total = 0.0
total_kl = 0.0
n_batches = 0
with torch.no_grad():
for global_x, node_x, y, baseline_long, baseline_short in loader:
pred, delta, gate, _mixed = model(
node_x.to(device),
global_x.to(device),
baseline_long.to(device),
baseline_short.to(device),
)
loss, parts = pv_base._dist_loss(pred, y.to(device), 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_attack(train: SplitData, val: SplitData, node_dim: int, global_dim: int, num_zones: int) -> tuple[ZoneTransformer, list[dict[str, float]]]:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = ZoneTransformer(node_dim=node_dim, global_dim=global_dim, num_zones=num_zones).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]] = []
for epoch in range(1, MAX_EPOCHS + 1):
model.train()
running = 0.0
n_batches = 0
for global_x, node_x, y, baseline_long, baseline_short in train_loader:
optimizer.zero_grad()
pred, delta, gate, _mixed = model(
node_x.to(device),
global_x.to(device),
baseline_long.to(device),
baseline_short.to(device),
)
loss, _parts = attack_base._distribution_loss(pred, y.to(device), delta, gate)
loss.backward()
optimizer.step()
running += float(loss.item())
n_batches += 1
val_metrics = _evaluate_attack_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 _train_pv(train: SplitData, val: SplitData, node_dim: int, global_dim: int, num_zones: int) -> tuple[ZoneTransformer, list[dict[str, float]]]:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = ZoneTransformer(node_dim=node_dim, global_dim=global_dim, num_zones=num_zones).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]] = []
for epoch in range(1, MAX_EPOCHS + 1):
model.train()
running = 0.0
n_batches = 0
for global_x, node_x, y, baseline_long, baseline_short in train_loader:
optimizer.zero_grad()
pred, delta, gate, _mixed = model(
node_x.to(device),
global_x.to(device),
baseline_long.to(device),
baseline_short.to(device),
)
loss, _parts = pv_base._dist_loss(pred, y.to(device), delta, gate)
loss.backward()
optimizer.step()
running += float(loss.item())
n_batches += 1
val_metrics = _evaluate_pv_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: ZoneTransformer, split: SplitData) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.eval()
preds, gates, mixeds = [], [], []
loader = _make_loader(split, shuffle=False)
with torch.no_grad():
for global_x, node_x, _y, baseline_long, baseline_short in loader:
pred, _delta, gate, mixed = model(
node_x.to(device),
global_x.to(device),
baseline_long.to(device),
baseline_short.to(device),
)
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 _mae_rows(df: pd.DataFrame, target_prefix: str) -> pd.DataFrame:
target_cols = [c for c in df.columns if c.startswith(target_prefix)]
model_cols = [c.replace(target_prefix, f"pred_model__{target_prefix}") for c in target_cols]
season_cols = [c.replace(target_prefix, f"pred_season__{target_prefix}") for c in target_cols]
short8_cols = [c.replace(target_prefix, f"pred_short8__{target_prefix}") for c in target_cols]
out = df[["fecha", "team_name", "opponent_name"]].copy()
out["model_mae"] = np.abs(df[target_cols].to_numpy() - df[model_cols].to_numpy()).mean(axis=1)
out["season_mae"] = np.abs(df[target_cols].to_numpy() - df[season_cols].to_numpy()).mean(axis=1)
out["short8_mae"] = np.abs(df[target_cols].to_numpy() - df[short8_cols].to_numpy()).mean(axis=1)
return out
def _last3_rows(df: pd.DataFrame) -> list[dict[str, float | str]]:
sub = df[df["team_name"].eq("Racing de Santander")].sort_values("fecha").tail(3)
rows = sub.to_dict(orient="records")
for row in rows:
if "fecha" in row and pd.notna(row["fecha"]):
row["fecha"] = pd.Timestamp(row["fecha"]).strftime("%Y-%m-%d")
return rows
def _report_html(summary: dict) -> str:
def rows_html(rows: list[dict[str, float | str]]) -> str:
return "".join(
f"
| {html.escape(str(r['fecha']))} | {html.escape(str(r['opponent_name']))} | "
f"{r['model_mae']:.4f} | {r['season_mae']:.4f} | {r['short8_mae']:.4f} | {r['previous_model_mae']:.4f} |
"
for r in rows
)
return f"""
Zone Transformer
Transformer por zonas
Cada zona funciona como un token y el modelo aprende relaciones libres entre zonas mediante self-attention. La salida sigue siendo residual sobre baseline temporada y ultimos 8 para mantener una comparacion justa con el GNN.
Ataque - test
| Metrica | Transformer | Temporada | Ultimos 8 | GNN previo |
| MAE | {summary['attack']['test_metrics']['model_mae']:.4f} | {summary['attack']['test_metrics']['season_baseline_mae']:.4f} | {summary['attack']['test_metrics']['short8_baseline_mae']:.4f} | {summary['attack']['previous_test_metrics']['model_mae']:.4f} |
| JSD | {summary['attack']['test_metrics']['model_jsd']:.4f} | {summary['attack']['test_metrics']['season_baseline_jsd']:.4f} | {summary['attack']['test_metrics']['short8_baseline_jsd']:.4f} | {summary['attack']['previous_test_metrics']['model_jsd']:.4f} |
| KL | {summary['attack']['test_metrics']['model_kl_proxy']:.4f} | {summary['attack']['test_metrics']['season_baseline_kl_proxy']:.4f} | {summary['attack']['test_metrics']['short8_baseline_kl_proxy']:.4f} | {summary['attack']['previous_test_metrics']['model_kl_proxy']:.4f} |
PV - test
| Metrica | Transformer | Temporada | Ultimos 8 | GNN previo |
| MAE | {summary['pv']['test_metrics']['model_mae']:.4f} | {summary['pv']['test_metrics']['season_baseline_mae']:.4f} | {summary['pv']['test_metrics']['short8_baseline_mae']:.4f} | {summary['pv']['previous_test_metrics']['model_mae']:.4f} |
| JSD | {summary['pv']['test_metrics']['model_jsd']:.4f} | {summary['pv']['test_metrics']['season_baseline_jsd']:.4f} | {summary['pv']['test_metrics']['short8_baseline_jsd']:.4f} | {summary['pv']['previous_test_metrics']['model_jsd']:.4f} |
| KL | {summary['pv']['test_metrics']['model_kl_proxy']:.4f} | {summary['pv']['test_metrics']['season_baseline_kl_proxy']:.4f} | {summary['pv']['test_metrics']['short8_baseline_kl_proxy']:.4f} | {summary['pv']['previous_test_metrics']['model_kl_proxy']:.4f} |
Ataque - ultimos 3 de Racing
| Fecha | Rival | Transformer | Temporada | Ultimos 8 | GNN previo |
{rows_html(summary['attack']['last3_racing'])}
PV - ultimos 3 de Racing
| Fecha | Rival | Transformer | Temporada | Ultimos 8 | GNN previo |
{rows_html(summary['pv']['last3_racing'])}
"""
def main() -> None:
_set_seed()
MODEL_DIR.mkdir(parents=True, exist_ok=True)
REPORTS_DIR.mkdir(parents=True, exist_ok=True)
attack_df, attack_targets = attack_base._load_data()
attack_train_idx, attack_val_idx, attack_test_idx, attack_val_start_date = attack_base._train_val_test_indices(attack_df)
attack_global_features, attack_node_tensor, _zone_order_array, attack_global_cols, attack_node_feat_names = attack_base._build_feature_matrices(attack_df, attack_targets)
attack_global_x, attack_global_bundle = attack_base._standardize_global(attack_train_idx, attack_global_features)
attack_node_x, attack_node_bundle = attack_base._standardize_node(attack_train_idx, attack_node_tensor)
attack_y = attack_df[attack_targets].to_numpy(dtype=np.float32)
attack_baseline_long = base._normalize_rows(attack_df[[f"long_mean__actual_attack_share__{zone}" for zone in attack_base.ZONE_ORDER]].to_numpy(dtype=float)).astype(np.float32)
attack_baseline_short = base._normalize_rows(attack_df[[f"short_mean__actual_attack_share__{zone}" for zone in attack_base.ZONE_ORDER]].to_numpy(dtype=float)).astype(np.float32)
attack_train = _make_split(attack_df, attack_train_idx, attack_global_x, attack_node_x, attack_y, attack_baseline_long, attack_baseline_short)
attack_val = _make_split(attack_df, attack_val_idx, attack_global_x, attack_node_x, attack_y, attack_baseline_long, attack_baseline_short)
attack_test = _make_split(attack_df, attack_test_idx, attack_global_x, attack_node_x, attack_y, attack_baseline_long, attack_baseline_short)
attack_model, attack_history = _train_attack(attack_train, attack_val, attack_train.node_x.shape[2], attack_train.global_x.shape[1], len(attack_base.ZONE_ORDER))
attack_val_pred, _attack_val_gates, _attack_val_mixed = _predict(attack_model, attack_val)
attack_test_pred, attack_test_gates, attack_test_mixed = _predict(attack_model, attack_test)
attack_val_metrics = attack_base._metrics_against_baselines(attack_val.y.astype(float), attack_val_pred, attack_val.baseline_long.astype(float), attack_val.baseline_short.astype(float))
attack_test_metrics = attack_base._metrics_against_baselines(attack_test.y.astype(float), attack_test_pred, attack_test.baseline_long.astype(float), attack_test.baseline_short.astype(float))
attack_pred_df = attack_test.metadata.copy().reset_index(drop=True)
for i, zone in enumerate(attack_base.ZONE_ORDER):
attack_pred_df[f"target_attack_share__{zone}"] = attack_test.y[:, i]
attack_pred_df[f"pred_model__target_attack_share__{zone}"] = attack_test_pred[:, i]
attack_pred_df[f"pred_season__target_attack_share__{zone}"] = attack_test.baseline_long[:, i]
attack_pred_df[f"pred_short8__target_attack_share__{zone}"] = attack_test.baseline_short[:, i]
attack_pred_df[f"gate__{zone}"] = attack_test_gates[:, i]
attack_pred_df[f"mixed_base__{zone}"] = attack_test_mixed[:, i]
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 attack_base.ZONE_ORDER]
attack_keep += [f"pred_model__target_attack_share__{zone}" for zone in attack_base.ZONE_ORDER]
attack_keep += [f"pred_season__target_attack_share__{zone}" for zone in attack_base.ZONE_ORDER]
attack_keep += [f"pred_short8__target_attack_share__{zone}" for zone in attack_base.ZONE_ORDER]
attack_keep += [f"gate__{zone}" for zone in attack_base.ZONE_ORDER]
attack_keep += [f"mixed_base__{zone}" for zone in attack_base.ZONE_ORDER]
attack_pred_df[attack_keep].to_parquet(ATTACK_PRED_PATH, index=False)
pv_df = pv_base._load_data()
pv_train_idx, pv_val_idx, pv_test_idx, pv_val_start_date = pv_base._train_val_test_indices(pv_df)
pv_y, pv_baseline_long, pv_baseline_short = pv_base._build_distributions(pv_df)
pv_global_features, pv_node_tensor, pv_node_feature_names = pv_base._build_feature_matrices(pv_df)
pv_global_x, pv_global_bundle = pv_base._standardize_global(pv_train_idx, pv_global_features)
pv_node_x, pv_node_bundle = pv_base._standardize_node(pv_train_idx, pv_node_tensor)
pv_train = _make_split(pv_df, pv_train_idx, pv_global_x, pv_node_x, pv_y, pv_baseline_long, pv_baseline_short)
pv_val = _make_split(pv_df, pv_val_idx, pv_global_x, pv_node_x, pv_y, pv_baseline_long, pv_baseline_short)
pv_test = _make_split(pv_df, pv_test_idx, pv_global_x, pv_node_x, pv_y, pv_baseline_long, pv_baseline_short)
pv_model, pv_history = _train_pv(pv_train, pv_val, pv_train.node_x.shape[2], pv_train.global_x.shape[1], len(pv_base.ZONE_ORDER))
pv_val_pred, _pv_val_gates, _pv_val_mixed = _predict(pv_model, pv_val)
pv_test_pred, pv_test_gates, pv_test_mixed = _predict(pv_model, pv_test)
pv_val_metrics = pv_base._metrics(pv_val.y.astype(float), pv_val_pred, pv_val.baseline_long.astype(float), pv_val.baseline_short.astype(float))
pv_test_metrics = pv_base._metrics(pv_test.y.astype(float), pv_test_pred, pv_test.baseline_long.astype(float), pv_test.baseline_short.astype(float))
pv_pred_df = pv_test.metadata.copy().reset_index(drop=True)
for i, zone in enumerate(pv_base.ZONE_ORDER):
pv_pred_df[f"target_pv_dist__{zone}"] = pv_test.y[:, i]
pv_pred_df[f"pred_model__target_pv_dist__{zone}"] = pv_test_pred[:, i]
pv_pred_df[f"pred_season__target_pv_dist__{zone}"] = pv_test.baseline_long[:, i]
pv_pred_df[f"pred_short8__target_pv_dist__{zone}"] = pv_test.baseline_short[:, i]
pv_pred_df[f"gate__{zone}"] = pv_test_gates[:, i]
pv_pred_df[f"mixed_base__{zone}"] = pv_test_mixed[:, i]
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 pv_base.ZONE_ORDER]
pv_keep += [f"pred_model__target_pv_dist__{zone}" for zone in pv_base.ZONE_ORDER]
pv_keep += [f"pred_season__target_pv_dist__{zone}" for zone in pv_base.ZONE_ORDER]
pv_keep += [f"pred_short8__target_pv_dist__{zone}" for zone in pv_base.ZONE_ORDER]
pv_keep += [f"gate__{zone}" for zone in pv_base.ZONE_ORDER]
pv_keep += [f"mixed_base__{zone}" for zone in pv_base.ZONE_ORDER]
pv_pred_df[pv_keep].to_parquet(PV_PRED_PATH, index=False)
attack_prev_metrics = json.loads((MODEL_DIR / "attack_distribution_gnn_metrics.json").read_text(encoding="utf-8"))["test_metrics"]
pv_prev_metrics = json.loads((MODEL_DIR / "pv_distribution_gnn_metrics.json").read_text(encoding="utf-8"))["test_metrics"]
attack_last3 = _mae_rows(attack_pred_df, "target_attack_share__")
attack_prev_last3 = _mae_rows(pd.read_parquet(MODEL_DIR / "attack_distribution_gnn_test_predictions.parquet"), "target_attack_share__")
attack_last3 = attack_last3.merge(
attack_prev_last3.rename(columns={"model_mae": "previous_model_mae", "season_mae": "previous_season_mae", "short8_mae": "previous_short8_mae"}),
on=["fecha", "team_name", "opponent_name"],
how="left",
)
pv_last3 = _mae_rows(pv_pred_df, "target_pv_dist__")
pv_prev_last3 = _mae_rows(pd.read_parquet(MODEL_DIR / "pv_distribution_gnn_test_predictions.parquet"), "target_pv_dist__")
pv_last3 = pv_last3.merge(
pv_prev_last3.rename(columns={"model_mae": "previous_model_mae", "season_mae": "previous_season_mae", "short8_mae": "previous_short8_mae"}),
on=["fecha", "team_name", "opponent_name"],
how="left",
)
summary = {
"attack": {
"train_rows": int(len(attack_train.metadata)),
"val_rows": int(len(attack_val.metadata)),
"test_rows": int(len(attack_test.metadata)),
"val_start_date": attack_val_start_date,
"test_metrics": attack_test_metrics,
"val_metrics": attack_val_metrics,
"previous_test_metrics": attack_prev_metrics,
"last3_racing": _last3_rows(attack_last3),
},
"pv": {
"train_rows": int(len(pv_train.metadata)),
"val_rows": int(len(pv_val.metadata)),
"test_rows": int(len(pv_test.metadata)),
"val_start_date": pv_val_start_date,
"test_metrics": pv_test_metrics,
"val_metrics": pv_val_metrics,
"previous_test_metrics": pv_prev_metrics,
"last3_racing": _last3_rows(pv_last3),
},
"artifacts": {
"attack_predictions": str(ATTACK_PRED_PATH),
"pv_predictions": str(PV_PRED_PATH),
},
"training": {
"attack_epochs": len(attack_history),
"pv_epochs": len(pv_history),
},
"config": {
"learning_rate": LEARNING_RATE,
"batch_size": BATCH_SIZE,
"max_epochs": MAX_EPOCHS,
"patience": PATIENCE,
"hidden_dim": 96,
"num_layers": 3,
"num_heads": 4,
},
}
torch.save(
{
"attack_state_dict": attack_model.state_dict(),
"pv_state_dict": pv_model.state_dict(),
"attack_global_bundle": attack_global_bundle,
"attack_node_bundle": attack_node_bundle,
"attack_global_cols": attack_global_cols,
"attack_node_feat_names": attack_node_feat_names,
"pv_global_bundle": pv_global_bundle,
"pv_node_bundle": pv_node_bundle,
"pv_node_feature_names": pv_node_feature_names,
"zone_order": attack_base.ZONE_ORDER,
"config": summary["config"],
},
MODEL_PATH,
)
JSON_PATH.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
REPORT_PATH.write_text(_report_html(summary), 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()