Spaces:
Running
Running
| """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 | |