如何为matplotlib mplot3d绘制的3D散点图正确添加图例
3D散点图图例添加方案
报错原因
直接调用ax.legend()返回No artists with labels found to put in legend提示,是因为原代码通过c = y_train_new传入分类数组批量上色的写法,不会为每个分类生成带标签的独立绘图对象,matplotlib无法自动识别分类和散点标记的对应关系。
实现方案
按分类标签分组调用scatter3D绘制,每个分类单独传入对应颜色和标签参数,即可正常生成图例,配色完全匹配原代码使用的tab10_r色板,同时修复了原代码重复创建画布的小问题。
原代码连续调用两次画布创建语句,会额外生成一张空白画布,修正版本已移除该冗余逻辑。
完整可运行代码(包含提供的测试数据):
import numpy as np import pandas as pd import matplotlib.pyplot as plt from mpl_toolkits import mplot3d import seaborn as sns # 加载测试数据,使用自有PCA数据时可删除该段 data = pd.DataFrame([ [-3.8481877, -0.47685334, 0.63422906, 1.0396314, 1], [-2.320888, 0.65347993, 1.1519914, 0.12997247, 1], [1.5827686, 1.4119303, -1.7410104, -4.6962333, 1], [-0.1337152, 0.13315737, -1.6648949, -1.4205348, 1], [-0.4028037, 1.332986, 1.3618442, 0.3292255, 1], [-0.015517877, 1.346349, 1.4083523, 0.87017965, 1], [-0.2669228, 0.5478992, -0.06730786, -1.5959451, 1], [-0.03318152, 0.3263167, -2.116833, -5.4616213, 1], [0.4588691, 0.6723614, -1.617398, -4.3511734, 1], [0.5899199, 0.66525555, -1.694493, -3.9452586, 1], [1.610061, 2.4186094, 1.8807093, 1.3764497, 0], [1.7985699, 2.4387648, 1.6306056, 1.1184534, 0], [-9.222036, -9.9776, -9.832, -9.909746, 0], [0.21364458, -1.0171559, -4.9093766, -6.2154694, 0], [-0.019955145, -1.1677283, -4.6549516, -5.9503417, 0], [0.44730473, -0.77167743, -4.7527356, -5.971007, 0], [-0.16508447, -0.005777468, -1.5020386, -4.49326, 0], [-0.8654994, -0.54387957, -1.300646, -4.621529, 0], [-1.7471086, -2.0005553, -1.7533782, -2.6065414, 0], [-1.5313624, -1.6995796, -1.4394685, -2.600004, 0] ], columns=['x','y','z','not_used','Label']) # 匹配原代码的配色设置 colors = dict(zip(['2', '1', '0'], sns.color_palette('tab10_r', 3))) new_colors = list(colors.values()) del new_colors[0] fig = plt.figure(figsize=(15, 12)) ax = plt.axes(projection="3d") # 按标签分组绘制散点,可自行修改label_map里的文字为实际分类名称 label_map = {0:'类别0', 1:'类别1'} for label_val in sorted(data['Label'].unique()): subset = data[data['Label'] == label_val] ax.scatter3D( subset['x'], subset['y'], subset['z'], color = new_colors[label_val], marker = 'o', alpha=0.6, s=55, edgecolor='k', label=label_map[label_val] ) plt.title("3D Scatterplot: 95% of the variability captured", pad = 15) ax.set_xlabel('First principal component') ax.set_ylabel('Second principal component') ax.set_zlabel('Third principal component') # 调用图例即可正常显示 ax.legend(fontsize=12) plt.show()
使用说明
- 替换为自有PCA数据时,删除测试数据加载段,将循环内的子集筛选逻辑改为按
y_train_new的取值切分对应的x_pca数组即可,label_map中的文字可替换为实际业务分类名。 - 分组绘制的方式不需要额外手动匹配色值,和原代码的视觉效果完全一致,不会出现颜色错位问题。

内容的提问来源于stack exchange,提问作者Joe
相关产品推荐
相关产品推荐

