You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Seaborn高效绘制距离矩阵?解决空图与内存溢出问题

问题解决:11000样本距离矩阵热图为空+Colab内存崩溃

一、热图为空的核心原因与解决

  1. 缺少图像渲染触发代码
    seaborn依赖matplotlib渲染图像,在Colab中必须显式调用plt.show()才能显示热图,否则绘图对象仅存在于内存中,不会输出到界面。

  2. 全量11000x11000矩阵无法可视化
    11000个样本的距离矩阵包含超过1.2亿个元素,直接绘制热图时,每个像素对应一个距离值,图像分辨率远不足以区分这么多数据,最终视觉上呈现“空图”。解决方式二选一:

    • 绘制聚类中心距离矩阵:用每个簇的均值代表整个簇,矩阵规模缩小为k×k(k为聚类数),清晰展示簇间差异。
    • 簇内采样后绘制:每个簇随机抽取少量样本(如50-100个),总样本量控制在几千以内,既保留簇结构,又能正常渲染热图。

二、Colab内存崩溃优化方案

  1. 避免原地修改DataFrame
    inplace=True会产生临时内存开销,改用赋值方式创建排序后的副本:

    df_sorted = df_labeled.sort_values(by=['cluster']).copy()
    
  2. 降低数据类型内存占用
    将距离矩阵从默认的float64转为float32(精度满足需求的前提下),内存占用直接减半:

    dist = pdist(df_sorted, metric).astype('float32')
    
  3. 延迟生成方阵
    pdist生成的是压缩格式的距离数组(内存占用为方阵的1/2),仅在绘图时转成squareform,用完立即释放:

    dist_square = squareform(dist)
    sns.heatmap(dist_square, cmap="mako")
    del dist_square  # 及时释放方阵内存
    
  4. 优化内存清理时机
    不要在绘图前删除距离数据,确保绘图完成后再清理变量,同时利用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 08:55:21