根据numpy数组长度自动生成matplotlib子图及修复索引越界报错
自适应imshow子图排布越界、变形问题修复
问题背景
需求为根据numpy数组dets的长度自动生成imshow彩色子图,排布规则如下:
- 若数组长度为完全平方数,生成正方形排布的子图矩阵
- 若数组长度为非完全平方数,在正方形排布基础上额外增加1行,放置剩余子图
初始实现代码如下:
data_f = np.random.rand(len(dets),2,5) dets = np.arange(-5,-0.75,0.25) x = np.array([1,5,6,3,8,9,2,3,10,12,3]) v = np.linspace(0,10,len(x)) square = np.sqrt(len(dets)) check_square = len(dets)%square non_square = 1 print(len(data_f)) if check_square == 0: nrows = int(np.sqrt(len(dets))) ncols = int(np.sqrt(len(dets))) else: nrows = int(np.sqrt(len(dets)))+non_square ncols = int(np.sqrt(len(dets))) fig, ax = plt.subplots(nrows, ncols, sharex='col', sharey='row') for i in range(nrows): for j in range(ncols): if i==0: im = ax[i,j].imshow(data_f[j],extent=(x.min(), x.max(), v.min(), v.max()),origin='lower',aspect='auto') else: im = ax[i,j].imshow(data_f[j+ncols*i],extent=(x.min(), x.max(), v.min(), v.max()),origin='lower',aspect='auto')
代码运行后绘制出17张子图就触发报错,子图存在异常挤压变形问题,初始运行输出效果如下:
抛出的报错信息:
--------------------------------------------------------------------------- IndexError Traceback (most recent call last) ~\AppData\Local\Temp\1/ipykernel_4560/3817292743.py in <module> 6 im = ax[i,j].imshow(data_f[j],extent=(x.min(), x.max(), v.min(), v.max()),origin='lower',aspect='auto') 7 else: ----> 8 im = ax[i,j].imshow(data_f[j+ncols*i],extent=(x.min(), x.max(), v.min(), v.max()),origin='lower',aspect='auto') 9 IndexError: index 17 is out of bounds for axis 0 with size 17
问题根因
代码存在3个核心问题:
- 变量顺序错误:
data_f初始化时调用了len(dets),但dets在后续行才定义,运行时会直接触发未定义错误。 - 判断与索引逻辑错误:用浮点数取模判断完全平方数存在精度风险,且双层循环遍历了所有子图位置(非平方数场景下总子图位置数大于有效数据长度),没有做边界判断,访问超出
data_f长度的索引时就会触发越界。 - 空轴未处理:最后一行多余的子图位置没有隐藏,加上没有做间距自适应调整,导致所有子图被挤压变形。
修复后代码
import numpy as np import matplotlib.pyplot as plt # 修正变量定义顺序,先声明dets再生成对应数据集 dets = np.arange(-5, -0.75, 0.25) n_plots = len(dets) data_f = np.random.rand(n_plots, 2, 5) x = np.array([1,5,6,3,8,9,2,3,10,12,3]) v = np.linspace(0,10,len(x)) # 用整数运算判断完全平方数,避免浮点数精度误差 side = int(np.floor(np.sqrt(n_plots))) if side * side == n_plots: nrows, ncols = side, side else: nrows, ncols = side + 1, side fig, ax = plt.subplots(nrows, ncols, sharex='col', sharey='row') # 将二维子图数组展平为一维,简化索引计算 ax_flat = ax.flatten() for idx in range(nrows * ncols): if idx < n_plots: # 有效索引位置正常绘制imshow ax_flat[idx].imshow( data_f[idx], extent=(x.min(), x.max(), v.min(), v.max()), origin='lower', aspect='auto' ) else: # 多余的空子图直接隐藏坐标轴 ax_flat[idx].axis('off') # 自动调整子图间距,解决挤压变形问题 fig.tight_layout() plt.show()
修复后逻辑不会触发索引越界,多余空轴自动隐藏,子图间距自适应调整,完全符合排布要求。
内容的提问来源于stack exchange,提问作者Kayla Jane Rodriguez
相关产品推荐
相关产品推荐

