求鸢尾花数据集1D PCA降维可视化按类别着色的Python代码
鸢尾花1D PCA降维结果按类别着色实现
我已使用PCA将鸢尾花(iris)数据集降维至1D,编写了基础1D绘图代码,且参考过Stack Overflow上的《1D plot matplotlib》问题。现希望参考自己实现的2D降维按类别着色绘图代码,为1D可视化中的数据点按setosa、versicolor、virginica三个类别着色,请求对应的实现代码。
现有1D降维绘图代码
import numpy as np import matplotlib.pyplot as plt from sklearn import datasets from sklearn.decomposition import PCA import matplotlib.cm as cm # 加载数据集 iris = datasets.load_iris() X = iris.data Y = iris.target # PCA降维至1D pca = PCA(n_components=1, whiten=False) transformed = pca.fit_transform(X) # 绘制基础1D图 plt.figure(figsize=(10, 2)) plt.hlines(1, -10, 20) # 绘制水平线 plt.xlim(-5, 5) plt.ylim(0.8, 1.2) y = np.ones(np.shape(transformed)) # 所有点的y值设为1 plt.plot(transformed, y, 'x', ms=10) # 绘制所有数据点 plt.show()
参考的2D降维按类别着色代码
import numpy as np import matplotlib.pyplot as plt from sklearn import datasets from sklearn.decomposition import PCA import matplotlib.cm as cm # 加载数据集 iris = datasets.load_iris() X = iris.data Y = iris.target # PCA降维至2D pca = PCA(n_components=2, whiten=False) transformed = pca.fit_transform(X) # 按类别着色绘图 for label in np.unique(Y): if label == 0: c = "red" elif label == 1: c = "blue" elif label == 2: c = "green" else: pass plt.scatter(transformed[Y == label, 0], transformed[Y == label, 1], c=c) plt.title('主成分分析结果') plt.xlabel('主成分1') plt.ylabel('主成分2') plt.show()
实现按类别着色的1D绘图代码
import numpy as np import matplotlib.pyplot as plt from sklearn import datasets from sklearn.decomposition import PCA # 加载数据集 iris = datasets.load_iris() X = iris.data Y = iris.target # PCA降维至1D pca = PCA(n_components=1, whiten=False) transformed = pca.fit_transform(X) # 设置绘图参数 plt.figure(figsize=(10, 2)) plt.hlines(1, -10, 20, color='gray') # 绘制灰色水平线 plt.xlim(-5, 5) plt.ylim(0.8, 1.2) plt.title('鸢尾花数据集1D PCA降维(按类别着色)') plt.xlabel('主成分1') # 定义类别与颜色映射 label_color_map = { 0: 'red', # setosa 1: 'blue', # versicolor 2: 'green' # virginica } # 按类别循环绘制数据点 for label in np.unique(Y): # 获取当前类别的数据点 class_points = transformed[Y == label, 0] # 生成对应长度的y值(固定为1) y = np.ones(len(class_points)) # 绘制当前类别的点,指定颜色和标记 plt.plot(class_points, y, 'x', ms=10, color=label_color_map[label], label=iris.target_names[label]) # 添加图例 plt.legend(loc='upper right') plt.show()
内容的提问来源于stack exchange,提问作者rrkk
相关产品推荐
相关产品推荐

