如何在Python中基于Pandas DataFrame绘制多组拼接类混淆矩阵?
实现多组拼接式类混淆矩阵图表
问题背景
我有如下Pandas DataFrame:
import pandas as pd df = pd.DataFrame({'cl1': ['A','A','A','A', 'A','A','A','A', 'D','D','D','D', 'D','D','D','D'], 'cl2': ['C','C','C','C', 'B','B','B','B', 'C','C','C','C', 'B','B','B','B'], 'p1p2': ['00','01','10','11', '00','01','10','11', '00','01','10','11', '00','01','10','11'], 'val':[1,2,3,4, 10,20,30,40, 5,6,7,8, 50,60,70,80]})
希望创建多组拼接式类混淆矩阵图表,请问如何在Python中实现?
实现方案
用matplotlib搭配seaborn热力图就能搞定,核心是按cl1和cl2分组,给每个组生成小混淆矩阵,再拼接成整体布局。
1. 导入依赖库
import matplotlib.pyplot as plt import seaborn as sns import pandas as pd
2. 预处理数据
把p1p2拆成行、列维度,方便转成混淆矩阵格式:
# 拆分p1p2的两位字符分别作为行和列 df['row'] = df['p1p2'].str[0] df['col'] = df['p1p2'].str[1] # 按cl1和cl2分组 groups = df.groupby(['cl1', 'cl2'])
3. 绘制拼接式热力图
创建2x2的子图网格,逐个绘制每个分组的混淆矩阵:
# 创建子图布局,这里4个分组用2行2列 fig, axes = plt.subplots(nrows=2, ncols=2, figsize=(10, 8)) axes = axes.flatten() # 把二维轴数组转一维,方便遍历 # 遍历每个分组画图 for idx, ((cl1_val, cl2_val), group_data) in enumerate(groups): # 把分组数据转成混淆矩阵的二维格式 matrix = group_data.pivot(index='row', columns='col', values='val') # 绘制热力图,显示数值,用蓝色系配色 sns.heatmap(matrix, annot=True, fmt='d', cmap='Blues', ax=axes[idx], cbar=False) # 设置子图标题 axes[idx].set_title(f'cl1={cl1_val}, cl2={cl2_val}') # 设置轴标签 axes[idx].set_xlabel('p2') axes[idx].set_ylabel('p1') # 自动调整子图间距,避免重叠 plt.tight_layout() plt.show()
可选优化
- 如果需要统一的颜色条,去掉
cbar=False,然后在代码末尾添加:# 添加共享颜色条 fig.subplots_adjust(right=0.85) cbar_ax = fig.add_axes([0.88, 0.15, 0.03, 0.7]) fig.colorbar(axes[0].collections[0], cax=cbar_ax) - 可以调整
figsize参数改变整体图表大小,修改cmap更换配色方案。
内容的提问来源于stack exchange,提问作者quant
相关产品推荐
相关产品推荐

