You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为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中的文字可替换为实际业务分类名。
  • 分组绘制的方式不需要额外手动匹配色值,和原代码的视觉效果完全一致,不会出现颜色错位问题。

3D散点图示例

内容的提问来源于stack exchange,提问作者Joe

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 18:15:51