如何为t-SNE/UMAP降维的2D散点图按索引中的国家/年份着色
解决t-SNE降维后散点图按索引标签着色的问题
核心思路
先从原sector_features_的元组索引中提取country和year信息,合并到降维后的DataFrame中,就能直接用这些列作为着色依据,分别用matplotlib、seaborn、Altair实现可视化。
步骤1:提取原索引的标签信息
假设降维后的DataFrame名为tsne_results(替换成你实际的变量名),执行以下代码拆分索引并合并:
# 将原DataFrame的元组索引转为单独列 index_df = sector_features_.index.to_frame(index=False) # 合并索引信息与降维结果 tsne_combined = pd.concat([tsne_results, index_df], axis=1)
处理后tsne_combined会包含t-SNE的2个维度列(比如tsne_1、tsne_2),以及country和year列。
步骤2:用Matplotlib实现着色可视化
按国家着色
plt.figure(figsize=(10,8)) # 遍历每个国家绘制散点 for country in tsne_combined['country'].unique(): subset = tsne_combined[tsne_combined['country'] == country] plt.scatter(subset['tsne_1'], subset['tsne_2'], label=country, alpha=0.7) plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') plt.title('t-SNE 可视化(按国家着色)') plt.xlabel('t-SNE 维度1') plt.ylabel('t-SNE 维度2') plt.show()
按年份着色
支持分类着色和连续数值着色两种方式:
# 分类着色(年份作为离散标签) plt.figure(figsize=(10,8)) for year in tsne_combined['year'].unique(): subset = tsne_combined[tsne_combined['year'] == year] plt.scatter(subset['tsne_1'], subset['tsne_2'], label=year, alpha=0.7) plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') plt.title('t-SNE 可视化(按年份分类着色)') plt.xlabel('t-SNE 维度1') plt.ylabel('t-SNE 维度2') plt.show() # 连续数值着色(年份作为连续变量) plt.figure(figsize=(10,8)) scatter = plt.scatter(tsne_combined['tsne_1'], tsne_combined['tsne_2'], c=tsne_combined['year'], cmap='viridis', alpha=0.7) plt.colorbar(scatter, label='年份') plt.title('t-SNE 可视化(按年份连续着色)') plt.xlabel('t-SNE 维度1') plt.ylabel('t-SNE 维度2') plt.show()
步骤3:用Seaborn快速实现
Seaborn的scatterplot可以直接通过hue参数指定着色列,代码更简洁:
按国家着色
plt.figure(figsize=(10,8)) sns.scatterplot(data=tsne_combined, x='tsne_1', y='tsne_2', hue='country', alpha=0.7) plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') plt.title('t-SNE 可视化(按国家着色)') plt.show()
按年份着色
# 分类着色 plt.figure(figsize=(10,8)) sns.scatterplot(data=tsne_combined, x='tsne_1', y='tsne_2', hue='year', alpha=0.7) plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') plt.title('t-SNE 可视化(按年份分类着色)') plt.show() # 连续数值着色 plt.figure(figsize=(10,8)) sns.scatterplot(data=tsne_combined, x='tsne_1', y='tsne_2', hue='year', palette='viridis', alpha=0.7, size=1, sizes=(20,20)) plt.colorbar(label='年份') plt.title('t-SNE 可视化(按年份连续着色)') plt.show()
步骤4:用Altair实现交互式可视化
Altair支持交互式hover显示详情,适合探索数据:
按国家着色
import altair as alt alt.Chart(tsne_combined).mark_circle(size=60, opacity=0.7).encode( x='tsne_1', y='tsne_2', color=alt.Color('country:N', legend=alt.Legend(title='国家')), tooltip=['country', 'year', 'tsne_1', 'tsne_2'] ).properties( title='t-SNE 可视化(按国家着色)', width=600, height=500 ).interactive()
按年份着色
# 分类着色 alt.Chart(tsne_combined).mark_circle(size=60, opacity=0.7).encode( x='tsne_1', y='tsne_2', color=alt.Color('year:N', legend=alt.Legend(title='年份')), tooltip=['country', 'year', 'tsne_1', 'tsne_2'] ).properties( title='t-SNE 可视化(按年份分类着色)', width=600, height=500 ).interactive() # 连续数值着色 alt.Chart(tsne_combined).mark_circle(size=60, opacity=0.7).encode( x='tsne_1', y='tsne_2', color=alt.Color('year:Q', scale=alt.Scale(scheme='viridis'), legend=alt.Legend(title='年份')), tooltip=['country', 'year', 'tsne_1', 'tsne_2'] ).properties( title='t-SNE 可视化(按年份连续着色)', width=600, height=500 ).interactive()
内容的提问来源于stack exchange,提问作者Yannick Silva
相关产品推荐
相关产品推荐

