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

如何用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张图片的网格布局展示

解决方案

报错原因

  1. 子图编号超出范围:原代码仅基于mask数量计算行列,未算上1张dicom图,导致网格总容量(rows*cols)小于实际需要展示的子图总数,触发编号越界错误。
  2. 索引越界:循环中使用的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()

关键修正点

  1. 行列计算逻辑:基于总子图数计算,确保网格容量足够。
  2. 索引匹配:用idx从0到count-1遍历mask,对应正确的张量索引。
  3. Tensor转Numpy:matplotlib无法直接显示torch张量,需转换为numpy数组,GPU张量需先移至CPU。
  4. 显示优化:添加axis('off')隐藏坐标轴,tight_layout()自动调整子图间距,提升视觉效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 14:15:33