class GraphSAGE(torch.nn.Module): """GraphSAGE""" def __init__(self, dim_in, dim_h, dim_out): super().__init__() self.sage1 = SAGEConv(dim_in, dim_h) self.sage2 = SAGEConv(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 = self.sage1(x, edge_index) h = torch.relu(h) h = F.dropout(h, p=0.5, training=self.training) h = self.sage2(h, edge_index) return h, F.log_softmax(h, dim=1) def fit(self, data, epochs): criterion = torch.nn.CrossEntropyLoss() optimizer = self.optimizer self.train() for epoch in range(epochs+1): acc = 0 val_loss = 0 val_acc = 0 # Train on batches for batch in train_loader: optimizer.zero_grad() _, out = self(batch.x, batch.edge_index) loss = criterion(out[batch.train_mask], batch.y[batch.train_mask]) acc += accuracy(out[batch.train_mask].argmax(dim=1), batch.y[batch.train_mask]) loss.backward() optimizer.step() # Validation val_loss += criterion(out[batch.val_mask], batch.y[batch.val_mask]) val_acc += accuracy(out[batch.val_mask].argmax(dim=1), batch.y[batch.val_mask]) # Print metrics every 10 epochs if(epoch % 10 == 0): print(f'Epoch {epoch:>3} | Train Loss: {loss/len(train_loader):.3f} ' f'| Train Acc: {acc/len(train_loader)*100:>6.2f}% | Val Loss: ' f'{val_loss/len(train_loader):.2f} | Val Acc: ' f'{val_acc/len(train_loader)*100:.2f}%')