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

运行VGG训练代码遇InvalidArgumentError,请求排查与修改建议

VGG16训练报错排查(InvalidArgumentError: Graph execution error)

问题背景

运行VGG16迁移训练代码时触发InvalidArgumentError: Graph execution error,无法定位是数据集还是代码逻辑问题,附上完整代码及训练异常情况。

完整代码

import tensorflow as tf
from tensorflow import keras
from keras import Sequential,optimizers
from keras.models import Model
from keras.applications.vgg16 import VGG16
from keras.preprocessing.image import ImageDataGenerator
from keras.preprocessing import image
from keras.applications.vgg16 import VGG16
from keras.layers import Dense,Conv2D,MaxPooling2D,Flatten,BatchNormalization,Dropout,Convolution2D,ZeroPadding2D

# generators
trdata = ImageDataGenerator(rescale = 1./255)
traindata = trdata.flow_from_directory(directory="/content/splited_data/train",
    batch_size=64,
    target_size=(224,224)
)

tsdata = ImageDataGenerator(rescale = 1./255)
testdata = tsdata.flow_from_directory(directory="/content/splited_data/test",
    batch_size=64,
    target_size=(224,224)
)

# create VGG Model
model = VGG16(weights="imagenet",include_top=True)

model.summary()

for layers in (model.layers)[:19]:
  print(layers)
  layers.trainable =False

X = model.layers[-2].output
predictions = Dense(3,activation="softmax")(X)
model_final = Model(inputs = model.input,outputs = predictions)

model_final.compile(optimizer=optimizers.SGD(lr=0.0001,momentum=0.9),loss="categorical_crossentropy",metrics=['accuracy',tf.keras.metrics.Precision(),tf.keras.metrics.Recall()])

model_final.summary()

# from keras.callbacks import ModelCheckpoint,EarlyStopping
# checkpoint =ModelCheckpoint("vgg16_1.h5",monitor="val_accuracy",verbose=1,save_best_only=True,save_weights_only=False,mode="auto",save_freq=1)
# early =EarlyStopping(monitor="val_accuracy",min_delta=0,patience=40,verbose=1,mode="auto")
# hist=model_final.fit_generator(generator=traindata,steps_per_epoch=2,epochs=10,validation_data=testdata,validation_steps=1,callbacks=[checkpoint,early])
hist=model_final.fit_generator(generator=traindata,steps_per_epoch=2,epochs=10,validation_data=testdata)

异常情况

训练轮次输出异常,触发图执行错误。

错误排查与修改方案

1. 替换废弃API

fit_generator在TensorFlow 2.1及以上版本已被废弃,改用fit方法自动处理生成器输入:

# 替换原fit_generator代码
hist = model_final.fit(
    traindata,
    steps_per_epoch=2,
    epochs=10,
    validation_data=testdata
)

2. 校验数据集类别匹配

  • 确认/content/splited_data/train和/test目录下的子文件夹数量为3个(对应输出层Dense(3)的类别数),若实际类别数不匹配会触发维度错误。
  • 查看flow_from_directory输出的Found X images belonging to Y classes日志,确保Y=3。

3. 修正冻结层范围

VGG16完整结构共19层,model.layers[:19]会冻结所有层(包括需训练的顶层),正确做法是冻结前18层,保留顶层可微调:

# 冻结前18层,保留顶层训练
for layers in model.layers[:18]:
    layers.trainable = False

4. 统一API使用规范

  • 避免混用keras和tf.keras优化器,统一使用tf.keras.optimizers:
optimizer=tf.keras.optimizers.SGD(learning_rate=0.0001, momentum=0.9)
  • 显式实例化metrics并命名,避免冲突:
metrics=[
    'accuracy',
    tf.keras.metrics.Precision(name='precision'),
    tf.keras.metrics.Recall(name='recall')
]

5. 完善数据生成器配置

  • 显式指定class_mode='categorical'确保多分类模式匹配:
traindata = trdata.flow_from_directory(
    directory="/content/splited_data/train",
    batch_size=64,
    target_size=(224,224),
    class_mode='categorical'
)
testdata = tsdata.flow_from_directory(
    directory="/content/splited_data/test",
    batch_size=64,
    target_size=(224,224),
    class_mode='categorical'
)
  • 调整steps_per_epoch为len(traindata),确保遍历完整训练集,避免数据耗尽错误。

6. 优化迁移学习初始化

迁移学习时建议使用include_top=False初始化VGG16,手动添加顶层分类器,减少冗余计算:

# 加载不带顶层的预训练VGG16
base_model = VGG16(weights="imagenet", include_top=False, input_shape=(224,224,3))
# 冻结基础层
for layer in base_model.layers:
    layer.trainable = False
# 添加自定义分类顶层
x = base_model.output
x = Flatten()(x)
x = Dense(256, activation='relu')(x)
x = Dropout(0.5)(x)
predictions = Dense(3, activation="softmax")(x)
model_final = Model(inputs=base_model.input, outputs=predictions)

总结

优先检查数据集类别与输出层维度匹配性,替换废弃API,修正冻结层范围,统一TF/Keras API使用,同时校验图像文件完整性,按上述步骤修改后可解决大部分图执行错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 15:57:04