如何对Pandas DataFrame按两列排序并绘制对应热力图?
问题描述
现有如下结构的Pandas DataFrame:
| methods | sub_1 | sub_2 | sub_3 | sub_4 | ..... | distance | count_0 |
|---|---|---|---|---|---|---|---|
| m1 | 0 | 2 | 2 | 1 | ..... | 0 | 20 |
| m2 | 1 | 2 | 2 | 1 | ..... | 7 | 22 |
| m3 | 1 | 2 | 2 | 1 | ..... | 26 | 12 |
| m4 | 0 | 2 | 2 | 0 | ..... | 21 | 10 |
| m5 | 0 | 2 | 2 | 0 | ..... | 17 | 5 |
字段说明:
methods:聚类方法名称sub_*:客户ID,对应值为该方法给客户分配的聚类标签distance:该方法聚类结果与集成结果的距离count_0:被分配到0类的客户数量
需求是按distance和count_0排序后绘制热力图,使同颜色(同标签)区域聚集,但当前代码仅对行排序,无法实现同颜色聚集效果。当前代码如下:
fig, ax = plt.subplots(figsize=(40,6)) df = df.sort_values(by=['distance', 'count_0'], ascending=[True, True]) svm = sns.heatmap(df.iloc[:,0:546], ax=ax, xticklabels=False, yticklabels=True, linewidths=0.4, annot_kws={"size": 16}, cbar_kws={"shrink": 0.3}) ax.set_xlabel("Customer IDs", fontsize=30) ax.hlines(y = 1, xmin = 0, xmax = 550, colors = 'white', lw = 10) for i in range(2,10): ax.hlines(y = i, xmin = 0, xmax = 550, colors = 'white', lw = 4) figure = svm.get_figure() ax.tick_params(axis='both', which='major', labelsize=30, labelbottom = False, bottom=False, top = False, labeltop=True) plt.xticks(rotation=90) ax.tick_params(labelsize=30) plt.show()
解决方案
要实现同颜色区域聚集,除了对行按distance和count_0排序外,还需要对**客户列(sub_*)**进行聚类排序,让标签模式相似的客户列相邻。可以用层次聚类重新排列列的顺序:
步骤1:准备数据并排序行
先按需求对行排序,提取用于绘制热力图的标签数据(排除methods列):
import pandas as pd import seaborn as sns import matplotlib.pyplot as plt from scipy.cluster.hierarchy import linkage, leaves_list # 按distance和count_0排序行 df_sorted = df.sort_values(by=['distance', 'count_0'], ascending=[True, True]) # 提取sub_*列的标签数据(假设这些列是第1列到第545列,根据实际列数调整) label_data = df_sorted.iloc[:, 1:546]
步骤2:对列进行层次聚类排序
使用SciPy的层次聚类计算列的相似性,获取重新排列的列索引:
# 计算列的层次聚类链接(ward方法最小化类内方差,可替换为single/complete等) col_linkage = linkage(label_data.T, method='ward') # 获取排序后的列索引 col_order = leaves_list(col_linkage) # 重新排列标签数据的列 label_data_sorted = label_data.iloc[:, col_order]
步骤3:绘制优化后的热力图
用排序后的行和列数据绘制热力图,保留原有的格式设置:
fig, ax = plt.subplots(figsize=(40,6)) # 绘制热力图,使用排序后的列数据 svm = sns.heatmap(label_data_sorted, ax=ax, xticklabels=False, yticklabels=df_sorted['methods'], linewidths=0.4, annot_kws={"size": 16}, cbar_kws={"shrink": 0.3}) ax.set_xlabel("Customer IDs", fontsize=30) # 动态添加行分隔线(根据实际排序后的行数调整) ax.hlines(y = 1, xmin = 0, xmax = label_data_sorted.shape[1], colors = 'white', lw = 10) for i in range(2, len(df_sorted)+1): ax.hlines(y = i, xmin = 0, xmax = label_data_sorted.shape[1], colors = 'white', lw = 4) ax.tick_params(axis='both', which='major', labelsize=30, labelbottom = False, bottom=False, top = False, labeltop=True) plt.xticks(rotation=90) plt.show()
关键修正与说明
- 修正原代码错误:原代码
df.iloc[:,0:546]包含了methods列,现在改为仅提取标签列,避免非标签数据干扰热力图。 - 列聚类逻辑:通过层次聚类让标签模式相似的客户列相邻,从根本上实现同颜色区域的聚集。
- 动态分隔线:循环范围根据实际行数生成,避免固定范围导致的错误。
内容的提问来源于stack exchange,提问作者Rajesh Ahir
相关产品推荐
相关产品推荐

