如何基于matplotlib生成符合要求的单通道numpy格式掩码用于目标检测
解决方案
你可以直接基于matplotlib的子图坐标计算生成目标numpy数组,不需要对生成的彩色掩码做颜色聚类,既高效也不会出现颜色识别误差,实现逻辑如下:
实现思路
- 保留你原有的
gridspec、on_off参数,生成掩码画布后先触发画布渲染,获取准确的子图坐标 - 初始化和画布像素尺寸一致的全0单通道numpy数组
- 遍历所有启用的子图,将子图的相对坐标转换为绝对像素坐标,给对应区域的数组赋值递增的整数标签即可
完整代码示例
import numpy as np import matplotlib.pyplot as plt import random # 此处保留你原有依赖定义:colors列表、draw_subplot函数实现 colors = ['#ff0000', '#00ff00', '#0000ff', '#ffff00', '#ff00ff', '#00ffff'] def draw_subplot(ax, number, number_buffer): ax.text(0.5, 0.5, str(number), fontsize=20) return (0, 0) # ------------------- 原图像生成逻辑 ------------------- on_off=[] gridspec=random.sample([1,1,1,1,1,1,1,1,1,1,1,1,2,2,2,2,2,2,3,4,5,6,7],3) fig, axes=plt.subplots(2,3,figsize=(7,7),gridspec_kw={'width_ratios': gridspec}) if random.random()<0.3: fig.set_facecolor(random.sample(colors,1)[0]) subplot_dict={} for i,ax in enumerate(axes.flatten()): subplot_dict[i]={} if random.random()<0.2: on_off.append(0) ax.axis('off') continue red, black=draw_subplot(ax, number=random.sample(range(1,10),1)[0], number_buffer=random.sample(range(0,3),1)[0]) on_off.append(1) subplot_dict[i]['red']=red subplot_dict[i]['black']=black plt.show() # ------------------- 单通道掩码数组生成逻辑 ------------------- mask_fig, mask_axes = plt.subplots(2,3,figsize=(7,7),gridspec_kw={'width_ratios': gridspec}) # 先触发画布渲染,保证坐标计算准确 mask_fig.canvas.draw() # 获取画布像素尺寸,初始化全0单通道掩码 width, height = mask_fig.get_size_inches() * mask_fig.dpi mask_arr = np.zeros((int(height), int(width)), dtype=np.uint8) current_label = 1 for i, ax in enumerate(mask_axes.flatten()): if not on_off[i]: ax.remove() continue # 子图相对坐标转绝对像素坐标 bbox = ax.get_window_extent() # 适配numpy数组坐标原点(左上角)和matplotlib坐标原点(左下角)的差异 x0, y0 = int(bbox.x0), int(height - bbox.y1) x1, y1 = int(bbox.x1), int(height - bbox.y0) # 给子图对应区域赋值标签 mask_arr[y0:y1, x0:x1] = current_label current_label += 1 # 保留你原有掩码可视化逻辑 ax.set_facecolor(colors[i]) for spine in ax.spines.values(): spine.set_edgecolor(colors[i]) ax.tick_params(axis='both', which='both',bottom=False, left=False,labelbottom=False, labelleft=False) plt.show() # 输出的mask_arr即为符合要求的单通道整数掩码 print(mask_arr)
注意事项
- 必须先调用
canvas.draw()再获取子图坐标,避免matplotlib懒加载导致坐标计算偏差 - 标签默认按子图行优先遍历顺序从1开始递增,你可以根据需求自行调整赋值逻辑
- 调整画布
figsize和dpi参数即可更改掩码分辨率,无需修改核心逻辑
内容的提问来源于stack exchange,提问作者vineeth venugopal
相关产品推荐
相关产品推荐

