1. 项目概述
今天我想分享一个基于PyTorch实现的图像分类项目,特别关注半监督学习在实际应用中的实现。这个项目使用food-11数据集,通过结合少量标注数据和大量未标注数据,构建了一个有效的食物分类系统。
作为一名长期从事计算机视觉开发的工程师,我发现半监督学习在实际业务场景中特别有价值——它能够显著减少对标注数据的依赖,同时保持不错的模型性能。下面我将详细解析这个项目的技术实现,包括数据准备、模型架构、训练策略等核心环节。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备与数据加载
2.1 基础环境配置
首先需要安装必要的Python包:
bash复制pip install opencv-python
pip install timm
这里特别说明两个关键库的选择原因:
- OpenCV(cv2):用于图像读取和基础处理,相比PIL更高效且功能丰富
- timm:PyTorch生态中优秀的预训练模型库,提供大量SOTA模型实现
提示:建议使用虚拟环境管理依赖,避免包版本冲突。我常用conda创建独立环境:
bash复制conda create -n semi_sup python=3.8 conda activate semi_sup
2.2 数据加载实现
项目中的数据加载器设计很有特色,支持三种模式:
python复制class food_Dataset(Dataset):
def __init__(self, path, mode="train"):
self.mode = mode # train/val/semi
...
关键设计点:
- 统一接口:通过mode参数区分不同数据场景
- 内存优化:使用np.uint8存储图像,减少内存占用
- 灵活扩展:支持动态添加未标注数据
数据增强策略也值得关注:
python复制train_transform = transforms.Compose([
transforms.ToPILImage(),
transforms.RandomResizedCrop(224),
transforms.RandomRotation(50),
transforms.ToTensor()
])
为什么选择这些增强?
- RandomResizedCrop:模拟不同拍摄视角
- RandomRotation:增强旋转不变性
- 注意保持验证集不做随机增强,确保评估稳定
3. 模型架构与迁移学习
3.1 自定义CNN模型
项目提供了两种模型选择:
- 自定义CNN(myModel)
- 预训练VGG(通过timm加载)
这里重点分析自定义CNN的设计:
python复制class myModel(nn.Module):
def __init__(self, num_class):
super().__init__()
self.conv1 = nn.Conv2d(3, 64, 3, 1, 1)
self.bn1 = nn.BatchNorm2d(64)
self.pool1 = nn.MaxPool2d(2)
...
模型特点:
- 渐进式下采样:224x224 → 7x7
- 每层包含Conv+BN+ReLU+Pooling
- 最终全连接层实现分类
经验分享:对于食物分类这种中等复杂度任务,建议通道数从64开始逐步加倍,避免过早丢失细节信息。
3.2 迁移学习实践
更推荐使用预训练模型:
python复制model, _ = initialize_model("vgg", 11, use_pretrained=True)
迁移学习的两种策略:
- 特征提取:冻结骨干网络,仅训练分类头
- 微调:全部参数参与训练,但使用较小学习率
本项目采用微调策略,因为:
- 食物图像与ImageNet有相似性但差异明显
- 数据量相对充足(相比纯特征提取)
4. 半监督学习实现
4.1 核心算法流程
半监督学习的关键在于:
- 用已标注数据训练初始模型
- 对未标注数据预测伪标签
- 高置信度预测加入训练集
实现代码:
python复制def get_semi_loader(no_label_loader, model, device, thres=0.99):
semiset = semiDataset(no_label_loader, model, device, thres)
...
4.2 置信度阈值选择
阈值设置是成败关键:
- 过高(如0.99):数据利用率低
- 过低(如0.9):噪声标签影响性能
建议策略:
- 初始阶段设高阈值(0.95-0.99)
- 随训练过程动态降低
- 不同类别可设置不同阈值
5. 训练优化技巧
5.1 优化器配置
项目使用AdamW优化器:
python复制optimizer = torch.optim.AdamW(
model.parameters(),
lr=0.001,
weight_decay=1e-4
)
为什么选择AdamW?
- 自适应学习率:适合不同参数
- 权重衰减解耦:更稳定的正则化
- 相比SGD更少需要学习率调度
5.2 训练监控
完善的训练日志很重要:
python复制print(f'[{epoch}/{epochs}] {time.time()-start_time:.2f} sec(s) '
f'TrainLoss: {train_loss/len(train_loader):.6f} | '
f'ValLoss: {val_loss/len(val_loader):.6f} '
f'TrainAcc: {train_acc/len(train_set):.6f} | '
f'ValAcc: {val_acc/len(val_set):.6f}')
建议增加:
- 学习率记录
- 显存监控
- 半监督数据量统计
6. 常见问题与解决方案
6.1 内存不足
症状:训练时出现CUDA out of memory
解决方法:
- 减小batch size(如16→8)
- 使用梯度累积:
python复制for i, (x, y) in enumerate(train_loader): loss.backward() if (i+1) % 2 == 0: # 每2步更新一次 optimizer.step() optimizer.zero_grad()
6.2 过拟合
症状:训练acc高但验证acc低
对策:
- 增加数据增强(如ColorJitter)
- 早停机制(监控验证loss)
- 更强的正则化(增大weight_decay)
6.3 半监督效果差
可能原因:
- 初始模型质量差
- 阈值设置不合理
- 数据分布不一致
改进方法:
- 先用更多标注数据训练强初始模型
- 实施课程学习(逐步降低阈值)
- 检查未标注数据质量
7. 项目扩展方向
基于这个基础框架,还可以尝试:
-
更先进的半监督算法
- MixMatch
- FixMatch
- 一致性正则化
-
模型架构改进
- 使用EfficientNet等现代架构
- 添加注意力机制
- 多尺度特征融合
-
部署优化
- 模型量化
- ONNX导出
- TensorRT加速
这个项目最让我惊喜的是通过简单的伪标签方法,就能显著提升模型性能。在实际业务中,标注成本往往很高,而半监督学习提供了很好的性价比方案。建议读者可以从这个基础实现出发,逐步尝试更复杂的算法变体。
