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

基于TensorFlow Keras的胸片CNN分类模型测试准确率仅50%求助

高度不平衡胸片分类模型准确率停滞50%的解决方案

核心问题分析

  1. 数据集极端不平衡:训练集阴性样本(20000张)是阳性样本(467张)的40+倍,模型天然倾向于预测多数类,测试集正负样本各300张的情况下,准确率停留在50%等同于随机猜测。
  2. 类权重计算不合理:手动设置的权重未贴合训练集实际分布,极端权重会直接让模型倒向少数类。
  3. 评估指标无效:准确率在不平衡数据集里不具备参考价值,无法反映模型对阳性样本的识别能力。
  4. 模型特征提取能力不足:当前CNN结构过于简单,无法捕捉胸片中心脏增大的细微特征。
  5. 数据加载可能存在标签映射错误:测试集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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 18:20:20