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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 05:36:04