Python中为MNIST数据集文本标记t-SNE可视化图添加图例
解决MNIST t-SNE可视化添加数字颜色图例的问题
嘿,我懂你的困扰——用文本标注数字后没法直接生成对应颜色的图例对吧?这是因为plt.annotate添加的文本并不是matplotlib图例默认识别的绘图元素(比如散点、线条),所以我们得手动创建代理艺术家来生成图例。
下面是修改后的完整代码,我修复了导入语句的错误,同时添加了图例的实现:
import pandas as pd import numpy as np import matplotlib.pyplot as plt from sklearn import (manifold, datasets) import seaborn as sns digits = datasets.load_digits(n_class=10) X = digits.data y = digits.target # 运行t-SNE降维 tsne = manifold.TSNE(n_components=2, random_state=0, perplexity=40) X_t = tsne.fit_transform(X) # 整理数据成DataFrame XY = pd.DataFrame(np.column_stack([X_t, y]), columns=['x', 'y', 'digit']) XY['digit'] = XY['digit'].astype(int) # 生成颜色调色板 customPalette = sns.hls_palette(10, l=.7, s=1) fig, axes = plt.subplots(figsize=(10,10)) axes.set_xlim(XY['x'].min()*0.98, XY['x'].max()*1.02) axes.set_ylim(XY['y'].min()*0.98, XY['y'].max()*1.02) # 标注每个数字文本 for digit in range(10): subset = XY[XY['digit'] == digit] for _, row in subset.iterrows(): plt.annotate( row['digit'], (row['x'], row['y']), ha='center', va='center', size=10, color=customPalette[digit] ) # 创建代理元素用于生成图例 proxy_artists = [] for digit in range(10): # 用带对应颜色的圆点作为代理(仅用于图例,不显示在主图) proxy = plt.Line2D( [], [], marker='o', color=customPalette[digit], linestyle='', markersize=12, label=str(digit) ) proxy_artists.append(proxy) # 添加图例,调整位置避免遮挡主图 plt.legend( handles=proxy_artists, loc='upper right', bbox_to_anchor=(1.2, 1), title='Digit', fontsize=10 ) plt.title('t-SNE on MNIST Digits') plt.show()
关键说明:
- 代理艺术家:因为我们没有使用
scatter这类能生成可识别图例元素的函数,所以手动创建了Line2D代理——它们是虚拟的圆点,只用来在图例中展示每个数字对应的颜色,不会出现在主可视化区域。 - 图例位置调整:用
bbox_to_anchor把图例放在图的右侧外侧,避免遮挡聚类的数字,你可以根据需要调整这个参数改变图例位置。 - 我还简化了数据合并的代码,用
np.column_stack替代了多个DataFrame操作,让代码更简洁。
这样运行后,你就能得到带有对应数字颜色图例的t-SNE可视化图啦~
内容的提问来源于stack exchange,提问作者Nicolai Iversen
相关产品推荐
相关产品推荐

