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

matplotlib中imshow结合自定义非线性colormap报错如何解决

问题描述

我参考相关方案自定义了非线性colormap,希望让0-50区间的颜色粒度比50-100区间更细,用于展示ConfusionMatrixDisplay对象,代码如下:

from sklearn.datasets import make_classification
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay
from sklearn.model_selection import train_test_split
from sklearn.svm import SVC
from matplotlib.colors import LinearSegmentedColormap
import matplotlib.pyplot as plt
import numpy as np

plt.rcParams["figure.figsize"] = (15, 15)
font = {'family' : 'DejaVu Sans',
    'weight' : 'bold',
    'size'   : 22}
plt.rc('font', **font)


class nlcmap(LinearSegmentedColormap):
    def __init__(self, cmap, levels):
        self.cmap = cmap
        self.N = cmap.N
        self.monochrome = self.cmap.monochrome
        self.levels = np.asarray(levels, dtype='float64')
        self._x = self.levels
        self.levmax = self.levels.max()
        self.transformed_levels = np.linspace(0.0, self.levmax, len(self.levels))

    def __call__(self, xi, alpha=1.0, **kw):
        yi = np.interp(xi, self._x, self.transformed_levels)
        return self.cmap(yi / self.levmax, alpha)


levels = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 100]

cmap_nonlin = nlcmap(plt.cm.viridis, levels)
X, y = make_classification(random_state=0)
X_train, X_test, y_train, y_test = train_test_split(X, y,
                                            random_state=0)
clf = SVC(random_state=0)
clf.fit(X_train, y_train)
SVC(random_state=0)
predictions = clf.predict(X_test)
cm = confusion_matrix(y_test, predictions, labels=clf.classes_)
disp = ConfusionMatrixDisplay(confusion_matrix=cm,
                             display_labels=clf.classes_)
lin_cmap = plt.cm.viridis
levels = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 100]
cmap_nonlin = nlcmap(plt.cm.viridis, levels)
fig, ax = plt.subplots()
im = disp.plot(cmap=cmap_nonlin, colorbar=False)
disp.ax_.get_images()[0].set_clim(0, 100)
disp.figure_.colorbar(disp.im_, orientation="horizontal", pad=0.1)
plt.savefig("test.png")

运行后触发如下报错:

Traceback (most recent call last):
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/backends/backend_macosx.py", line 61, in _draw
    self.figure.draw(renderer)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/artist.py", line 41, in draw_wrapper
    return draw(artist, renderer, *args, **kwargs)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/figure.py", line 1864, in draw
    renderer, self, artists, self.suppressComposite)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/image.py", line 131, in _draw_list_compositing_images
    a.draw(renderer)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/artist.py", line 41, in draw_wrapper
    return draw(artist, renderer, *args, **kwargs)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/cbook/deprecation.py", line 411, in wrapper
    return func(*inner_args, **inner_kwargs)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/axes/_base.py", line 2747, in draw
    mimage._draw_list_compositing_images(renderer, self, artists)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/image.py", line 131, in _draw_list_compositing_images
    a.draw(renderer)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/artist.py", line 41, in draw_wrapper
    return draw(artist, renderer, *args, **kwargs)
  File "/Users/me/anaconda3/envs/myenv/lib/python3.6/site-packages/matplotlib/image.py", line 646, in draw
    renderer.draw_image(gc, l, b, im)
TypeError: Cannot cast array data from dtype('float64') to dtype('uint8') according to the rule 'safe'

问题可独立复现,不依赖sklearn组件,复现代码如下:

fig, ax = plt.subplots()
ax.imshow(np.array([[10, 15], [20, 30]]), cmap=cmap_nonlin)

要求:在不修改原始数据、仅调整colormap的前提下解决问题。


问题原因

自定义的nlcmap继承LinearSegmentedColormap后,没有遵循Matplotlib的Colormap接口规范:Matplotlib底层渲染图像时,期望colormap输出为0-255范围的uint8类型RGBA数组,而当前__call__方法实现返回的是0-1范围的float64类型数组,触发了类型转换安全检查失败。


可行解决方案

两种方案均不需要修改原始数据,直接替换原有colormap定义即可正常运行。

方案1:修改自定义colormap类的实现,适配接口要求

只需要调整nlcmap的初始化和__call__方法,做接口兼容即可,原有业务逻辑无需改动:

class nlcmap(LinearSegmentedColormap):
    def __init__(self, cmap, levels):
        # 正确初始化父类属性
        super().__init__(cmap.name, cmap._segmentdata, cmap.N, cmap._gamma)
        self.cmap = cmap
        self.levels = np.asarray(levels, dtype='float64')
        self._x = self.levels
        self.levmax = self.levels.max()
        self.transformed_levels = np.linspace(0.0, self.levmax, len(self.levels))

    def __call__(self, xi, alpha=1.0, bytes=False, **kw):
        yi = np.interp(xi, self._x, self.transformed_levels)
        # 透传bytes参数,原生cmap会自动匹配返回uint8/float64类型的RGBA值
        return self.cmap(yi / self.levmax, alpha=alpha, bytes=bytes)

方案2:无需自定义类,直接用原生接口生成非线性colormap

用LinearSegmentedColormap.from_list直接构造符合粒度要求的colormap,代码更简洁,兼容性更好:

levels = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 100]
max_level = max(levels)
# 把自定义level映射到0-1的归一化区间
norm_levels = [l / max_level for l in levels]
# 从viridis中提取对应位置的颜色
colors = plt.cm.viridis(norm_levels)
# 生成非线性colormap
cmap_nonlin = LinearSegmentedColormap.from_list("nonlin_viridis", list(zip(norm_levels, colors)))

内容的提问来源于stack exchange,提问作者jeandut

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 04:00:04