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={})
| 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, |
nothing calls this directly
no test coverage detected