基于CNN的半监督图像分类:数据拆分与后续流程咨询
半监督CNN训练流程与数据集拆分指南
一、数据集拆分的正确时机
- 必须先拆分带标签数据,再用
ImageDataGenerator处理。绝对不能先训练模型再拆分。 - 核心原因:测试集要完全独立,不能参与任何训练环节(包括初始有监督训练、伪标注生成等),否则评估结果会严重失真;验证集用来监控模型过拟合、调整超参数,也必须和训练集严格分离。
- 拆分建议:带标签数据总量少(共150张),按7:2:1比例拆分训练集、验证集、测试集,保证每个类别分布一致,比如猫取35/10/5张,狮子取70/20/10张。
二、半监督学习后续训练流程
你当前用全量带标签数据训练的方式是错误的,正确流程如下:
- 拆分带标签数据为训练集(train_with_label)、验证集(val_with_label)、测试集(test_with_label),三个文件夹结构与原
with_label一致(各包含cat_images和lion_images子文件夹)。 - 用带标签训练集训练基线CNN模型,用验证集监控训练过程,及时停止防止过拟合。
- 用训练好的基线模型对无标签数据做伪标注:预测每张图的类别,筛选置信度高于阈值(比如0.9)的样本,按预测类别归类到伪标注训练集(如
pseudo_cat、pseudo_lion)。 - 合并带标签训练集与高置信度伪标注数据集,重新训练模型(可微调基线模型或从头训练),继续用验证集监控。
- 重复步骤3-4:用新模型再次对剩余无标签数据做伪标注,筛选高置信度样本合并后训练,直到模型性能不再提升。
- 最后用独立测试集评估最终模型的真实性能。
三、修正后的代码示例
1. 数据生成器配置(已完成数据集拆分)
from tensorflow.keras.preprocessing.image import ImageDataGenerator from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense # 基础数据生成器(归一化) datagen = ImageDataGenerator(rescale=1./255) # 带标签训练集生成器 train_gen = datagen.flow_from_directory( './train_with_label', # 预先拆分好的训练集路径 target_size=(150, 150), batch_size=32, class_mode='binary' ) # 带标签验证集生成器 val_gen = datagen.flow_from_directory( './val_with_label', target_size=(150, 150), batch_size=32, class_mode='binary' ) # 带标签测试集生成器(仅用于最终评估) test_gen = datagen.flow_from_directory( './test_with_label', target_size=(150, 150), batch_size=32, class_mode='binary', shuffle=False # 不打乱,方便后续匹配真实标签 ) # 无标签数据生成器 unlabeled_gen = datagen.flow_from_directory( './without_label', target_size=(150, 150), batch_size=32, class_mode=None, shuffle=False # 不打乱,方便对应伪标注结果 )
2. 训练基线CNN模型
# 构建CNN模型(与你原结构一致) model = Sequential() model.add(Conv2D(32, (3, 3), activation='relu', input_shape=(150, 150, 3))) model.add(MaxPooling2D((2, 2))) model.add(Conv2D(64, (3, 3), activation='relu')) model.add(MaxPooling2D((2, 2))) model.add(Conv2D(128, (3, 3), activation='relu')) model.add(MaxPooling2D((2, 2))) model.add(Flatten()) model.add(Dense(512, activation='relu')) model.add(Dense(1, activation='sigmoid')) # 编译模型 model.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy']) # 训练基线模型,用验证集监控过拟合 history = model.fit( train_gen, epochs=15, # 根据验证集准确率调整,提前停止防止过拟合 validation_data=val_gen )
3. 生成高置信度伪标注数据
import numpy as np import os from shutil import copyfile # 对无标签数据进行预测 unlabeled_preds = model.predict(unlabeled_gen) # 筛选置信度>0.9或<0.1的高置信样本 confident_indices = np.where((unlabeled_preds > 0.9) | (unlabeled_preds < 0.1))[0] # 获取无标签数据的文件路径 unlabeled_files = unlabeled_gen.filenames # 创建伪标注训练文件夹 pseudo_train_dir = './pseudo_train' os.makedirs(os.path.join(pseudo_train_dir, 'cat_images'), exist_ok=True) os.makedirs(os.path.join(pseudo_train_dir, 'lion_images'), exist_ok=True) # 复制高置信度样本到对应伪标注类别文件夹 for idx in confident_indices: file_path = unlabeled_files[idx] full_source_path = os.path.join('./without_label', file_path) if unlabeled_preds[idx] < 0.1: # 预测为猫 copyfile(full_source_path, os.path.join(pseudo_train_dir, 'cat_images', os.path.basename(file_path))) else: # 预测为狮子 copyfile(full_source_path, os.path.join(pseudo_train_dir, 'lion_images', os.path.basename(file_path)))
4. 合并数据集并重新训练模型
# 简单方式:将伪标注数据复制到带标签训练集的对应文件夹(也可自定义生成器合并) # 复制完成后,重新创建合并后的训练集生成器 merged_train_gen = datagen.flow_from_directory( './train_with_label', # 已合并伪标注数据的训练集路径 target_size=(150, 150), batch_size=32, class_mode='binary' ) # 重新训练模型(可选择微调基线模型或从头训练) model.fit( merged_train_gen, epochs=10, validation_data=val_gen )
5. 模型最终评估
# 测试集基础评估 test_loss, test_acc = model.evaluate(test_gen) print(f"测试集损失:{test_loss:.4f},测试集准确率:{test_acc:.4f}") # 生成混淆矩阵与分类报告 from sklearn.metrics import confusion_matrix, classification_report test_preds = model.predict(test_gen) test_preds_class = np.where(test_preds > 0.5, 1, 0) test_true_class = test_gen.classes print("\n混淆矩阵:") print(confusion_matrix(test_true_class, test_preds_class)) print("\n分类报告:") print(classification_report(test_true_class, test_preds_class, target_names=['cat', 'lion']))
四、关键注意事项
- 伪标注置信度阈值不能过低,否则错误标签会严重干扰模型训练,建议从0.9开始尝试调整。
- 每次伪标注仅加入高置信度样本,不要一次性导入所有无标签数据。
- 带标签数据量少,训练时可加入数据增强(如
ImageDataGenerator的rotation_range、width_shift_range参数),缓解过拟合。
内容的提问来源于stack exchange,提问作者file_csv
相关产品推荐
相关产品推荐

