File size: 3,047 Bytes
a61ae40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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)