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
相关产品推荐
相关产品推荐

