class GIN(torch.nn.Module): """GIN""" def __init__(self, dim_h): super(GIN, self).__init__() self.conv1 = GINConv( Sequential(Linear(dataset.num_node_features, dim_h), BatchNorm1d(dim_h), ReLU(), Linear(dim_h, dim_h), ReLU())) self.conv2 = GINConv( Sequential(Linear(dim_h, dim_h), BatchNorm1d(dim_h), ReLU(), Linear(dim_h, dim_h), ReLU())) self.conv3 = GINConv( Sequential(Linear(dim_h, dim_h), BatchNorm1d(dim_h), ReLU(), Linear(dim_h, dim_h), ReLU())) self.lin1 = Linear(dim_h*3, dim_h*3) self.lin2 = Linear(dim_h*3, dataset.num_classes) def forward(self, x, edge_index, batch): # Node embeddings h1 = self.conv1(x, edge_index) h2 = self.conv2(h1, edge_index) h3 = self.conv3(h2, edge_index) # Graph-level readout h1 = global_add_pool(h1, batch) h2 = global_add_pool(h2, batch) h3 = global_add_pool(h3, batch) # Concatenate graph embeddings h = torch.cat((h1, h2, h3), dim=1) # Classifier h = self.lin1(h) h = h.relu() h = F.dropout(h, p=0.5, training=self.training) h = self.lin2(h) return h, F.log_softmax(h, dim=1) gcn = GCN(dim_h=32) gin = GIN(dim_h=32) gcn = train(gcn, train_loader) gin = train(gin, train_loader)