1. 项目概述与背景
在计算机视觉领域,图片分类是最基础也最重要的任务之一。这个项目实现了一个基于深度学习的图片分类框架,特别之处在于它同时利用了有标签数据和无标签数据进行训练。这种半监督学习的方法在实际应用中非常有价值,因为获取大量标注数据往往成本高昂,而无标签数据则相对容易获得。
项目中使用的数据集包含11个类别的图片,主要涉及食品分类场景。这种细粒度分类任务对模型的特征提取能力提出了较高要求。框架的核心思路是:
- 先用有标签数据训练基础模型
- 然后用训练好的模型对无标签数据进行预测
- 对高置信度的预测结果打上"伪标签"
- 最后将这些伪标签数据重新加入训练集进行迭代优化
这种半监督学习方法能显著提升模型性能,特别是在标注数据有限的情况下。从技术实现来看,项目采用了PyTorch框架,并基于ResNet18进行迁移学习,同时加入了数据增强、模型验证等标准流程。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备与数据加载
2.1 开发环境配置
项目使用PyTorch作为深度学习框架,这是目前最流行的选择之一。首先需要设置随机种子以保证结果可复现:
python复制def seed_everything(seed):
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
random.seed(seed)
np.random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
seed_everything(0) # 固定随机种子
注意:设置随机种子是深度学习项目中容易被忽视但非常重要的一步。它确保了每次运行代码时,随机初始化、数据打乱等操作产生相同的结果,这对调试和结果复现至关重要。
2.2 数据加载与预处理
项目处理两种类型的数据:
- 有标签数据:包含图片(X)和对应的类别标签(Y)
- 无标签数据:只有图片(X)
数据加载通过自定义的FoodDataset类实现:
python复制class FoodDataset(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) # 标签转为LongTensor
# 设置数据增强
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([
transforms.ToPILImage(),
transforms.ToTensor()
])
实操技巧:验证集不应该使用过于激进的数据增强,否则会影响对模型真实性能的评估。通常只需保持与训练集相同的基础预处理即可。
3. 模型架构与迁移学习
3.1 模型选择
项目采用了ResNet18作为基础模型,这是一种经典的卷积神经网络架构:
python复制from torchvision.models import resnet18
model = resnet18(pretrained=True) # 使用预训练权重
in_features = model.fc.in_features
model.fc = nn.Linear(in_features, 11) # 修改最后的全连接层
ResNet18的优势在于:
- 深度适中,训练和推理速度较快
- 残差连接有效缓解了梯度消失问题
- ImageNet预训练权重提供了良好的特征提取能力
3.2 迁移学习策略
项目中提供了两种微调策略:
- 线性探测(linear probing):只训练最后的全连接层
- 完整微调:训练所有层参数
python复制# 初始化模型的工具函数
def initialize_model(model_name, num_classes, use_pretrained=True, linear_prob=False):
model = resnet18(pretrained=use_pretrained)
in_features = model.fc.in_features
model.fc = nn.Linear(in_features, num_classes)
if linear_prob: # 线性探测时冻结其他层
for param in model.parameters():
param.requires_grad = False
for param in model.fc.parameters():
param.requires_grad = True
return model, in_features
经验分享:当标注数据较少时,线性探测通常是更好的选择;而有足够标注数据时,完整微调往往能获得更好的性能。实际项目中可以两种方法都尝试,选择验证集表现更好的方案。
4. 半监督学习实现
4.1 伪标签生成
半监督学习的核心是为无标签数据生成可靠的伪标签:
python复制class semiDataset(Dataset):
def __init__(self, no_label_loader, model, device, thres):
x, y = self.get_label(no_label_loader, model, device, thres)
if x: # 有符合条件的样本
self.flag = True
self.X = np.array(x)
self.Y = torch.LongTensor(y)
self.transform = train_transform
else:
self.flag = False
def get_label(self, no_label_loader, model, device, thres):
model.eval()
soft = torch.nn.Softmax(dim=1)
x, y = [], []
with torch.no_grad():
for batch_x, _ in no_label_loader:
batch_x = batch_x.to(device)
pred = model(batch_x)
pred_soft = soft(pred)
pred_max, pred_label = pred_soft.max(1)
# 筛选高置信度样本
mask = pred_max > thres
if mask.any():
x.extend([no_label_loader.dataset[i][1]
for i, m in enumerate(mask) if m])
y.extend(pred_label[mask].cpu().numpy().tolist())
return x, y
4.2 置信度阈值选择
置信度阈值(thres)的选择对半监督学习效果影响很大:
python复制thres = 0.99 # 设置较高的置信度阈值
重要原则:宁可少用一些无标签数据,也不要引入大量低质量的伪标签。过低的阈值会导致噪声积累,反而降低模型性能。实际项目中可以通过实验选择最佳阈值。
5. 训练流程与优化
5.1 训练配置
项目使用了AdamW优化器和交叉熵损失函数:
python复制lr = 0.001
loss = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
device = "cuda" if torch.cuda.is_available() else "cpu"
AdamW相比经典Adam的主要改进是:
- 更正确的权重衰减实现方式
- 通常能获得更好的泛化性能
- 对超参数相对鲁棒
5.2 训练循环
训练过程包含以下几个关键步骤:
- 常规有监督训练
- 伪标签生成
- 半监督训练
- 模型验证与保存
python复制def train_val(model, train_loader, val_loader, no_label_loader,
optimizer, device, epochs, thres, save_path):
model = model.to(device)
best_acc = 0.0
semi_loader = None
for epoch in range(epochs):
# 常规训练
model.train()
train_loss, train_acc = 0.0, 0.0
for x, y in train_loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
pred = model(x)
loss = criterion(pred, y)
loss.backward()
optimizer.step()
train_loss += loss.item()
train_acc += (pred.argmax(1) == y).sum().item()
# 半监督训练
if semi_loader:
semi_loss, semi_acc = 0.0, 0.0
for x, y in semi_loader:
# 类似常规训练流程...
# 验证
model.eval()
val_loss, val_acc = 0.0, 0.0
with torch.no_grad():
for x, y in val_loader:
# 计算验证损失和准确率...
# 保存最佳模型
if val_acc > best_acc:
best_acc = val_acc
torch.save(model.state_dict(), save_path)
# 生成新伪标签
if val_acc > 0.6: # 当模型达到一定性能后才使用无标签数据
semi_loader = get_semi_loader(no_label_loader, model, device, thres)
训练技巧:不要一开始就使用无标签数据,等模型在有标签数据上表现较好(如准确率>60%)后再引入半监督学习,这样生成的伪标签质量更高。
6. 结果分析与可视化
训练过程中记录了损失和准确率的变化,并进行了可视化:
python复制plt.plot(train_losses, label='Train')
plt.plot(val_losses, label='Validation')
plt.title('Training and Validation Loss')
plt.legend()
plt.show()
plt.plot(train_accs, label='Train')
plt.plot(val_accs, label='Validation')
plt.title('Training and Validation Accuracy')
plt.legend()
plt.show()
典型的训练曲线应该呈现以下特征:
- 训练损失持续下降
- 验证损失先下降后趋于平稳或略有上升
- 训练和验证准确率逐步提升并趋于稳定
如果出现验证损失上升而准确率下降的情况,可能表明模型开始过拟合,需要调整正则化策略或提前停止训练。
7. 常见问题与解决方案
7.1 内存不足问题
当处理大型图像数据集时,可能会遇到GPU内存不足的情况。解决方法包括:
- 减小batch size
- 使用梯度累积技巧
- 尝试混合精度训练
- 优化数据加载流程
7.2 过拟合问题
如果模型在训练集上表现很好但在验证集上较差,可以尝试:
- 增加数据增强的强度
- 添加更多的正则化(如Dropout、权重衰减)
- 使用早停策略
- 减少模型复杂度
7.3 半监督学习效果不佳
当引入无标签数据后性能下降时,可以检查:
- 置信度阈值是否设置合理
- 伪标签数据的质量
- 有标签数据和无标签数据的分布是否一致
- 模型在有标签数据上的表现是否足够好
8. 项目扩展与改进方向
这个基础框架可以进一步扩展和优化:
- 模型架构:尝试更先进的网络如EfficientNet、Vision Transformer等
- 数据增强:加入更强大的增强策略如MixUp、CutMix等
- 半监督算法:实现更先进的半监督学习方法如FixMatch、FlexMatch等
- 类别不平衡:对于不平衡数据集,可以采用加权损失或过采样技术
- 部署优化:使用ONNX或TensorRT进行模型优化和加速
在实际应用中,我发现合理设置伪标签的置信度阈值对最终效果影响最大。通常需要多次实验才能找到最佳值。另外,数据质量比数量更重要,确保无标签数据与目标任务相关是成功应用半监督学习的前提。
