MCPcopy Create free account
hub / github.com/IBM/Project_CodeNet / GNN_node

Class GNN_node

model-experiments/gnn-based-experiments/src/model/conv.py:71–140  ·  view source on GitHub ↗

Output: node representations

Source from the content-addressed store, hash-verified

69
70### GNN to generate node embedding
71class GNN_node(torch.nn.Module):
72 """
73 Output:
74 node representations
75 """
76 def __init__(self, num_layer, emb_dim, node_encoder, drop_ratio = 0.5, JK = "last", residual = False, gnn_type = 'gin', edge_attr_dim=2):
77 '''
78 emb_dim (int): node embedding dimensionality
79 num_layer (int): number of GNN message passing layers
80
81 '''
82
83 super(GNN_node, self).__init__()
84 self.num_layer = num_layer
85 self.drop_ratio = drop_ratio
86 self.JK = JK
87 ### add residual connection or not
88 self.residual = residual
89
90 if self.num_layer < 2:
91 raise ValueError("Number of GNN layers must be greater than 1.")
92
93 self.node_encoder = node_encoder
94
95 ###List of GNNs
96 self.convs = torch.nn.ModuleList()
97 self.batch_norms = torch.nn.ModuleList()
98
99 for layer in range(num_layer):
100 if gnn_type == 'gin':
101 self.convs.append(GINConv(emb_dim, edge_attr_dim))
102 elif gnn_type == 'gcn':
103 self.convs.append(GCNConv(emb_dim, edge_attr_dim))
104 else:
105 ValueError('Undefined GNN type called {}'.format(gnn_type))
106
107 self.batch_norms.append(torch.nn.BatchNorm1d(emb_dim))
108
109 def forward(self, batched_data):
110 x, edge_index, edge_attr, node_depth, batch = batched_data.x, batched_data.edge_index, batched_data.edge_attr, batched_data.node_depth, batched_data.batch
111
112
113 ### computing input node embedding
114
115 h_list = [self.node_encoder(x, node_depth.view(-1,))]
116 for layer in range(self.num_layer):
117
118 h = self.convs[layer](h_list[layer], edge_index, edge_attr)
119 h = self.batch_norms[layer](h)
120
121 if layer == self.num_layer - 1:
122 #remove relu for the last layer
123 h = F.dropout(h, self.drop_ratio, training = self.training)
124 else:
125 h = F.dropout(F.relu(h), self.drop_ratio, training = self.training)
126
127 if self.residual:
128 h += h_list[layer]

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected