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

Method forward

python/lbann/modules/graph/sparse/GINConv.py:41–79  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

GraphExpandFunction · 0.90
GraphReduceFunction · 0.90
activationFunction · 0.85

Tested by

no test coverage detected