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.axhline(y=0.4, xmin=0, xmax=0.59, label="split at y=0.4", c = "pink", linestyle="--") plt.legend() plt.show()