import torch.nn.functional as F from torch.nn import Linear, Dropout from torch_geometric.nn import GCNConv, GATv2Conv class GCN(torch.nn.Module): """Graph Convolutional Network""" def __init__(self, dim_in, dim_h, dim_out): super().__init__() self.gcn1 = GCNConv(dim_in, dim_h) self.gcn2 = GCNConv(dim_h, dim_out) self.optimizer = torch.optim.Adam(self.parameters(), lr=0.01, weight_decay=5e-4) def forward(self, x, edge_index): h = F.dropout(x, p=0.5, training=self.training) h = self.gcn1(h, edge_index) h = torch.relu(h) h = F.dropout(h, p=0.5, training=self.training) h = self.gcn2(h, edge_index) return h, F.log_softmax(h, dim=1) class GAT(torch.nn.Module): """Graph Attention Network""" def __init__(self, dim_in, dim_h, dim_out, heads=8): super().__init__() self.gat1 = GATv2Conv(dim_in, dim_h, heads=heads) self.gat2 = GATv2Conv(dim_h*heads, dim_out, heads=1) self.optimizer = torch.optim.Adam(self.parameters(), lr=0.005, weight_decay=5e-4) def forward(self, x, edge_index): h = F.dropout(x, p=0.6, training=self.training) h = self.gat1(x, edge_index) h = F.elu(h) h = F.dropout(h, p=0.6, training=self.training) h = self.gat2(h, edge_index) return h, F.log_softmax(h, dim=1) def accuracy(pred_y, y): """Calculate accuracy.""" return ((pred_y == y).sum() / len(y)).item() def train(model, data): """Train a GNN model and return the trained model.""" criterion = torch.nn.CrossEntropyLoss() optimizer = model.optimizer epochs = 200 model.train() for epoch in range(epochs+1): # Training optimizer.zero_grad() _, out = model(data.x, data.edge_index) loss = criterion(out[data.train_mask], data.y[data.train_mask]) acc = accuracy(out[data.train_mask].argmax(dim=1), data.y[data.train_mask]) loss.backward() optimizer.step() # Validation val_loss = criterion(out[data.val_mask], data.y[data.val_mask]) val_acc = accuracy(out[data.val_mask].argmax(dim=1), data.y[data.val_mask]) # Print metrics every 10 epochs if(epoch % 10 == 0): print(f'Epoch {epoch:>3} | Train Loss: {loss:.3f} | Train Acc: ' f'{acc*100:>6.2f}% | Val Loss: {val_loss:.2f} | ' f'Val Acc: {val_acc*100:.2f}%') return model def test(model, data): """Evaluate the model on test set and print the accuracy score.""" model.eval() _, out = model(data.x, data.edge_index) acc = accuracy(out.argmax(dim=1)[data.test_mask], data.y[data.test_mask]) return acc