Keras迁移学习模型测试精度停滞问题求助
道路卫星图像四分类模型测试精度停滞40%的排查与解决
问题背景
针对道路卫星图像做四分类(Good/Fair/Poor/Bad对应标签0-3),采用Keras迁移学习方案,训练精度可达99%但测试精度始终停滞在40%左右,已尝试增加epochs、更换模型架构、调整学习率/优化器、修改图像尺寸等手段,问题仍未解决。
核心问题定位
训练精度极高但测试精度极低,本质是严重过拟合,结合数据集、模型架构、训练流程等维度,可从以下方向逐一排查优化:
1. 数据集层面排查与优化
- 类别不平衡检查:统计训练/测试集中四类样本的数量,若某类占比过高(如超过60%),模型会倾向于预测该类导致整体精度偏低。
- 解决:使用
class_weight参数给少数类赋予更高权重,或对样本少的类做过采样、样本多的类做欠采样。
- 解决:使用
- 数据集分布一致性:确认训练/验证/测试集的图像场景是否一致(如是否存在训练集全是城市道路、测试集全是乡村道路的情况),分布差异会导致模型泛化能力差。
- 解决:采用随机划分方式拆分数据集,避免手动拆分带来的分布偏差。
- 标签正确性验证:抽查测试集标签,确认是否存在标注错误(如Bad类被误标为Good),错误标签会直接拉低评估精度。
2. 模型架构优化(重点)
- 替换Flatten为全局池化层:当前用
Flatten将Xception输出的(5,5,2048)转为51200维向量,后续Dense层参数高达512万,极易引发过拟合。改用GlobalAveragePooling2D或GlobalMaxPooling2D,可将特征维度压缩至2048维,大幅减少参数数量:model = keras.models.Sequential() model.add(keras.applications.Xception(include_top=False, weights='imagenet', input_shape=(150, 150, 3))) model.add(keras.layers.GlobalAveragePooling2D()) # 替换Flatten model.add(keras.layers.Dense(100)) model.add(keras.layers.Activation(keras.activations.relu)) model.add(keras.layers.BatchNormalization()) model.add(keras.layers.Dropout(0.5)) model.add(keras.layers.Dense(4)) model.add(keras.layers.Activation(keras.activations.softmax)) - 分层解冻预训练层:不要一次性解冻全部Xception层,优先解冻顶层10-20层(底层是通用图像特征,顶层是ImageNet专属特征),既保留预训练优势,又适配道路图像任务,同时减少训练参数:
# 解冻Xception最后20层 for layer in model.layers[0].layers[-20:]: layer.trainable = True # 重新编译,使用极小学习率微调 model.compile(loss='categorical_crossentropy', optimizer=tf.optimizers.Adam(learning_rate=1e-5), metrics=['accuracy'])
3. 数据增强提升泛化能力
训练时添加针对卫星图像的数据增强,打破模型对训练图像细节的记忆:
data_augmentation = keras.Sequential([ keras.layers.RandomFlip("horizontal_and_vertical"), keras.layers.RandomRotation(0.2), keras.layers.RandomZoom(0.2), keras.layers.RandomContrast(0.1), ]) # 对训练集应用增强 train_ds = train_ds.map(lambda x, y: (data_augmentation(x, training=True), y))
4. 训练流程正则化优化
- 添加早停机制:监控验证精度,当精度连续多轮不提升时停止训练,自动恢复最优权重:
early_stopping = keras.callbacks.EarlyStopping(monitor='val_accuracy', patience=5, restore_best_weights=True) model.fit(train_ds, validation_data=val_ds, epochs=50, callbacks=[early_stopping]) - 添加L2正则化:在Dense层加入权重衰减,限制权重大小避免过拟合:
model.add(keras.layers.Dense(100, kernel_regularizer=keras.regularizers.l2(0.001)))
5. 预处理与评估流程验证
- 预训练模型专属预处理:Xception等预训练模型要求特定的图像归一化,需在加载数据集时应用对应预处理函数:
def preprocess_image(image, label): image = keras.applications.xception.preprocess_input(image) return image, label train_ds = train_ds.map(preprocess_image) test_ds = test_ds.map(preprocess_image) - 标签格式一致性:使用
image_dataset_from_directory时,若设置label_mode='categorical',会自动生成one-hot标签,无需手动调用tf.onehot,避免标签格式不匹配。
内容的提问来源于stack exchange,提问作者laptopbox2023
相关产品推荐
相关产品推荐

