运行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
相关产品推荐
相关产品推荐

