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

Keras训练VGG16遇切片索引越界错误,求技术帮助

问题:VGG16训练Imagenette数据集时出现切片索引越界错误

我首次使用Keras和TensorFlow,尝试让VGG16模型在imagenette数据集上训练,但遇到了切片索引越界错误,已经调试很久没解决,希望得到帮助。我已经按照VGG16官方文档将图片调整到了正确尺寸。

我的代码如下:

tfds_name = 'imagenette'
(ds_train, ds_validation), ds_info= tfds.load(
    name=tfds_name,
    split=['train', 'validation'],
    with_info = True,
    as_supervised=True)

#model from assignment link
ourModel = tf.keras.applications.VGG16(
    include_top=True,                 #3 fill layers on top
    weights="imagenet",               #use imagenet
    input_tensor=None,                #use another layer as input 
    input_shape=None,                 #inly set if include to false 
    pooling=None,                     #use with include top false 
    classes=1000,                     #number of classes to set, we use imagenet values
    classifier_activation="softmax",  # classifier on input can only be none or softmax on pretrained 
)

#make it so layers frozen 
#for layer in ourModel.layers[:-1]:
#  layer.trainable = False

loss_fn = tf.keras.losses.SparseCategoricalCrossentropy()
ourModel.compile(optimizer="adam",
              loss=loss_fn,
              metrics=['accuracy'])

def reshape(img,label):
  img = tf.cast(img, tf.float32)
  img = tf.image.resize(img, (224,224))
  resize_image = tf.reshape(img, [-1, 224, 224, 3])
  resize_image = preprocess_input(resize_image)
  return resize_image, label

ds_train = ds_train.map(reshape)
ds_validation = ds_validation.map(reshape)
ourModel.fit(ds_train,
             epochs=10,
             validation_data = ds_validation)

错误信息:

ValueError: in user code:

    File "/usr/local/lib/python3.7/dist-packages/keras/engine/training.py", line 1051, in train_function  *
        return step_function(self, iterator)
    File "/usr/local/lib/python3.7/dist-packages/keras/engine/training.py", line 1040, in step_function  **
        outputs = model.distribute_strategy.run(run_step, args=(data,))
    File "/usr/local/lib/python3.7/dist-packages/keras/engine/training.py", line 1030, in run_step  **
        outputs = model.train_step(data)
    File "/usr/local/lib/python3.7/dist-packages/keras/engine/training.py", line 890, in train_step
        loss = self.compute_loss(x, y, y_pred, sample_weight)
    File "/usr/local/lib/python3.7/dist-packages/keras/engine/training.py", line 949, in compute_loss
        y, y_pred, sample_weight, regularization_losses=self.losses)
    File "/usr/local/lib/python3.7/dist-packages/keras/engine/compile_utils.py", line 212, in __call__
        batch_dim = tf.shape(y_t)[0]

    ValueError: slice index 0 of dimension 0 out of bounds. for '{{node strided_slice}} = StridedSlice[Index=DT_INT32, T=DT_INT32, begin_mask=0, ellipsis_mask=0, end_mask=0, new_axis_mask=0, shrink_axis_mask=1](Shape, strided_slice/stack, strided_slice/stack_1, strided_slice/stack_2)' with input shapes: [0], [1], [1], [1] and with computed input tensors: input[1] = <0>, input[2] = <1>, input[3] = <1>.

问题根源及解决步骤

  1. 输入维度冗余:你在reshape函数里给单张图片额外增加了batch维度(tf.reshape(img, [-1, 224, 224, 3])),导致每个样本输入形状变成(1,224,224,3),但模型期望单样本输入为(224,224,3),批量输入为(batch_size,224,224,3),多余维度引发损失计算时的索引错误。
  2. 缺少批次处理:TensorFlow Dataset训练时需要按批次喂入数据,你没有对数据集执行batch()操作,导致模型接收的是单个带冗余维度的样本,进一步触发错误。

修改后的完整代码

tfds_name = 'imagenette'
(ds_train, ds_validation), ds_info= tfds.load(
    name=tfds_name,
    split=['train', 'validation'],
    with_info = True,
    as_supervised=True)

ourModel = tf.keras.applications.VGG16(
    include_top=True,
    weights="imagenet",
    input_tensor=None,
    input_shape=None,
    pooling=None,
    classes=1000,
    classifier_activation="softmax",
)

# 如需冻结特征层,取消下面注释
# for layer in ourModel.layers[:-1]:
#   layer.trainable = False

loss_fn = tf.keras.losses.SparseCategoricalCrossentropy()
ourModel.compile(optimizer="adam",
              loss=loss_fn,
              metrics=['accuracy'])

def preprocess(img,label):
  img = tf.cast(img, tf.float32)
  img = tf.image.resize(img, (224,224))
  # 调用VGG16专属预处理函数,避免导入错误
  img = tf.keras.applications.vgg16.preprocess_input(img)
  return img, label

# 添加批次处理和预取,提升训练效率
ds_train = ds_train.map(preprocess).batch(32).prefetch(tf.data.AUTOTUNE)
ds_validation = ds_validation.map(preprocess).batch(32).prefetch(tf.data.AUTOTUNE)

ourModel.fit(ds_train,
             epochs=10,
             validation_data = ds_validation)

额外提示

Imagenette数据集只有10个类别,但你当前使用的VGG16默认适配1000类ImageNet数据集,会导致分类结果与标签不匹配,最终准确率极低。建议替换顶层分类器适配10类任务,示例代码如下:

# 加载不含顶层分类器的VGG16特征提取器
base_model = tf.keras.applications.VGG16(
    include_top=False,
    weights="imagenet",
    input_shape=(224,224,3),
    pooling='avg'
)
base_model.trainable = False  # 冻结特征层

# 构建适配10类的新模型
inputs = tf.keras.Input(shape=(224,224,3))
x = base_model(inputs, training=False)
outputs = tf.keras.layers.Dense(10, activation='softmax')(x)
ourModel = tf.keras.Model(inputs, outputs)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 22:05:44