Xception模型修复求助:猫罐头分类模型预测结果异常
猫罐头分类模型问题排查与优化方案
一、代码中的明显错误修复
- ModelCheckpoint模式错误:你设置了
monitor='val_loss'但mode='max',这会导致保存val_loss最大的最差模型,应改为mode='min':checkpoint = ModelCheckpoint('Xception_checkpoint.h5', verbose=1, monitor='val_loss', save_best_only=True, mode='min') - 未启用冻结层配置:定义了
FREEZE_LAYERS=2但没实际冻结模型层,导致整个Xception网络都在训练,易过拟合且效率低,添加冻结代码:# 冻结除最后2层外的所有预训练层 for layer in model.layers[:-FREEZE_LAYERS]: layer.trainable = False # 确保最后几层可训练 for layer in model.layers[-FREEZE_LAYERS:]: layer.trainable = True - 冗余Flatten层:
GlobalAveragePooling2D已输出一维张量,后续Flatten()完全多余,直接删除:x = model.output x = GlobalAveragePooling2D()(x) # x = Flatten()(x) # 删除该行 x = Dropout(0.5)(x) predictions = Dense(26, activation='softmax')(x) - 替换过时的fit_generator:Keras中
fit_generator已弃用,改用fit方法:history = model.fit(train_generator, epochs=epoch, verbose=1, steps_per_epoch=train_generator.samples//batch_size, validation_data=valid_generator, validation_steps=valid_generator.samples//batch_size, callbacks=[checkpoint, estop, reduce_lr]) - Xception预处理修正:预训练Xception要求输入像素缩放到
[-1,1],而非[0,1],替换数据增强的rescale为官方预处理函数:train_datagen = ImageDataGenerator( preprocessing_function=tf.keras.applications.xception.preprocess_input, rotation_range=30, width_shift_range=0.2, height_shift_range=0.2, shear_range=0.2, zoom_range=0.2, channel_shift_range=10, horizontal_flip=True, fill_mode='nearest' ) val_datagen = ImageDataGenerator( preprocessing_function=tf.keras.applications.xception.preprocess_input ) - Adam学习率参数与数值调整:新版本Keras中
lr改为learning_rate,且预训练模型初始学习率不宜过高,建议设为1e-4:model.compile(optimizer=Adam(learning_rate=1e-4), loss='categorical_crossentropy', metrics=['accuracy'])
二、数据集层面的排查与优化
- 修正数据集划分逻辑:当前用
test_path作为验证集,测试集应留到最终评估用,建议从训练集中拆分15%-20%作为独立验证集,避免测试集数据泄露。 - 平衡类别样本量:统计26个类别的样本数,若存在样本量差异极大的类别(比如某类仅10张,另一类500张),模型会偏向样本多的类,可通过以下方式解决:
- 对样本少的类别针对性增强(比如多做旋转、亮度调整)
- 使用
class_weight给少数类更高权重:from sklearn.utils.class_weight import compute_class_weight import numpy as np class_labels = train_generator.class_indices.values() class_weights = compute_class_weight(class_weight='balanced', classes=np.unique(class_labels), y=train_generator.classes) class_weights = dict(zip(np.unique(class_labels), class_weights)) # 在fit中加入class_weight=class_weights
- 提升数据质量一致性:
- 删除模糊、过曝/欠曝、背景杂乱的干扰图片
- 检查是否有同一款罐头被误归为不同类,或不同款罐头特征过于相似
- 验证训练集与测试集是否存在重复图片,杜绝数据泄露
三、模型训练策略优化
- 分阶段训练:
- 先冻结Xception所有预训练层,只训练自定义顶部全连接层,让模型先适配分类任务(训练10-15个epoch)
- 解冻Xception最后10-20层,用更小的学习率(比如1e-5)微调,让预训练特征适配你的数据集
- 调整正则化强度:在Dense层前加入BatchNormalization稳定训练,同时可将Dropout比例调至0.3避免过度正则化:
x = model.output x = GlobalAveragePooling2D()(x) x = BatchNormalization()(x) x = Dropout(0.3)(x) predictions = Dense(26, activation='softmax')(x) - 增加训练epoch上限:当前epoch=30可能不足,可设为50个epoch,配合EarlyStopping自动停止不必要的训练
四、诊断与验证工具
- 可视化训练曲线:通过loss和accuracy曲线判断过拟合/欠拟合:
plt.plot(history.history['accuracy'], label='训练准确率') plt.plot(history.history['val_accuracy'], label='验证准确率') plt.xlabel('轮次') plt.ylabel('准确率') plt.legend() plt.show() plt.plot(history.history['loss'], label='训练损失') plt.plot(history.history['val_loss'], label='验证损失') plt.xlabel('轮次') plt.ylabel('损失') plt.legend() plt.show() - 混淆矩阵分析:找出易混淆的类别,针对性优化数据或特征:
from sklearn.metrics import confusion_matrix import seaborn as sns # 获取验证集预测结果 y_pred = model.predict(valid_generator) y_pred_classes = np.argmax(y_pred, axis=1) y_true = valid_generator.classes # 绘制混淆矩阵 cm = confusion_matrix(y_true, y_pred_classes) plt.figure(figsize=(12,12)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=valid_generator.class_indices.keys(), yticklabels=valid_generator.class_indices.keys()) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.show()
内容的提问来源于stack exchange,提问作者Chanel Chen
相关产品推荐
相关产品推荐

