import numpy as np import pytest import torch from gnn4colliders.features import NODE_FEATURE_NAMES, build_node_features @pytest.fixture def schema(): return ( [ ["jet_pt", "ele_pt", "mu_pt", "ph_pt", "MET_met"], ["jet_eta", "ele_eta", "mu_eta", "ph_eta", 0], ["jet_phi", "ele_phi", "mu_phi", "ph_phi", "MET_phi"], "CALC_E", ["jet_btag", 0, 0, 0, 0], [0, "ele_charge", "mu_charge", 0, 0], "NODE_TYPE", ], ["vector", "vector", "vector", "vector", "single"], [0.1, 1, 1, 0.1, 1, 1, 1], ) @pytest.fixture def event(): return { "jet_pt": np.array([100.0, 50.0], dtype=np.float32), "ele_pt": np.array([20.0], dtype=np.float32), "mu_pt": np.array([30.0], dtype=np.float32), "ph_pt": np.array([40.0], dtype=np.float32), "MET_met": np.float32(25.0), "jet_eta": np.array([1.0, -0.5], dtype=np.float32), "ele_eta": np.array([0.25], dtype=np.float32), "mu_eta": np.array([-0.75], dtype=np.float32), "ph_eta": np.array([0.5], dtype=np.float32), "jet_phi": np.array([3.0, -3.0], dtype=np.float32), "ele_phi": np.array([0.2], dtype=np.float32), "mu_phi": np.array([-0.4], dtype=np.float32), "ph_phi": np.array([1.0], dtype=np.float32), "MET_phi": np.float32(-1.2), "jet_btag": np.array([0.8, 0.1], dtype=np.float32), "ele_charge": np.array([-1.0], dtype=np.float32), "mu_charge": np.array([1.0], dtype=np.float32), } def test_schema_and_values(event, schema): names, object_types, scales = schema features, lengths = build_node_features(event, names, object_types, scales) assert NODE_FEATURE_NAMES == ( "pt", "eta", "phi", "energy", "btag", "charge", "node_type", ) assert lengths == [2, 1, 1, 1, 1] assert features.shape == (6, 7) assert features.dtype == torch.float32 np.testing.assert_allclose( features.numpy(), [ [10.0, 1.0, 3.0, 15.431, 0.8, 0.0, 0.0], [5.0, -0.5, -3.0, 5.638, 0.1, 0.0, 0.0], [2.0, 0.25, 0.2, 2.063, 0.0, -1.0, 1.0], [3.0, -0.75, -0.4, 3.884, 0.0, 1.0, 2.0], [4.0, 0.5, 1.0, 4.511, 0.0, 0.0, 3.0], [2.5, 0.0, -1.2, 2.500, 0.0, 0.0, 4.0], ], rtol=0, atol=2e-3, ) def test_empty_vectors_and_input_immutability(event, schema): names, object_types, scales = schema event = dict(event) event["jet_pt"] = np.array([], dtype=np.float32) event["jet_eta"] = np.array([], dtype=np.float32) event["jet_phi"] = np.array([], dtype=np.float32) event["jet_btag"] = np.array([], dtype=np.float32) before = event["ele_pt"].copy() features, lengths = build_node_features(event, names, object_types, scales) assert lengths == [0, 1, 1, 1, 1] assert features.shape == (4, 7) np.testing.assert_array_equal(event["ele_pt"], before)