如何修改3D t-SNE散点图的图例标签、指定类别固定颜色及排查报错
问题解答
首先说明你修改后代码的错误原因
你现在的代码共调用了5次ax.scatter(),每次都传入了全部样本的坐标数据,相当于把所有点重复绘制了5次、每层覆盖不同颜色,最终所有点的颜色都会被最后一次绘制的紫色覆盖,自然会出现颜色重叠异常的问题。你需要按类别筛选出对应的数据点,每次仅绘制单个类别的点。
完整解决方案(同时满足你两个原始需求)
第一步:预定义统一的类别映射规则(2D、3D图共用,保证颜色完全一致)
提前把类别数值、自定义类别名、对应颜色的映射关系写死,避免两边绘图时自动生成的调色盘不一致:
# 请根据你实际Y的类别数值调整key,下面默认你的Y是0-4的五个数值 class_mapping = { 0: {"name": "类别1", "color": "blue"}, 1: {"name": "类别2", "color": "orange"}, 2: {"name": "类别3", "color": "gray"}, 3: {"name": "类别4", "color": "cyan"}, 4: {"name": "类别5", "color": "purple"} } # 提取统一调色盘给2D图使用 custom_palette = [class_mapping[k]["color"] for k in sorted(class_mapping.keys())]
第二步:修改2D散点图代码
仅需要把原来的palette参数替换为上面定义的custom_palette即可,如果想要图例直接显示自定义名称,可以额外加一步把Y映射为名称:
f, ax = plt.subplots(figsize=(12, 12)) # 可选:如果要直接在2D图例显示自定义名称,把Y换成映射后的名称列 # Y_name = Y.map(lambda x: class_mapping[x]["name"]) sns.scatterplot(x = tsne_x, y = tsne_y, hue = Y, # 要显示自定义名就把这里换成Y_name palette = custom_palette, alpha = 0.8)
第三步:修改3D散点图代码
按类别筛选数据后分别绘制,同时传入自定义图例标签:
import seaborn as sns, numpy as np, pandas as pd import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D sns.set_style("whitegrid", {'axes.grid' : False}) fig = plt.figure(figsize=(12,12)) ax = Axes3D(fig) # 把坐标和类别合并成DataFrame方便按类别筛选 plot_df = pd.DataFrame({ "x": tsne_x, "y": tsne_y, "z": tsne_z, "class": Y }) handles = [] labels = [] for class_val, info in class_mapping.items(): # 仅筛选当前类别的数据点绘制 subset = plot_df[plot_df["class"] == class_val] sc = ax.scatter(subset["x"], subset["y"], subset["z"], color=info["color"], alpha=0.8) handles.append(sc) labels.append(info["name"]) ax.set_xlabel('X Label') ax.set_ylabel('Y Label') ax.set_zlabel('Z Label') # 传入自定义的图例元素和名称 ax.legend(handles, labels, loc="upper right", title="类别") plt.show()
原始问题对应解答
- 修改图例名称:可以通过两种方式实现,一是提前把Y列的数值映射为你需要的自定义名称字符串,二是绘图后手动给legend方法传入自定义的labels列表。
- 固定类别颜色:只需要提前预定义类别到颜色的映射关系,2D和3D绘图时都使用这套固定的映射即可,不要使用自动生成的调色盘。
内容的提问来源于stack exchange,提问作者U23r
相关产品推荐
相关产品推荐

