基于TensorFlow Keras的胸片CNN分类模型测试准确率仅50%求助
高度不平衡胸片分类模型准确率停滞50%的解决方案
核心问题分析
- 数据集极端不平衡:训练集阴性样本(20000张)是阳性样本(467张)的40+倍,模型天然倾向于预测多数类,测试集正负样本各300张的情况下,准确率停留在50%等同于随机猜测。
- 类权重计算不合理:手动设置的权重未贴合训练集实际分布,极端权重会直接让模型倒向少数类。
- 评估指标无效:准确率在不平衡数据集里不具备参考价值,无法反映模型对阳性样本的识别能力。
- 模型特征提取能力不足:当前CNN结构过于简单,无法捕捉胸片中心脏增大的细微特征。
- 数据加载可能存在标签映射错误:测试集
class_names命名逻辑与训练集不一致,可能导致标签索引错乱。
具体修复步骤
1. 修正数据加载的标签一致性
确保训练集、验证集、测试集的类别命名和索引完全对应:
# 统一类别命名逻辑:索引0为阳性,1为阴性 classNames = ["pos", "neg"] # 训练集/验证集文件夹结构应为 ./data/trainup/pos/、./data/trainup/neg/ train_ds = tf.keras.preprocessing.image_dataset_from_directory( './data/trainup/', labels='inferred', label_mode='categorical', class_names=classNames, color_mode='grayscale', batch_size=32, image_size=(256, 256), shuffle=True, seed=123, validation_split=0.2, subset="training", interpolation='gaussian', ) val_ds = tf.keras.preprocessing.image_dataset_from_directory( './data/trainup/', labels='inferred', label_mode='categorical', class_names=classNames, color_mode='grayscale', batch_size=32, image_size=(256, 256), shuffle=True, seed=123, # 与训练集用相同seed,保证划分一致 validation_split=0.2, subset="validation", interpolation='gaussian', ) # 测试集文件夹结构应为 ./data/testup/pos/、./data/testup/neg/ test_ds = tf.keras.preprocessing.image_dataset_from_directory( './data/testup/', labels='inferred', label_mode='categorical', class_names=classNames, color_mode='grayscale', batch_size=32, image_size=(256, 256), shuffle=False, # 测试集不需要打乱,方便后续分析 interpolation='gaussian', )
2. 自动计算合理类权重
用sklearn工具根据训练集分布计算平衡权重,避免手动估算误差:
from sklearn.utils.class_weight import compute_class_weight # 提取训练集所有标签的索引 train_labels = [] for _, labels in train_ds: train_labels.extend(np.argmax(labels.numpy(), axis=1)) # 计算平衡类权重 class_weights = compute_class_weight('balanced', classes=np.unique(train_labels), y=train_labels) class_weight_dict = {i: class_weights[i] for i in range(len(class_weights))}
3. 更换有效评估指标
放弃准确率,改用对不平衡数据集更有意义的指标:
opt = keras.optimizers.Adam(learning_rate=1e-4) model.compile( optimizer=opt, loss=tf.keras.losses.CategoricalCrossentropy(from_logits=True), metrics=[ tf.keras.metrics.Precision(name='precision'), # 阳性预测准确率 tf.keras.metrics.Recall(name='recall'), # 阳性样本召回率 tf.keras.metrics.AUC(name='auc') # 整体分类能力 ] )
4. 增强阳性样本的数据扩充
针对阳性样本添加数据增强,减少分布差异:
# 数据增强层(仅训练时生效) data_augmentation = tf.keras.Sequential([ tf.keras.layers.experimental.preprocessing.RandomFlip("horizontal"), tf.keras.layers.experimental.preprocessing.RandomRotation(0.1), tf.keras.layers.experimental.preprocessing.RandomZoom(0.1), tf.keras.layers.experimental.preprocessing.RandomContrast(0.2) ])
5. 升级模型结构,增强特征提取能力
优化CNN结构,增加深度和特征维度,加入BatchNormalization稳定训练:
model = tf.keras.Sequential([ tf.keras.layers.experimental.preprocessing.Rescaling(1./255, input_shape=(256, 256, 1)), data_augmentation, # 加入数据增强 tf.keras.layers.Conv2D(32, 3, padding='same', activation='relu'), tf.keras.layers.BatchNormalization(), tf.keras.layers.MaxPooling2D(), tf.keras.layers.Conv2D(64, 3, padding='same', activation='relu'), tf.keras.layers.BatchNormalization(), tf.keras.layers.MaxPooling2D(), tf.keras.layers.Conv2D(128, 3, padding='same', activation='relu'), tf.keras.layers.BatchNormalization(), tf.keras.layers.MaxPooling2D(), tf.keras.layers.Dropout(0.3), # 防止过拟合 tf.keras.layers.Flatten(), tf.keras.layers.Dense(256, activation='relu'), tf.keras.layers.Dropout(0.5), tf.keras.layers.Dense(2) # 输出层,对应二分类 ])
6. 优化训练策略,防止过拟合
加入早停和学习率调度,提升训练稳定性:
# 早停:监控验证集召回率,5轮无提升则停止,恢复最优权重 early_stopping = tf.keras.callbacks.EarlyStopping( monitor='val_recall', patience=5, restore_best_weights=True, verbose=1 ) # 学习率调度:验证集损失3轮无下降则减半学习率 lr_scheduler = tf.keras.callbacks.ReduceLROnPlateau( monitor='val_loss', factor=0.5, patience=3, min_lr=1e-6, verbose=1 ) # 训练模型 history = model.fit( train_ds, validation_data=val_ds, epochs=30, class_weight=class_weight_dict, callbacks=[early_stopping, lr_scheduler] )
7. 验证预测结果,排查标签问题
通过混淆矩阵确认模型的实际分类效果:
import sklearn.metrics as le_me # 获取测试集真实标签和预测标签 test_true = [] test_pred = [] for imgs, labels in test_ds: test_true.extend(np.argmax(labels.numpy(), axis=1)) test_pred.extend(np.argmax(model.predict(imgs, verbose=0), axis=1)) # 打印混淆矩阵和关键指标 print("混淆矩阵:") print(le_me.confusion_matrix(test_true, test_pred)) print("\n分类报告:") print(le_me.classification_report(test_true, test_pred, target_names=["阳性", "阴性"]))
额外建议
- 优先验证标签映射:如果混淆矩阵显示全0或全300的极端情况,说明数据集文件夹命名或
class_names设置错误,需先修正。 - 尝试欠采样阴性样本:如果训练资源有限,可以随机采样部分阴性样本(如采样467*5=2335张),缩小正负样本差距后再训练。
- 迁移学习:使用预训练的图像模型(如VGG16、ResNet)作为特征提取器,在胸片数据集上微调,提升特征捕捉能力。
内容的提问来源于stack exchange,提问作者Jeff Stoomble
相关产品推荐
相关产品推荐

