1. 半监督图像分类框架概述
在计算机视觉领域,图像分类是最基础也最重要的任务之一。传统的监督学习需要大量标注数据,而数据标注往往成本高昂。半监督学习(Semi-Supervised Learning)正是为了解决这个问题而提出的,它能够同时利用少量标注数据和大量未标注数据来训练模型。
这个半监督图像分类框架基于PyTorch实现,核心思路是"伪标签法"(Pseudo Labeling)。具体来说,就是先用有标签数据训练一个基础模型,然后用这个模型对无标签数据进行预测,筛选出高置信度的预测结果作为"伪标签",最后将这些伪标签数据重新加入训练集进行训练。
注意:伪标签法的关键在于置信度阈值的选择。阈值过高会导致可用数据太少,阈值过低则可能引入太多错误标签。本框架默认使用0.99的高阈值,确保伪标签质量。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备与数据加载
2.1 环境配置
首先需要安装必要的Python库。建议使用conda创建虚拟环境:
bash复制conda create -n semi_supervised python=3.8
conda activate semi_supervised
pip install torch torchvision numpy matplotlib pillow tqdm
2.2 数据加载器实现
框架的核心数据结构是自定义的food_Dataset类,它继承自PyTorch的Dataset类,支持三种模式:
- 训练模式(train):读取图片和标签,应用数据增强
- 验证模式(val):读取图片和标签,不应用数据增强
- 半监督模式(semi):只读取图片,没有标签
python复制class food_Dataset(Dataset):
def __init__(self, path, mode="train"):
self.mode = mode
if mode == "semi":
self.X = self.read_file(path) # 只读图片
else:
self.X, self.Y = self.read_file(path) # 读图片和标签
self.Y = torch.LongTensor(self.Y)
# 根据模式选择不同的数据预处理
self.transform = train_transform if mode == "train" else val_transform
数据增强对于防止过拟合至关重要。本框架为训练集和验证集设计了不同的预处理流程:
python复制train_transform = transforms.Compose([
transforms.ToPILImage(),
transforms.RandomResizedCrop(224), # 随机裁剪和缩放
transforms.RandomRotation(50), # 随机旋转
transforms.ToTensor() # 转为张量
])
val_transform = transforms.Compose([ # 验
