二进制热力图配色异常解决方案问询:如何基于逻辑值自动选择配色方案
解决二进制热力图全0/全1场景的配色异常问题
问题背景
我需要绘制多张二进制热力图,多数热力图同时包含0和1两种数值,但有少数热力图的数值全部为0或全部为1。在这类全0或全1的场景下,热力图的配色和色条显示会出现异常。我已经修改了代码实现了预期效果,但希望能获得更优的方案建议。
初始代码
import matplotlib.pyplot as plt from matplotlib.colors import LinearSegmentedColormap from matplotlib import gridspec import numpy as np import seaborn as snn import random M=np.zeros((2,5)) for j in range(2): for i in range(5): r=random.random() M[j,i]=r print(M) #[[0.61060519 0.04500793 0.74199826 0.22084509 0.31493589] #[0.3432519 0.59293327 0.50043671 0.07201856 0.65329049]] M0=M>0.5 M1=M>0 M2=M>2 M0 = M0.astype(float) M1 = M1.astype(float) M2 = M2.astype(float) fig=plt.figure(figsize=(10,7)) gs = fig.add_gridspec(ncols=1, nrows=4, height_ratios=[1,1, 1,1]) ax0=fig.add_subplot(gs[0]) ax1=fig.add_subplot(gs[1]) ax2=fig.add_subplot(gs[2]) ax3=fig.add_subplot(gs[3]) colors = ((1.0, 1.0, 0.0), (1, 0.0, 1.0)) cmap = LinearSegmentedColormap.from_list('Custom', colors, len(colors)) b=snn.heatmap(M,ax=ax0) b=snn.heatmap(M0,cmap=cmap,ax=ax1) colorbar = b.collections[0].colorbar colorbar.set_ticks([0.25,0.75]) colorbar.set_ticklabels(['0', '1']) b=snn.heatmap(M1,cmap=cmap,ax=ax2) colorbar = b.collections[0].colorbar colorbar.set_ticks([0.25,0.75]) colorbar.set_ticklabels(['0', '1']) b=snn.heatmap(M2,cmap=cmap,ax=ax3) colorbar = b.collections[0].colorbar colorbar.set_ticks([0.25,0.75]) colorbar.set_ticklabels(['0', '1']) fig.tight_layout() fig.savefig('a.png',dpi=300) plt.close()
已实现的修改后代码
import matplotlib.pyplot as plt from matplotlib.colors import LinearSegmentedColormap from matplotlib import gridspec import numpy as np import seaborn as snn import random M=np.zeros((2,5)) for j in range(2): for i in range(5): r=random.random() M[j,i]=r M0=M>0.5 M1=M>0 M2=M>2 M0 = M0.astype(float) M1 = M1.astype(float) M2 = M2.astype(float) fig=plt.figure(figsize=(10,7)) gs = fig.add_gridspec(ncols=1, nrows=4, height_ratios=[1,1, 1,1]) ax0=fig.add_subplot(gs[0]) ax1=fig.add_subplot(gs[1]) ax2=fig.add_subplot(gs[2]) ax3=fig.add_subplot(gs[3]) def chooseColorBar(ax,M): if np.mean(M)==0: colors = ((1.0, 1.0, 0.0),(1.0, 1.0, 0.0)) cmap = LinearSegmentedColormap.from_list('Custom', colors, len(colors)) b=snn.heatmap(M,cmap=cmap,ax=ax) colorbar=b.collections[0].colorbar colorbar.set_ticks([0]) colorbar.set_ticklabels(['0']) elif np.mean(M)==1: colors = ((1, 0.0, 1.0),(1, 0.0, 1.0)) cmap = LinearSegmentedColormap.from_list('Custom', colors, len(colors)) b=snn.heatmap(M,cmap=cmap,ax=ax) colorbar=b.collections[0].colorbar colorbar.set_ticks([1]) colorbar.set_ticklabels(['1']) else: colors = ((1.0, 1.0, 0.0), (1, 0.0, 1.0)) cmap = LinearSegmentedColormap.from_list('Custom', colors, len(colors)) b=snn.heatmap(M,cmap=cmap,ax=ax) colorbar=b.collections[0].colorbar colorbar.set_ticks([0.25,0.75]) colorbar.set_ticklabels(['0', '1']) b=snn.heatmap(M,ax=ax0) chooseColorBar(ax1,M0) chooseColorBar(ax2,M1) chooseColorBar(ax3,M2) fig.tight_layout() fig.savefig('a.png',dpi=300) plt.close()
优化方案建议
你的现有代码已经解决了核心问题,这里给你几个更简洁、健壮的优化方向:
1. 复用自定义色板,避免重复创建
不用在每个分支里重新定义色板,提前创建好基础的黄紫双色cmap。全0或全1的场景下,只需要通过vmin和vmax锁定数值范围,热力图会自动匹配对应的颜色,无需重新生成单一颜色的cmap:
# 提前定义好基础色板 base_colors = ((1.0, 1.0, 0.0), (1, 0.0, 1.0)) base_cmap = LinearSegmentedColormap.from_list('Custom', base_colors, len(base_colors))
2. 更精准的判断逻辑
用np.all()替代np.mean()来判断全0/全1,逻辑更直观,也避免极端情况下的误判(比如数据中存在非0非1值的意外情况):
if np.all(M == 0): # 全0逻辑 elif np.all(M == 1): # 全1逻辑 else: # 正常二进制逻辑
3. 提取公共逻辑,减少代码冗余
把色条设置的重复代码抽离出来,让函数更简洁。
优化后的完整代码
import matplotlib.pyplot as plt from matplotlib.colors import LinearSegmentedColormap from matplotlib import gridspec import numpy as np import seaborn as snn import random # 生成测试数据 M = np.random.rand(2, 5) M0 = (M > 0.5).astype(float) M1 = (M > 0).astype(float) M2 = (M > 2).astype(float) # 初始化画布 fig = plt.figure(figsize=(10,7)) gs = fig.add_gridspec(ncols=1, nrows=4, height_ratios=[1,1,1,1]) axes = [fig.add_subplot(gs[i]) for i in range(4)] # 提前定义基础色板 base_colors = ((1.0, 1.0, 0.0), (1, 0.0, 1.0)) base_cmap = LinearSegmentedColormap.from_list('Custom', base_colors, len(base_colors)) def plot_binary_heatmap(ax, data): if np.all(data == 0): # 全0场景:锁定数值范围为0,自动显示黄色 snn_plot = snn.heatmap(data, cmap=base_cmap, vmin=0, vmax=0, ax=ax) cbar = snn_plot.collections[0].colorbar cbar.set_ticks([0]) cbar.set_ticklabels(['0']) elif np.all(data == 1): # 全1场景:锁定数值范围为1,自动显示紫色 snn_plot = snn.heatmap(data, cmap=base_cmap, vmin=1, vmax=1, ax=ax) cbar = snn_plot.collections[0].colorbar cbar.set_ticks([1]) cbar.set_ticklabels(['1']) else: # 正常二进制场景 snn_plot = snn.heatmap(data, cmap=base_cmap, ax=ax) cbar = snn_plot.collections[0].colorbar cbar.set_ticks([0.25, 0.75]) cbar.set_ticklabels(['0', '1']) # 绘制各个子图 snn.heatmap(M, ax=axes[0]) plot_binary_heatmap(axes[1], M0) plot_binary_heatmap(axes[2], M1) plot_binary_heatmap(axes[3], M2) fig.tight_layout() fig.savefig('a.png', dpi=300) plt.close()
优化点说明
- 减少重复创建cmap:只生成一次基础色板,所有场景复用,提升代码效率
- 更可靠的判断:
np.all()直接检查所有元素是否为0/1,逻辑清晰无歧义 - 代码更简洁:用列表推导式初始化子图,抽离色条设置的重复逻辑,可读性更强
- 自动匹配颜色:通过
vmin和vmax锁定数值范围,让热力图自动使用色板对应位置的颜色,无需手动创建单一颜色的色板
内容的提问来源于stack exchange,提问作者ankit agrawal
相关产品推荐
相关产品推荐

