TensorFlow图像分类绘图时StridedSlice切片越界报错如何解决?
问题根源
你触发的越界报错核心原因是设置的BATCH_SIZE = 5,每个批次最多只有索引为0~4的5张图像,但绘图代码写了range(9),尝试访问第5个索引位置的图像时,超出了批次的维度范围。
排查验证步骤
- 数据集总共有30张图像,按20%比例拆分验证集后,训练集共24张,按batch size 5拆分,单个批次最多只有5张图像,不存在索引为5的元素
- 你代码中打印的
image_batch.shape输出应为(5, 256, 256, 3),第一个维度就是单批次的图像数量,和你设置的batch size完全匹配
解决方法
方法1:调整循环次数适配batch size
把绘图循环的迭代次数改成和batch size一致,同步调整子图布局即可:
plt.figure(figsize=(10, 10)) for images, labels in train_ds.take(1): for i in range(BATCH_SIZE): ax = plt.subplot(1, BATCH_SIZE, i + 1) plt.imshow(images[i].numpy().astype('uint8')) plt.title(class_names[labels[i]]) plt.axis("off")
方法2:修改batch size满足绘图需求
如果确实需要展示9张图像,直接把BATCH_SIZE调整为不小于9的数值即可,优先选择2的幂次可以提升训练效率:
BATCH_SIZE = 16
方法3:拼接多批次图像凑够展示数量
如果不想修改训练用的batch size,可以取多个批次的图像拼接后再选择9张展示:
plt.figure(figsize=(10, 10)) all_images = [] all_labels = [] # 取2个批次共10张图像,足够选出9张用于展示 for images, labels in train_ds.take(2): all_images.extend(images.numpy()) all_labels.extend(labels.numpy()) for i in range(9): ax = plt.subplot(3, 3, i + 1) plt.imshow(all_images[i].astype('uint8')) plt.title(class_names[all_labels[i]]) plt.axis("off")
内容的提问来源于stack exchange,提问作者Aditya Aryan
相关产品推荐
相关产品推荐

