h, _ = gat(data.x, data.edge_index) # Train TSNE tsne = TSNE(n_components=2, learning_rate='auto', init='pca').fit_transform(h.detach()) # Plot TSNE plt.figure(figsize=(10, 10)) plt.axis('off') plt.scatter(tsne[:, 0], tsne[:, 1], s=50, c=data.y) plt.show()