| 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) |
|
|