Output: node representations
| 69 | |
| 70 | ### GNN to generate node embedding |
| 71 | class 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] |