如何用Matplotlib自动适配行列显示数量不固定的图片?
问题描述
我想用matplotlib的fig.add_subplot方法展示图片,但运行时出现两个报错:
报错1:子图编号超出范围
Traceback (most recent call last): File "/home/---/Documents/---/---/dataset.py", line 134, in <module> display_dicom(dicom,target["mask"]) File "/home/---/Documents/---/---/dataset.py", line 123, in display_dicom fig.add_subplot(rows,cols,i+2) File "/home/---/.pyenv/versions/3.8.13/envs/---/lib/python3.8/site-packages/matplotlib/figure.py", line 745, in add_subplot ax = subplot_class_factory(projection_class)(self, *args, **pkw) File "/home/---/.pyenv/versions/3.8.13/envs/---/lib/python3.8/site-packages/matplotlib/axes/_subplots.py", line 36, in __init__ self.set_subplotspec(SubplotSpec._from_subplot_args(fig, args)) File "/home/---/.pyenv/versions/3.8.13/envs/---/lib/python3.8/site-packages/matplotlib/gridspec.py", line 612, in _from_subplot_args raise ValueError( ValueError: num必须满足1 <= num <= 2,不能为3
报错2:索引越界
Traceback (most recent call last): File "/home/---/Documents/---/---/dataset.py", line 133, in <module> display_dicom(dicom,target["mask"]) File "/home/---/Documents/---/---/dataset.py", line 123, in display_dicom plt.imshow(mask[i],cmap=plt.cm.bone) IndexError:索引69超出了维度0的范围,该维度大小为69
我需要自动计算plt.figure的行数和列数,求不会导致代码崩溃的计算公式。目前手动指定行列时代码正常运行,我的尝试代码如下:
def display_dicom(dicom,mask): count,width,height = mask.shape if count == 0: count = 1 fig = plt.figure(figsize=(10,10)) rows= int(math.sqrt(count)+1) cols = int(math.sqrt(count)+1) fig.add_subplot(rows,cols, 1) plt.imshow(dicom, cmap=plt.cm.bone) # set the color map to bone plt.title("dicom") for i in range(2,count+2): fig.add_subplot(rows,cols,i) plt.imshow(mask[i],cmap=plt.cm.bone) plt.title(f"mask {i+1}") plt.show()
期望展示效果:
- 83张图片的网格布局展示
- 10张图片的网格布局展示
- 2张图片的网格布局展示
解决方案
报错原因
- 子图编号超出范围:原代码仅基于mask数量计算行列,未算上1张dicom图,导致网格总容量(rows*cols)小于实际需要展示的子图总数,触发编号越界错误。
- 索引越界:循环中使用的
i值从2开始,与mask的索引(0到count-1)不匹配,导致访问不存在的mask索引。
稳定的行列计算公式
基于**总子图数(1张dicom + count张mask)**计算行列,确保网格容量足够容纳所有子图:
total_plots = count + 1 # 处理count=0时单独设为1 cols = math.ceil(math.sqrt(total_plots)) rows = math.ceil(total_plots / cols)
该公式能保证rows * cols >= total_plots,不会出现子图编号超出范围的问题。
修正后的代码
import math import matplotlib.pyplot as plt def display_dicom(dicom, mask): count, width, height = mask.shape # 处理mask数量为0的情况 if count == 0: total_plots = 1 else: total_plots = count + 1 # 1张dicom + count张mask # 计算合适的行列数 cols = math.ceil(math.sqrt(total_plots)) rows = math.ceil(total_plots / cols) fig = plt.figure(figsize=(10,10)) # 显示dicom图 ax1 = fig.add_subplot(rows, cols, 1) # Tensor转numpy,若在GPU上需先调用.cpu() ax1.imshow(dicom.cpu().numpy(), cmap=plt.cm.bone) ax1.set_title("dicom") ax1.axis('off') # 关闭坐标轴优化显示 # 显示所有mask for idx in range(count): subplot_idx = idx + 2 # 从第2个位置开始排列mask ax = fig.add_subplot(rows, cols, subplot_idx) ax.imshow(mask[idx].cpu().numpy(), cmap=plt.cm.bone) ax.set_title(f"mask {idx+1}") ax.axis('off') plt.tight_layout() # 自动调整子图间距,避免标题重叠 plt.show()
关键修正点
- 行列计算逻辑:基于总子图数计算,确保网格容量足够。
- 索引匹配:用
idx从0到count-1遍历mask,对应正确的张量索引。 - Tensor转Numpy:matplotlib无法直接显示torch张量,需转换为numpy数组,GPU张量需先移至CPU。
- 显示优化:添加
axis('off')隐藏坐标轴,tight_layout()自动调整子图间距,提升视觉效果。
内容的提问来源于stack exchange,提问作者Alican Kartal
相关产品推荐
相关产品推荐

