如何用Bokeh将np.array格式混淆矩阵绘制成分类热图?附报错解决
Bokeh 生成混淆矩阵热力图问题解决
问题背景
需要用Bokeh生成类似seaborn.heatmap的混淆矩阵可视化,原始数据如下:
short_labels = ["PERSON", "CITY", "COUNTRY", "O"] cm = np.array([[0.53951528, 0. , 0. , 0.46048472], [0.06407323, 0.18077803, 0. , 0.75514874], [0.06442577, 0.00560224, 0.08963585, 0.84033613], [ nan, nan, nan, nan]])
编写的代码运行后报错:
colors = ["#0B486B", "#79BD9A", "#CFF09E", "#79BD9A", "#0B486B", "#79BD9A", "#CFF09E", "#79BD9A", "#0B486B"] rand_cm_df = pd.DataFrame(rand_cm, columns=short_labels) rand_hm = figure(title=f"Title", toolbar_location=None, x_range=short_labels, y_range=short_labels, output_backend="svg") rand_hm.rect(source=ColumnDataSource(rand_cm_df), color=colors, width=1, height=1)
错误信息:
RuntimeError: Expected line_color, hatch_color and fill_color to reference fields in the supplied data source. When a 'source' argument is passed to a glyph method, values that are sequences (like lists or arrays) must come from references to data columns in the source. For instance, as an example: source = ColumnDataSource(data=dict(x=a_list, y=an_array)) p.circle(x='x', y='y', source=source, ...) # pass column names and a source Alternatively, *all* data sequences may be provided as literals as long as a source is *not* provided: p.circle(x=a_list, y=an_array, ...) # pass actual sequences and no source
解决方案
错误原因
- 使用
ColumnDataSource时,color参数必须是数据源中的列名,不能直接传入列表; - 原始代码未指定
rect的x、y坐标,无法定位每个方块; - 固定颜色列表无法实现基于混淆矩阵值的渐变效果,不符合热力图需求。
修正后完整代码
import numpy as np import pandas as pd from bokeh.plotting import figure, show, ColumnDataSource from bokeh.transform import linear_cmap from bokeh.models import ColorBar, LinearColorMapper # 原始数据 short_labels = ["PERSON", "CITY", "COUNTRY", "O"] cm = np.array([[0.53951528, 0. , 0. , 0.46048472], [0.06407323, 0.18077803, 0. , 0.75514874], [0.06442577, 0.00560224, 0.08963585, 0.84033613], [ nan, nan, nan, nan]]) # 1. 转换数据格式:宽表转长表,适配Bokeh的rect glyph cm_df = pd.DataFrame(cm, columns=short_labels, index=short_labels) melted_df = cm_df.reset_index().melt(id_vars='index', var_name='x', value_name='value') melted_df = melted_df.rename(columns={'index': 'y'}) # 2. 创建颜色映射:模拟seaborn热力图的渐变 color_mapper = LinearColorMapper(palette=["#CFF09E", "#79BD9A", "#0B486B"], low=melted_df['value'].min(), high=melted_df['value'].max(), nan_color='lightgray') # 处理NaN值 # 3. 创建figure hm = figure(title="混淆矩阵热力图", toolbar_location=None, x_range=short_labels, y_range=list(reversed(short_labels)), # 反转y轴,和seaborn对齐 output_backend="svg") # 4. 添加热力图方块 hm.rect(x='x', y='y', width=1, height=1, source=ColumnDataSource(melted_df), fill_color=linear_cmap(field_name='value', mapper=color_mapper), line_color='white') # 5. 添加颜色条 color_bar = ColorBar(color_mapper=color_mapper, label_standoff=12, location=(0,0)) hm.add_layout(color_bar, 'right') # 6. 调整样式 hm.xaxis.axis_label = '预测标签' hm.yaxis.axis_label = '真实标签' hm.grid.grid_line_color = None # 显示图形 show(hm)
关键说明
- 数据格式转换:用
pd.melt将宽格式的混淆矩阵转为长格式,每个行对应一个热力方块的x、y坐标和值; - 颜色映射:
LinearColorMapper根据混淆矩阵的值自动映射渐变颜色,nan_color处理最后一行的NaN值; - 坐标轴对齐:反转y轴让真实标签从上到下排列,和seaborn heatmap的布局一致;
- 样式优化:去掉网格线,添加轴标签和颜色条,还原seaborn热力图的视觉效果。
内容的提问来源于stack exchange,提问作者jedrix
相关产品推荐
相关产品推荐

