基于Keras ImageDataGenerator训练带ImageNet预权重的ResNet50图像分类模型
针对ResNet50图像分类训练的完整流程建议
嘿,我来帮你把这个ResNet50训练的流程补全并优化一下~首先先把你没写完的验证集生成代码补上,和训练集的逻辑基本一致,只需要替换路径和几个参数:
test_datagen = ImageDataGenerator() validation_generator = test_datagen.flow_from_directory( './val_qcut_2_classes', # 记得换成你的验证集实际目录路径 batch_size=batch_size, shuffle=False, # 验证集一般不用打乱,这样后续评估的时候能对应上真实标签 target_size=input_size[1:], class_mode='categorical' )
接下来给你梳理从加载预训练模型到完成训练的完整步骤,都是实际项目里常用的操作,能帮你少踩坑:
1. 加载预训练ResNet50并搭建自定义分类头
因为我们是做自己的二分类任务,所以要去掉ResNet50自带的ImageNet顶层分类层,换成适配我们任务的结构:
from tensorflow.keras.applications import ResNet50 from tensorflow.keras.layers import Dense, GlobalAveragePooling2D from tensorflow.keras.models import Model # 加载预训练权重,排除顶层的分类层 base_model = ResNet50(weights='imagenet', include_top=False, input_shape=input_size) # 添加我们自己的分类头 x = base_model.output x = GlobalAveragePooling2D()(x) # 用全局平均池化替代Flatten,减少参数数量 x = Dense(1024, activation='relu')(x) # 可选的中间隐藏层,你可以根据数据集大小调整 predictions = Dense(2, activation='softmax')(x) # 二分类任务,输出维度设为2 # 组装成完整的训练模型 model = Model(inputs=base_model.input, outputs=predictions)
2. 分层训练(冻结+解冻,提升效率和精度)
一开始建议先冻结所有预训练层,只训练我们新增的分类头——这样既能快速让模型适配你的数据,又不会破坏ResNet50已经学到的通用特征:
# 冻结base_model的所有层 for layer in base_model.layers: layer.trainable = False
等分类头训练个5-10个epoch,精度稳定后,再解冻一部分顶层的预训练层进行微调,能进一步榨取模型的性能:
# 比如解冻base_model的最后10层,你也可以根据情况调整层数 for layer in base_model.layers[-10:]: layer.trainable = True
3. 编译模型
针对二分类任务,选对损失函数和优化器很重要:
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
如果是微调阶段,建议把学习率调小一点,比如用Adam(learning_rate=1e-5),避免预训练层的权重被冲得太厉害。
4. 启动训练
用fit配合我们的生成器开始训练就行:
history = model.fit( train_generator, epochs=15, # 这个数值你可以根据训练过程中的精度变化调整 validation_data=validation_generator, steps_per_epoch=train_generator.samples // batch_size, validation_steps=validation_generator.samples // batch_size )
5. 额外的实用优化建议
- 加数据增强:你现在用的是无增强的
ImageDataGenerator,可以加一些简单的增强来提升模型泛化能力,比如:
注意!验证集别加增强,保持和实际测试场景一致。train_datagen = ImageDataGenerator( rotation_range=20, width_shift_range=0.2, height_shift_range=0.2, horizontal_flip=True, zoom_range=0.15 ) - 处理类不平衡:如果你的两类样本数量差很多,可以在
flow_from_directory里加上class_weight参数,给样本少的类别更高的权重。 - 监控训练过程:可以用
matplotlib画出训练的精度和损失曲线,方便判断是否过拟合或者欠拟合。
内容的提问来源于stack exchange,提问作者MenorcanOrange
相关产品推荐
相关产品推荐

