如何根据DataFrame另一列的值为Matplotlib直方图着色
解决方案
核心思路
避免逐行循环,改为按直方图分组(bin)循环(循环次数等于bin数量,通常远小于数据行数),对每个bin内的样本按distance值映射颜色,再堆叠绘制,既保证效率,又实现单bin多颜色的效果。
完整代码
import pandas as pd import matplotlib.pyplot as plt import numpy as np # 创建目标DataFrame df_dict = { "test_predictions": [0.1, 0.1, 0.2, 0.2, 0.3, 0.3, 0.4, 0.4, 0.4, 0.4, 0.4, 0.5, 0.5, 0.6, 0.6, 0.6, 0.7, 0.7, 0.7, 0.7, 0.7, 0.8, 0.8, 0.8, 0.9, 0.9, 0.9, 0.9, 0.9, 0.9, 0.9, 0.9], "y_true": [0, 0, 0, 1, 1, 1, 0, 0, 1, 1, 0, 1, 0, 1, 0, 0, 0, 1, 0, 1, 0, 1, 1, 1, 0, 0, 0, 1, 1, 0, 0, 1], "distance" : [-0.1, -0.09, -0.08, -0.08, -0.07, -0.05, -0.05, -0.05, -0.05, -0.05, -0.04, -0.04, -0.04, -0.03, -0.02, -0.01, 0.01, 0.01, 0.01, 0.02, 0.03, 0.03, 0.04, 0.05, 0.05, 0.06, 0.06, 0.07, 0.08, 0.08, 0.09, 0.1] } df = pd.DataFrame(df_dict) # 初始化图形与坐标轴 fig, ax1 = plt.subplots(figsize=(10,6)) # 绘制折线/散点图(将原折线改为散点,避免连线混乱) ax1.plot([0, 1], [0, 1], color="red", linestyle=":", label="Perfect Model") ax1.scatter(df['test_predictions'], df['y_true'], label="NN3", color='blue', s=20) ax1.set_xlabel('Test Predictions') ax1.set_ylabel('True Labels') ax1.legend(loc='upper left') # 创建双y轴用于直方图 ax2 = ax1.twinx() # 设置直方图参数 num_bins = 10 bins = np.linspace(df['test_predictions'].min(), df['test_predictions'].max(), num_bins) bin_width = np.diff(bins)[0] # 归一化distance值,适配颜色映射范围 norm = plt.Normalize(df['distance'].min(), df['distance'].max()) cmap = plt.cm.viridis # 一次性获取所有样本的bin索引 bin_indices = np.digitize(df['test_predictions'], bins) # 按bin循环处理(仅num_bins次循环,效率极高) for i in range(1, num_bins): # 提取当前bin内的所有distance值 current_distances = df.loc[bin_indices == i, 'distance'] if len(current_distances) == 0: continue # 生成当前bin内每个样本的颜色 colors = cmap(norm(current_distances)) # 堆叠绘制每个样本对应的小条,bottom控制堆叠位置 ax2.bar(bins[i-1], np.ones(len(current_distances)), width=bin_width, bottom=np.arange(len(current_distances)), color=colors, alpha=0.7) # 设置直方图轴属性 ax2.set_ylabel('Sample Count') ax2.set_ylim(0, df['test_predictions'].value_counts().max() + 1) # 添加颜色条,标注distance与颜色的对应关系 sm = plt.cm.ScalarMappable(norm=norm, cmap=cmap) sm.set_array([]) fig.colorbar(sm, ax=ax2, label='Distance') # 调整布局,避免标签重叠 plt.tight_layout() plt.show()
关键改进点
- 效率优化:将逐行循环(O(n))改为按bin循环(O(num_bins)),num_bins通常远小于数据行数,10000行数据仅需10次循环,性能提升显著。
- 堆叠效果实现:通过
bottom=np.arange(len(current_distances))为每个bin内的样本设置递增的底部位置,实现堆叠式直方图。 - 颜色映射标准化:使用
plt.Normalize统一将distance值映射到[0,1]区间,确保颜色映射的准确性,不受distance取值范围影响。 - 可视化优化:将原折线图改为散点图,避免连线导致的视觉混乱;添加颜色条明确颜色与
distance的对应关系。
内容的提问来源于stack exchange,提问作者Caesar
相关产品推荐
相关产品推荐

