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

如何用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

解决方案

错误原因

  1. 使用ColumnDataSource时,color参数必须是数据源中的列名,不能直接传入列表;
  2. 原始代码未指定rect的x、y坐标,无法定位每个方块;
  3. 固定颜色列表无法实现基于混淆矩阵值的渐变效果,不符合热力图需求。

修正后完整代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 23:01:37