GNN4Colliders / tests /unit /features /test_objects.py
ho22joshua's picture
feat: implement shared collider feature construction
a61ae40
Raw
History Blame
3.05 kB
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)