from torch_geometric.utils import degree # Get model's classifications _, out = gat(data.x, data.edge_index) # Calculate the degree of each node degrees = degree(data.edge_index[0]).numpy() # Store accuracy scores and sample sizes accuracies = [] sizes = [] # Accuracy for degrees between 0 and 5 for i in range(0, 6): mask = np.where(degrees == i)[0] accuracies.append(accuracy(out.argmax(dim=1)[mask], data.y[mask])) sizes.append(len(mask)) # Accuracy for degrees > 5 mask = np.where(degrees > 5)[0] accuracies.append(accuracy(out.argmax(dim=1)[mask], data.y[mask])) sizes.append(len(mask)) # Bar plot fig, ax = plt.subplots(figsize=(18, 9)) ax.set_xlabel('Node degree') ax.set_ylabel('Accuracy score') ax.set_facecolor('#EFEEEA') plt.bar(['0','1','2','3','4','5','>5'], accuracies, color='#0A047A') for i in range(0, 7): plt.text(i, accuracies[i], f'{accuracies[i]*100:.2f}%', ha='center', color='#0A047A') for i in range(0, 7): plt.text(i, accuracies[i]//2, sizes[i], ha='center', color='white')