File size: 2,872 Bytes
9da4e9c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
"""Verifica el contrato del modelo de ataque rediseñado (matchup GNN) tal como lo
consume la app: importa el módulo vendorizado, carga el bundle y predice una
distribución válida de ataque sobre las 10 zonas (suma ~1).

Se saltea si los artifacts NN no están disponibles localmente (cold CI sin Azure).
"""

from __future__ import annotations

import numpy as np
import pytest

matplotlib = pytest.importorskip("matplotlib")
matplotlib.use("Agg")
torch = pytest.importorskip("torch")

from racing_reports import vendor_env  # noqa: E402

ATTACK_BUNDLE = vendor_env.DATA_DIR / "modeling" / "attack_matchup_gnn_bundle.pt"
MODELING_DATASET = vendor_env.DATA_DIR / "modeling" / "attack_prediction_dataset.parquet"

pytestmark = pytest.mark.skipif(
    not (ATTACK_BUNDLE.exists() and MODELING_DATASET.exists()),
    reason="Artifacts NN (bundle/dataset) no disponibles localmente.",
)


def test_app_imports_matchup_model() -> None:
    pro = vendor_env.import_script("build_head_to_head_report_pro")
    assert pro.attack_model_mod.__name__ == "train_attack_matchup_gnn"
    assert pro.ATTACK_MODEL_PATH.name == "attack_matchup_gnn_bundle.pt"


def test_bundle_loads_without_key_mismatch() -> None:
    attack_mod = vendor_env.import_script("train_attack_matchup_gnn")
    bundle = torch.load(ATTACK_BUNDLE, map_location="cpu", weights_only=False)
    global_dim = len(bundle["global_bundle"]["global_feature_columns"])
    node_dim = len(bundle["node_bundle"]["means"][0])
    model = attack_mod.ResidualAttackGNN(node_dim=node_dim, global_dim=global_dim)
    missing, unexpected = model.load_state_dict(bundle["model_state_dict"], strict=False)
    assert not missing, f"faltan claves: {missing}"
    assert not unexpected, f"claves inesperadas: {unexpected}"
    assert bundle["zone_order"] == attack_mod.ZONE_ORDER


def test_prediction_is_a_valid_distribution() -> None:
    pro = vendor_env.import_script("build_head_to_head_report_pro")
    df_model = pro._load_modeling_dataset()
    racing = "bzkwzatvwahmbzok1ymm5vqa1"
    sub = df_model[df_model["teamId"] == racing]
    if sub.empty:
        pytest.skip("Racing no presente en el dataset local.")
    row = sub.iloc[-1]
    opp_id = (
        row["opponent_team_id"]
        if "opponent_team_id" in sub.columns and row.get("opponent_team_id")
        else df_model[df_model["team_name"] == row["opponent_name"]]["teamId"].iloc[-1]
    )
    pred = pro._predict_single_team_future(
        df_model, racing, opp_id, "Racing de Santander", str(row["opponent_name"]),
        row["league"], row["season"], is_home=True,
    )
    zones = list(pred.attack.keys())
    assert len(zones) == 10
    total = sum(pred.attack.values())
    assert total == pytest.approx(1.0, abs=1e-4), f"la distribución no suma 1: {total}"
    assert all(0.0 <= v <= 1.0 for v in pred.attack.values())
    assert len(pred.pv) == 10