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

基于CNN的半监督图像分类:数据拆分与后续流程咨询

半监督CNN训练流程与数据集拆分指南

一、数据集拆分的正确时机

  • 必须先拆分带标签数据,再用ImageDataGenerator处理。绝对不能先训练模型再拆分。
  • 核心原因:测试集要完全独立,不能参与任何训练环节(包括初始有监督训练、伪标注生成等),否则评估结果会严重失真;验证集用来监控模型过拟合、调整超参数,也必须和训练集严格分离。
  • 拆分建议:带标签数据总量少(共150张),按7:2:1比例拆分训练集、验证集、测试集,保证每个类别分布一致,比如猫取35/10/5张,狮子取70/20/10张。

二、半监督学习后续训练流程

你当前用全量带标签数据训练的方式是错误的,正确流程如下:

  1. 拆分带标签数据为训练集(train_with_label)、验证集(val_with_label)、测试集(test_with_label),三个文件夹结构与原with_label一致(各包含cat_images和lion_images子文件夹)。
  2. 用带标签训练集训练基线CNN模型,用验证集监控训练过程,及时停止防止过拟合。
  3. 用训练好的基线模型对无标签数据做伪标注:预测每张图的类别,筛选置信度高于阈值(比如0.9)的样本,按预测类别归类到伪标注训练集(如pseudo_cat、pseudo_lion)。
  4. 合并带标签训练集与高置信度伪标注数据集,重新训练模型(可微调基线模型或从头训练),继续用验证集监控。
  5. 重复步骤3-4:用新模型再次对剩余无标签数据做伪标注,筛选高置信度样本合并后训练,直到模型性能不再提升。
  6. 最后用独立测试集评估最终模型的真实性能。

三、修正后的代码示例

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 13:45:15