Apply GIN Layer. Args: node_feature_mat (Layer): Node feature matrix with the shape of (num_nodes,input_channels) source_indices (Layer): Source node indices of the edges with shape (num_nodes) target_indices (Layer): Target node indices of the edges wit
(self,
node_feature_mat,
source_indices,
target_indices,
activation = lbann.Relu)
| 39 | self.num_edges = num_edges |
| 40 | |
| 41 | def forward(self, |
| 42 | node_feature_mat, |
| 43 | source_indices, |
| 44 | target_indices, |
| 45 | activation = lbann.Relu): |
| 46 | """Apply GIN Layer. |
| 47 | |
| 48 | Args: |
| 49 | node_feature_mat (Layer): Node feature matrix with the shape of (num_nodes,input_channels) |
| 50 | source_indices (Layer): Source node indices of the edges with shape (num_nodes) |
| 51 | target_indices (Layer): Target node indices of the edges with shape (num_nodes |
| 52 | activation (Layer): Activation layer for the node features. If None, then no activation is |
| 53 | applied. (default: lbann.Relu) |
| 54 | Returns: |
| 55 | (Layer) : The output after kernel ops. The output can passed into another Graph Conv layer |
| 56 | directly |
| 57 | """ |
| 58 | eps = lbann.Constant(value=(1+self.eps), |
| 59 | num_neurons = [self.num_nodes, self.input_channel_size]) |
| 60 | |
| 61 | eps_node_features = lbann.Multiply(node_feature_mat, eps, name=self.name+"_epl_mult") |
| 62 | |
| 63 | node_feature_mat = lbann.Sum(eps_node_features, node_feature_mat) |
| 64 | |
| 65 | # Transform with the sequence of linear layers |
| 66 | for layer in self.nn: |
| 67 | node_feature_mat = layer(node_feature_mat) |
| 68 | |
| 69 | neighborhoods = GraphExpand(node_feature_mat, target_indices) |
| 70 | |
| 71 | neighborhoods = lbann.Reshape(neighborhoods, dims=[self.num_edges, self.output_channel_size]) |
| 72 | |
| 73 | aggregated_node_features = GraphReduce(neighborhoods, source_indices, [self.num_nodes, |
| 74 | self.output_channel_size]) |
| 75 | ## Apply activation |
| 76 | if activation: |
| 77 | aggregated_node_features = activation(aggregated_node_features) |
| 78 | |
| 79 | return aggregated_node_features |
nothing calls this directly
no test coverage detected