G = to_networkx(data, to_undirected=True) plt.figure(figsize=(12,12)) plt.axis('off') nx.draw_networkx(G, pos=nx.spring_layout(G, seed=0), with_labels=True, node_size=800, node_color=data.y, cmap="hsv", vmin=-2, vmax=3, width=0.8, edge_color="grey", font_size=14 ) plt.show()