58 for _
in range(NUM_LAYERS):
60 num_inputs = LATENT_SIZE
68 def __init__(self, num_node_inputs, num_edge_inputs, num_global_inputs):
74 def forward(self, node_data, edge_data, global_data):
75 return self.node_fn(node_data), self.edge_fn(edge_data), self.global_fn(global_data)
83 def __init__(self, num_node_inputs, num_edge_inputs, num_global_inputs):
85 self.edge_fn =
make_mlp_model(num_edge_inputs + 2 * num_node_inputs + num_global_inputs,
True)
86 self.node_fn =
make_mlp_model(LATENT_SIZE + num_node_inputs + num_global_inputs,
True)
87 self.global_fn =
make_mlp_model(2 * LATENT_SIZE + num_global_inputs,
True)
89 def forward(self, node_data, edge_data, global_data, receivers, senders):
93 [edge_data, node_data[receivers], node_data[senders],
global_data.expand(n_edges, -1)], dim=1
95 edge_output = self.edge_fn(edge_input)
101 node_output = self.node_fn(node_input)
105 global_output = self.global_fn(global_input)
106 return node_output, edge_output, global_output
114 self._core =
MLPGraphNetwork(2 * LATENT_SIZE, 2 * LATENT_SIZE, 2 * LATENT_SIZE)
118 def forward(self, node_data, edge_data, global_data, receivers, senders, num_processing_steps):
119 latent = self._encoder(node_data, edge_data, global_data)
122 for _
in range(num_processing_steps):
124 latent = self._core(*core_input, receivers, senders)
125 decoded_op = self._decoder(*latent)
147 input_names = [
"node_data",
"edge_data",
"global_data"]
149 "node_data": {0: num_nodes_dim},
150 "edge_data": {0: num_edges_dim},
158 input_names += [
"receivers",
"senders"]
164 input_names=input_names,
165 output_names=[
"node_output",
"edge_output",
"global_output"],
166 dynamic_shapes=dynamic_shapes,
178for name
in [
"encoder",
"core",
"decoder",
"output_transform"]:
182 print(
"generated SOFIE model", name +
".hxx")
196tree.Branch(
"node_data",
"std::vector<float>", node_data)
197tree.Branch(
"edge_data",
"std::vector<float>", edge_data)
198tree.Branch(
"global_data",
"std::vector<float>", global_data)
199tree.Branch(
"receivers",
"std::vector<int>", receivers)
202print(
"\n\nSaving data in a ROOT File:")
203h1 =
ROOT.TH1D(
"h1",
"GNN nodes output", 40, 1, 0)
204h2 =
ROOT.TH1D(
"h2",
"GNN edges output", 40, 1, 0)
205h3 =
ROOT.TH1D(
"h3",
"GNN global output", 40, 1, 0)
207for i
in range(numevts):
221for graphData
in dataset:
235print(
"time to evaluate ", numevts,
" events", end - start)
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.