如何用Seaborn高效绘制距离矩阵?解决空图与内存溢出问题
问题解决:11000样本距离矩阵热图为空+Colab内存崩溃
一、热图为空的核心原因与解决
缺少图像渲染触发代码
seaborn依赖matplotlib渲染图像,在Colab中必须显式调用plt.show()才能显示热图,否则绘图对象仅存在于内存中,不会输出到界面。全量11000x11000矩阵无法可视化
11000个样本的距离矩阵包含超过1.2亿个元素,直接绘制热图时,每个像素对应一个距离值,图像分辨率远不足以区分这么多数据,最终视觉上呈现“空图”。解决方式二选一:- 绘制聚类中心距离矩阵:用每个簇的均值代表整个簇,矩阵规模缩小为
k×k(k为聚类数),清晰展示簇间差异。 - 簇内采样后绘制:每个簇随机抽取少量样本(如50-100个),总样本量控制在几千以内,既保留簇结构,又能正常渲染热图。
- 绘制聚类中心距离矩阵:用每个簇的均值代表整个簇,矩阵规模缩小为
二、Colab内存崩溃优化方案
避免原地修改DataFrame
inplace=True会产生临时内存开销,改用赋值方式创建排序后的副本:df_sorted = df_labeled.sort_values(by=['cluster']).copy()降低数据类型内存占用
将距离矩阵从默认的float64转为float32(精度满足需求的前提下),内存占用直接减半:dist = pdist(df_sorted, metric).astype('float32')延迟生成方阵
pdist生成的是压缩格式的距离数组(内存占用为方阵的1/2),仅在绘图时转成squareform,用完立即释放:dist_square = squareform(dist) sns.heatmap(dist_square, cmap="mako") del dist_square # 及时释放方阵内存优化内存清理时机
不要在绘图前删除距离数据,确保绘图完成后再清理变量,同时利用gc.collect()强制回收内存。
优化后的代码示例
示例1:绘制聚类中心距离矩阵(推荐)
from scipy.spatial.distance import pdist, squareform import seaborn as sns import matplotlib.pyplot as plt import gc import pandas as pd def cluster_center_distance_matrix(df_labeled, metric="euclidean"): # 计算每个簇的中心(均值) cluster_centers = df_labeled.groupby('cluster').mean().reset_index() # 生成中心距离矩阵 dist = pdist(cluster_centers.drop('cluster', axis=1), metric).astype('float32') dist_square = squareform(dist) # 绘制热图 plt.figure(figsize=(8, 6)) sns.heatmap(dist_square, cmap="mako", xticklabels=cluster_centers['cluster'], yticklabels=cluster_centers['cluster']) plt.title('Cluster Center Distance Matrix') plt.show() # 清理内存 del dist, dist_square, cluster_centers gc.collect() cluster_center_distance_matrix(finalDf)
示例2:簇内采样后绘制热图
from scipy.spatial.distance import pdist, squareform import seaborn as sns import matplotlib.pyplot as plt import gc import pandas as pd def sampled_distance_matrix(df_labeled, sample_size=100, metric="euclidean"): # 每个簇随机采样sample_size个样本 sampled_df = df_labeled.groupby('cluster').apply(lambda x: x.sample(min(sample_size, len(x)))).reset_index(drop=True) sampled_df = sampled_df.sort_values(by=['cluster']).copy() # 生成距离矩阵 dist = pdist(sampled_df.drop('cluster', axis=1), metric).astype('float32') dist_square = squareform(dist) # 绘制热图 plt.figure(figsize=(12, 10)) sns.heatmap(dist_square, cmap="mako") plt.title(f'Sampled Distance Matrix (each cluster sampled {sample_size} samples)') plt.show() # 清理内存 del sampled_df, dist, dist_square gc.collect() sampled_distance_matrix(finalDf, sample_size=50)
内容的提问来源于stack exchange,提问作者JayJona
相关产品推荐
相关产品推荐

