plt.figure(dpi=200) plt.scatter(X[:,0], X[:,1], c=c) plt.xlabel("x") plt.ylabel("y") plt.axvline(x=0.9, label="split at x=0.9", c = "k", linestyle="--") plt.legend() plt.show()