np.where引发IndexError索引越界:数组维度不匹配问题求助
问题根源与解决方案
这个错误的核心原因很直接:你的im_pred和img的第二维度(列数)不匹配,导致索引超出了img的有效范围。
让我们一步步拆解:
数组形状不匹配
- 你生成的
im_pred是(224, 244)的数组,意味着它有224行、244列,列索引的范围是0~243。 - 而
img是(224, 224, 3)的数组,只有224列,列索引的有效范围是0~223。
- 你生成的
np.where返回的索引问题
当你执行np.where(im_pred==i)时,它返回两个数组:x:匹配元素的行索引,范围是0~223,这和img的行维度完全兼容。y:匹配元素的列索引,范围是0~243——这里的问题就来了,y中大于等于224的索引,在img的列维度里根本不存在,所以赋值时就会触发IndexError(你看到的227只是其中一个超出范围的索引例子)。
你拆分后打印的
np.max(y) = 243正好验证了这一点,这个值明显超过了img列维度的最大有效索引223。修复方案
根据你的需求,有两种常见的解决方式:- 修正
im_pred的形状:如果244是笔误,直接让它和img的前两个维度一致:im_pred = np.random.randint(0, num_classes, (224, 224)) # 将244改为224 - 切片匹配
img的列数:如果im_pred的244列是有意为之,只取它的前224列来对应img:for i in range(num_classes): x, y = np.where(im_pred[:, :224] == i) # 仅取im_pred的前224列 img[x, y, :] = [225, 0, 0]
- 修正
另外提个小建议:直接用x和y数组来索引img(img[x, y, :])比img[np.where(im_pred==i), :]更清晰,也更容易排查索引问题。
内容的提问来源于stack exchange,提问作者abhinavkulkarni
相关产品推荐
相关产品推荐

