fig = plt.figure(figsize=(12, 12)) ax = fig.add_subplot(projection='3d') ax.patch.set_alpha(0) plt.tick_params(left=False, bottom=False, labelleft=False, labelbottom=False) ax.scatter(embed[:, 0], embed[:, 1], embed[:, 2], s=200, c=data.y, cmap="hsv", vmin=-2, vmax=3)