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

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
      
  • 提升数据质量一致性:
    • 删除模糊、过曝/欠曝、背景杂乱的干扰图片
    • 检查是否有同一款罐头被误归为不同类,或不同款罐头特征过于相似
    • 验证训练集与测试集是否存在重复图片,杜绝数据泄露

三、模型训练策略优化

  • 分阶段训练:
    1. 先冻结Xception所有预训练层,只训练自定义顶部全连接层,让模型先适配分类任务(训练10-15个epoch)
    2. 解冻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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 09:21:35