MCPcopy Create free account
hub / github.com/LBANN/lbann / __init__

Method __init__

python/lbann/modules/graph/sparse/NNConv.py:11–64  ·  view source on GitHub ↗

Inititalize the edge conditioned graph kernel with edge data represented with pseudo-COO format. The reduction over edge features are performed via the scatter layer The update function of the kernel is: .. math:: X^{\prime}_{i} = \Theta

(self,
                 sequential_nn,
                 num_nodes,
                 num_edges,
                 input_channels,
                 output_channels,
                 edge_input_channels,
                 activation=lbann.Relu,
                 name=None,
                 parallel_strategy={})

Source from the content-addressed store, hash-verified

9 global_count = 0
10
11 def __init__(self,
12 sequential_nn,
13 num_nodes,
14 num_edges,
15 input_channels,
16 output_channels,
17 edge_input_channels,
18 activation=lbann.Relu,
19 name=None,
20 parallel_strategy={}):
21 """Inititalize the edge conditioned graph kernel with edge data
22 represented with pseudo-COO format. The reduction over edge
23 features are performed via the scatter layer
24 The update function of the kernel is:
25 .. math::
26 X^{\prime}_{i} = \Theta x_i + \sum_{j \in \mathcal{N(i)}}x_j \cdot h_{\Theta}(e_{i,j})
27 where :math:`h_{\mathbf{\Theta}}` denotes a channel-wise NN module
28 Args:
29 sequential_nn ([Module] or (Module)): A list or tuple of layer
30 modules for updating the
31 edge feature matrix
32 num_nodes (int): Number of vertices of each graph
33 (max number in the batch padded by 0)
34 num_edges (int): Number of edges of each graph
35 (max in the batch padded by 0)
36 output_channels (int): The output size of each node feature after
37 transformed with learnable weights
38 activation (type): The activation function of the node features
39 name (str): Default name of the layer is NN_{number}
40 parallel_strategy (dict): Data partitioning scheme.
41 """
42 NNConv.global_count += 1
43
44 self.name = (name
45 if name
46 else 'NNConv_{}'.format(NNConv.global_count))
47
48 self.output_channels = output_channels
49 self.input_channels = input_channels
50
51 self.num_nodes = num_nodes
52 self.num_edges = num_edges
53 self.edge_input_channels = edge_input_channels
54
55 self.node_activation = activation
56 self.parallel_strategy = parallel_strategy
57
58 self.node_nn = \
59 ChannelwiseFullyConnectedModule(self.output_channels,
60 bias=False,
61 activation=self.node_activation,
62 name=self.name+"_node_weights",
63 parallel_strategy=self.parallel_strategy)
64 self.edge_nn = sequential_nn
65
66 def message(self,
67 node_features,

Callers

nothing calls this directly

Tested by

no test coverage detected