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

如何在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数组的转换,能避免潜在的类型匹配问题。

额外注意事项

  1. 微调最后4层时,一定要确保前面的层被正确冻结,否则训练会很慢,还容易破坏预训练模型学到的特征提取能力。
  2. 训练时建议使用较小的学习率(比如把Adam的学习率设为1e-5,而非默认的1e-3),这样微调时不会打乱预训练的权重。
  3. flow_from_directory的target_size必须和VGG16的输入尺寸一致(默认是(224,224)),否则模型会因输入维度不匹配报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 06:27:37