1. 深度学习图像分类项目概述
这个项目实现了一个基于深度学习的图像分类系统,主要针对食品图片进行分类任务。项目采用了半监督学习策略,结合了有标签数据和无标签数据进行模型训练,有效提升了在小样本数据集上的分类性能。整个系统基于PyTorch框架构建,包含了数据预处理、模型构建、训练优化和评估验证等完整流程。
核心特点:
- 使用交叉熵损失函数进行多分类任务
- 采用AdamW优化器进行模型参数更新
- 实现了半监督学习流程,充分利用无标签数据
- 支持多种数据增强技术
- 提供了完整的训练和验证流程
- 包含模型性能可视化功能
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心组件解析
2.1 交叉熵损失函数
CrossEntropyLoss()是项目中使用的损失函数,专门用于多分类任务。它的核心作用体现在三个方面:
-
差异计算:能够精确计算模型预测值与真实标签之间的差异。在数学上,交叉熵损失可以表示为:
code复制L = -∑(y_i * log(p_i))其中y_i是真实标签的one-hot编码,p_i是模型预测的概率分布。
-
准确性衡量:损失值直接反映了模型分类的准确程度。损失值越小,说明模型预测结果与真实标签越接近。
-
梯度信号:为反向传播算法提供清晰的梯度信号,指导模型参数的更新方向。
在实际应用中,PyTorch的CrossEntropyLoss已经将Softmax操作和负对数似然损失整合在一起,因此模型最后一层不需要额外添加Softmax激活函数。
2.2 优化器配置
项目中使用AdamW优化器,这是一种改进版的Adam优化器,主要配置如下:
python复制optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
关键参数说明:
lr=0.01:设置较大的初始学习率,配合适当的学习率调度策略weight_decay=1e-4:权重衰减系数,本质上是L2正则化项
AdamW与传统Adam的主要区别在于它正确实现了权重衰减(weight decay)与自适应学习率的解耦,使得正则化效果更加稳定。这在深度学习模型中尤为重要,可以有效防止过拟合。
提示:在计算机视觉任务中,AdamW通常比原始Adam表现更好,特别是在使用预训练模型进行微调时。
3. 半监督学习实现
3.1 半监督学习流程
项目实现了典型的半监督学习流程,主要步骤如下:
-
数据准备:
- 收集数据集D,划分为:
- Dₗ={(xᵢ,yᵢ)}ᵢ₌₁ᴸ:带标签数据集(数量L通常较小)
- Dᵤ={xⱼ}ⱼ₌₁ᵁ:未带标签数据集(数量U通常远大于L)
- 收集数据集D,划分为:
-
模型选择:
- 使用预训练的VGG模型作为基础分类器
- 修改最后一层全连接层,适配11分类任务
-
损失函数设计:
- 监督损失Lₛ:带标签数据上的交叉熵损失
- 无监督损失Lᵤ:基于高置信度伪标签的一致性损失
- 总损失:L = Lₛ + λLᵤ(项目中λ动态调整)
-
训练过程:
- 先在带标签数据上训练基础模型
- 当验证准确率>0.6时,开始使用无标签数据
- 对无标签数据生成高置信度(>0.99)的伪标签
- 将伪标签数据加入训练集进行迭代训练
-
评估:
- 在独立的验证集上评估模型性能
- 监控训练损失和验证准确率曲线
3.2 伪标签生成机制
项目中伪标签生成的核心代码如下:
python复制def get_label(self, no_label_loder, model, device, thres):
model = model.to(device)
pred_prob = [] # 保存模型输出的概率
labels = [] # 预测的标签
x = []
y = []
soft = nn.Softmax() # 将模型输出转化成概率分布
with torch.no_grad(): # 禁用梯度
for bat_x, _ in no_label_loder:
bat_x = bat_x.to(device)
pred = model(bat_x)
pred_soft = soft(pred)
pred_max, pred_value = pred_soft.max(1) # 获得最大概率
pred_prob.extend(pred_max.cpu().numpy().tolist())
labels.extend(pred_value.cpu().numpy().tolist())
for index, prob in enumerate(pred_prob): # 遍历标签和概率
if prob > thres: # 只处理高于阈值的样本
x.append(no_label_loder.dataset[index][1]) # 调用到原始的getitem
y.append(labels[index]) # 将预测标签添加到列表
return x, y
关键设计点:
- 使用高阈值(0.99)确保伪标签质量
- 只在模型验证准确率>0.6后才启用伪标签
- 每3个epoch重新生成一次伪标签
- 使用Softmax归一化确保概率解释性
注意:伪标签的质量直接影响半监督学习效果。阈值设置过高可能导致可用样本太少,设置过低则会引入噪声。需要根据具体任务调整。
4. 数据预处理与增强
4.1 数据转换配置
项目为训练集和验证集配置了不同的数据转换策略:
python复制train_transform = transforms.Compose([
transforms.ToPILImage(), # 转换为PIL图像格式
transforms.RandomResizedCrop(224), # 随机裁剪并调整大小
transforms.RandomRotation(50), # 随机旋转(-50°~50°)
transforms.ToTensor() # 转换为Tensor
])
val_transform = transforms.Compose([
transforms.ToPILImage(), # 转换为PIL图像格式
transforms.ToTensor() # 转换为Tensor
])
训练集使用了两种重要的数据增强技术:
- RandomResizedCrop:随机裁剪图像区域并缩放到224×224,增加位置不变性
- RandomRotation:随机旋转图像±50度,增强旋转鲁棒性
验证集则只进行最基本的格式转换,保持数据的原始性以便准确评估模型性能。
4.2 自定义数据集类
项目实现了food_Dataset类来管理数据加载:
python复制class food_Dataset(Dataset):
def __init__(self, path, mode="train"):
self.mode = mode
if mode == "semi":
self.X = self.read_file(path) # 半监督数据只有X
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
def read_file(self, path):
if self.mode == "semi": # 无标签数据读取
file_list = os.listdir(path)
xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)
for j, img_name in enumerate(file_list):
img_path = os.path.join(path, img_name)
img = Image.open(img_path).resize((HW, HW))
xi[j, ...] = img
return xi
else: # 有标签数据读取
for i in tqdm(range(11)): # 11个类别
file_dir = path + "/%02d" % i
file_list = os.listdir(file_dir)
xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)
yi = np.zeros(len(file_list), dtype=np.uint8)
for j, img_name in enumerate(file_list):
img_path = os.path.join(file_dir, img_name)
img = Image.open(img_path).resize((HW, HW))
xi[j, ...] = img
yi[j] = i # 类别标签
if i == 0:
X, Y = xi, yi
else:
X = np.concatenate((X, xi), axis=0)
Y = np.concatenate((Y, yi), axis=0)
return X, Y
数据集类的主要特点:
- 统一处理有标签和无标签数据
- 支持不同的数据转换策略
- 自动完成图像resize和格式转换
- 保持图像与标签的正确对应关系
5. 模型架构与训练
5.1 模型初始化
项目提供了灵活的模型初始化方式,支持自定义模型和预训练模型:
python复制def initialize_model(model_name, num_classes, linear_prob=False, use_pretrained=True):
if model_name == "resnet18":
model_ft = models.resnet18(pretrained=use_pretrained)
set_parameter_requires_grad(model_ft, linear_prob)
num_ftrs = model_ft.fc.in_features
model_ft.fc = nn.Linear(num_ftrs, num_classes)
input_size = 224
# ...其他模型类似...
return model_ft, input_size
关键功能:
- 支持从PyTorch模型库加载预训练模型
- 可冻结特征提取器(linear_prob=True时)
- 自动修改最后一层适配具体分类任务
- 返回适合模型输入的图像尺寸
5.2 训练循环实现
核心训练逻辑在train_val函数中实现:
python复制def train_val(model, train_loader, val_loader, no_label_loader, device, epochs, optimizer, loss, thres, save_path):
model = model.to(device)
plt_train_loss, plt_val_loss = [], []
plt_train_acc, plt_val_acc = [], []
max_acc = 0.0
for epoch in range(epochs):
# 训练阶段
model.train()
for batch_x, batch_y in train_loader:
x, target = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad()
pred = model(x)
train_bat_loss = loss(pred, target)
train_bat_loss.backward()
optimizer.step()
# ...记录损失和准确率...
# 半监督训练(条件触发)
if epoch%3 == 0 and plt_val_acc[-1] > 0.6:
semi_loader = get_semi_loader(no_label_loader, model, device, thres)
if semi_loader:
for batch_x, batch_y in semi_loader:
# ...半监督训练步骤...
# 验证阶段
model.eval()
with torch.no_grad():
for batch_x, batch_y in val_loader:
# ...验证步骤...
# 保存最佳模型
if val_acc > max_acc:
torch.save(model, save_path)
max_acc = val_acc
# 打印训练日志
print(f'[{epoch}/{epochs}] TrainLoss: {plt_train_loss[-1]:.6f} | ValLoss: {plt_val_loss[-1]:.6f} | TrainAcc: {plt_train_acc[-1]:.6f} | ValAcc: {plt_val_acc[-1]:.6f}')
# 绘制训练曲线
plt.plot(plt_train_loss)
plt.plot(plt_val_loss)
plt.title("loss")
plt.legend(["train", "val"])
plt.show()
plt.plot(plt_train_acc)
plt.plot(plt_val_acc)
plt.title("acc")
plt.legend(["train", "val"])
plt.show()
训练流程的关键点:
- 完整的训练-验证循环
- 条件触发的半监督训练
- 模型性能持续监控
- 最佳模型保存机制
- 训练过程可视化
6. 实用技巧与注意事项
6.1 提高训练稳定性的技巧
-
学习率设置:
- 初始学习率设为0.01,适合微调预训练模型
- 可以添加学习率调度器(如ReduceLROnPlateau)动态调整
-
权重衰减:
- 使用1e-4的weight_decay值防止过拟合
- AdamW优化器能正确处理权重衰减与自适应学习率的关系
-
批次大小:
- 设置为16,平衡GPU内存使用和梯度稳定性
- 对于更大的模型可以适当减小
-
随机种子固定:
python复制def seed_everything(seed): torch.manual_seed(seed) torch.cuda.manual_seed(seed) # ...其他随机种子设置...确保实验可重复性
6.2 半监督学习实践建议
-
伪标签质量控制:
- 高置信度阈值(0.99)确保伪标签可靠性
- 只在模型达到一定准确率(>0.6)后使用伪标签
-
数据平衡:
- 监控伪标签的类别分布,避免类别不平衡
- 可以设置每个类别的最大样本数
-
迭代策略:
- 每3个epoch重新生成伪标签
- 随着模型改进,逐步降低置信度阈值
-
无标签数据使用:
- 保持无标签数据的顺序不变(shuffle=False)
- 确保伪标签与原始图像的对应关系正确
6.3 常见问题排查
-
损失不下降:
- 检查学习率是否合适
- 验证数据预处理是否正确
- 确认模型参数是否正常更新
-
过拟合:
- 增加权重衰减系数
- 添加更多数据增强
- 提前停止训练
-
GPU内存不足:
- 减小批次大小
- 使用梯度累积技巧
- 尝试混合精度训练
-
半监督效果不佳:
- 提高伪标签置信度阈值
- 延迟开始半监督训练的时间
- 检查无标签数据的质量
7. 项目扩展与改进方向
7.1 模型层面的改进
-
尝试更先进的模型架构:
- 替换VGG为ResNet、EfficientNet等现代架构
- 使用Transformer-based模型如ViT
-
改进半监督策略:
- 实现Mean Teacher等一致性正则化方法
- 引入MixMatch等先进的半监督算法
- 尝试伪标签与数据增强的组合
-
模型压缩与优化:
- 添加知识蒸馏
- 实现模型量化
- 进行剪枝优化
7.2 数据层面的改进
-
更丰富的数据增强:
- 添加颜色抖动
- 使用随机擦除
- 尝试MixUp/CutMix等混合增强
-
更好的数据采样策略:
- 实现类别平衡采样
- 设计难例挖掘策略
- 根据学习进度动态调整采样权重
-
数据质量提升:
- 添加数据清洗步骤
- 进行异常检测
- 可视化检查数据增强效果
7.3 训练优化改进
-
学习率调度:
- 添加余弦退火学习率
- 实现热重启策略
- 根据验证损失动态调整
-
正则化增强:
- 添加Dropout层
- 使用Label Smoothing
- 尝试Stochastic Depth
-
训练监控与分析:
- 添加TensorBoard日志
- 实现混淆矩阵可视化
- 进行错误案例分析
在实际应用中,可以根据具体任务需求和计算资源,选择最适合的改进方向逐步优化系统性能。这个项目提供了很好的基础框架,可以方便地进行各种扩展和实验。
