如何在Keras 2.1.5的ImageDataGenerator中使用VGG16的preprocess_input
解决VGG16迁移学习中ImageDataGenerator使用preprocess_input的报错问题
我刚看到你遇到的这个'JpegImageFile' object is not subscriptable错误,其实原因很明确:preprocess_input是为numpy数组设计的,但ImageDataGenerator默认传递给预处理函数的是PIL图像对象,而PIL对象不支持像数组那样的下标访问(比如img[i,j]),所以就抛出了这个错误。
下面给你两种可行的解决办法,附完整代码示例:
方法一:包装preprocess_input函数,先转numpy数组
我们可以写一个简单的包装函数,先把PIL图像转换成numpy数组,再传入preprocess_input处理:
import numpy as np from keras.applications.vgg16 import VGG16, preprocess_input from keras.preprocessing.image import ImageDataGenerator from keras.models import Model from keras.layers import Dense, Flatten # 定义包装后的预处理函数 def vgg_preprocess(img): # 将PIL图像转为numpy数组 img_array = np.array(img) # 应用VGG16的预处理逻辑 return preprocess_input(img_array) # 加载预训练VGG16模型,去掉顶层全连接层 base_model = VGG16(weights='imagenet', include_top=False, input_shape=(224,224,3)) # 冻结前面的层,只微调最后4层 for layer in base_model.layers[:-4]: layer.trainable = False # 构建自定义分类头 x = base_model.output x = Flatten()(x) x = Dense(256, activation='relu')(x) predictions = Dense(13, activation='softmax')(x) model = Model(inputs=base_model.input, outputs=predictions) # 编译模型,用小学习率避免破坏预训练权重 model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy']) # 初始化ImageDataGenerator,使用包装后的预处理函数 datagen = ImageDataGenerator( preprocessing_function=vgg_preprocess, rotation_range=20, width_shift_range=0.2, height_shift_range=0.2, horizontal_flip=True ) # 从目录加载训练数据 train_generator = datagen.flow_from_directory( 'train_dir', target_size=(224,224), batch_size=32, class_mode='categorical' ) # 开始训练 model.fit( train_generator, epochs=10, steps_per_epoch=train_generator.samples // train_generator.batch_size )
方法二:先rescale再直接用preprocess_input(另一种思路)
如果你不想写包装函数,也可以先让ImageDataGenerator把像素值维持在0-255范围,再直接传入preprocess_input,因为它的逻辑是基于0-255像素值设计的:
# 这里不做rescale(默认就是0-255),直接传入preprocess_input datagen = ImageDataGenerator( preprocessing_function=preprocess_input, rotation_range=20, width_shift_range=0.2, height_shift_range=0.2, horizontal_flip=True )
不过更推荐第一种方法,因为它明确处理了PIL到numpy数组的转换,能避免潜在的类型匹配问题。
额外注意事项
- 微调最后4层时,一定要确保前面的层被正确冻结,否则训练会很慢,还容易破坏预训练模型学到的特征提取能力。
- 训练时建议使用较小的学习率(比如把Adam的学习率设为
1e-5,而非默认的1e-3),这样微调时不会打乱预训练的权重。 flow_from_directory的target_size必须和VGG16的输入尺寸一致(默认是(224,224)),否则模型会因输入维度不匹配报错。
内容的提问来源于stack exchange,提问作者prashantgpt91
相关产品推荐
相关产品推荐

