ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于PyTorch迁移学习的垃圾分类图像分类实战:从数据预处理到模型调优

基于PyTorch迁移学习的垃圾分类图像分类实战:从数据预处理到模型调优 简介面向机器学习课程设计的Python垃圾分类系统源码包适合高校学生完成课程项目或入门图像分类实战。项目基于TensorFlow 2.3实现涵盖MobileNet模型训练、模型测试与评估矩阵等模块能帮助使用者理解从数据准备到模型部署的完整流程。压缩包共32个文件包含Python脚本、jpg/jpeg样本图片、png结果图、xml配置及xls分类结果表等整体仅2.27MB轻量易用。作者自述评审分达95分以上且经过严格调试确保可运行适合直接下载研习。目前已有1884人学习下载。读者可从中获得完整的垃圾分类识别思路、迁移学习调参方法以及代码组织结构参考便于快速复现并迁移到其他图像分类场景也是课程设计报告撰写的实用素材。1. 课程设计跑通很容易答辩不翻车才是关键每年这个季节都有不少人在做同一个题目基于机器学习的垃圾分类系统。很多同学拿着源码跑通训练准确率也刷上去了结果答辩时被老师一问“为什么用这个模型”“数据怎么处理的”“阈值怎么调的”就卡壳。原因很直接光会跑train.py不算会能把数据流、训练策略、预测逻辑讲清楚才是这门课真正要考察的东西。这篇博客就基于一套可运行的 Python 垃圾分类系统源码把从数据准备、模型选型、训练调参到界面封装的完整链路拆开说。这套方案以图像分类为骨架采用 PyTorch 实现迁移学习硬件门槛不高普通笔记本 CPU 也能完成小规模训练适合作为机器学习课程设计的代码基底也适合想快速上手图像分类实战的开发者。接下来按实际项目推进顺序来讲每一步的取舍和踩坑记录。2. 数据组织与预处理小规模垃圾分类数据集的正确打开方式2.1 目录结构设计与自定义 Dataset垃圾分类图像数据集的常见组织形式是按类别建文件夹但课程设计规模下数据量普遍不大每类几百张到上千张目录结构是否规范直接决定后面能否省心。拿到压缩包后先按下面这个结构检查数据目录dataset/ ├── train/ │ ├── recyclable/ # 可回收垃圾 │ ├── kitchen/ # 厨余垃圾 │ ├── hazardous/ # 有害垃圾 │ └── other/ # 其他垃圾 └── val/ ├── recyclable/ ├── kitchen/ ├── hazardous/ └── other/四分类是最常见的课程设计粒度覆盖面够、数据也相对好凑。如果你的需求是六分类或十分类把可回收拆成纸类、塑料、玻璃等目录结构按同样方式扩展即可代码层面不需要大改。PyTorch 的torchvision.datasets.ImageFolder可以直接按上述目录结构加载数据但如果要实现数据增强、缓存等自定义逻辑建议还是写一个 Dataset 类from torch.utils.data import Dataset from PIL import Image import os class GarbageDataset(Dataset): def __init__(self, root_dir, transformNone): self.root_dir root_dir self.transform transform self.classes sorted(os.listdir(root_dir)) self.class_to_idx {cls: idx for idx, cls in enumerate(self.classes)} self.samples [] for cls in self.classes: cls_dir os.path.join(root_dir, cls) for fname in os.listdir(cls_dir): if fname.lower().endswith((.jpg, .jpeg, .png)): self.samples.append((os.path.join(cls_dir, fname), self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] image Image.open(path).convert(RGB) if self.transform: image self.transform(image) return image, label这段代码有几个关键点。sorted(os.listdir(root_dir))保证类别顺序固定避免不同机器上os.listdir返回顺序不一致导致标签错位convert(RGB)把灰度图或带透明通道的 PNG 统一转成三通道否则 batch 维度对不上class_to_idx在后续推理时需要反向映射回来保存模型时不要丢掉这个映射字典。2.2 数据增强策略小数据集防过拟合的第一道防线垃圾分类图像和 ImageNet 那种自然图像不太一样垃圾袋、塑料瓶、纸箱在照片里往往带有角度、光照和背景噪声的变化但类别本身的共性特征足够明显。因此增强策略要平衡“增加多样性”和“不改变语义”这两点。from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])参数说明RandomResizedCrop的scale(0.7, 1.0)控制裁剪面积占原图比例下限设 0.7 而不是默认的 0.08是因为垃圾分类物体通常是画面主体裁得太多会丢失关键特征RandomRotation(15)旋转角度控制在正负 15 度以内过度旋转会让“瓶盖朝上”这类弱特征失效ColorJitter的 brightness 和 contrast 增强模拟不同光照环境但 saturation 系数不要调太大否则塑料和纸类的颜色区分会被破坏。验证集只用Resize加CenterCrop不做随机增强保证评估指标的稳定性。这里有一个常见的错误做法把增强后的图像直接可视化给老师看然后发现图像被裁切得面目全非——这是正常的增强是作用在训练过程中的不代表最终送入模型的数据分布是扭曲的。2.3 训练集/验证集切分与类别不平衡处理如果原始数据只有一个总目录而没有划分好 train/val最常见的切分方式是按类别分层划分保证每个类别在训练集和验证集中的比例一致。from sklearn.model_selection import train_test_split import shutil # 假设 file_paths 是按类别分组的 (path, label) 列表 train_paths, val_paths train_test_split( file_paths, test_size0.2, stratifylabels, random_state42 )stratifylabels表示按标签比例分层抽样random_state 固定随机种子便于实验复现。如果某个类别图像数量只有几十张建议手动将这个类别的验证集比例降到 0.1 或在报告中说明该类别的结果置信度有限。对于类别不平衡问题比如厨余垃圾样本数是有害垃圾的 3 倍先别急着上过采样看一下训练损失曲线——如果模型在少数类上的 F1 分数明显偏低再给损失函数加类别权重class_counts [len(os.listdir(os.path.join(train_dir, cls))) for cls in classes] total sum(class_counts) weights [total / (len(classes) * c) for c in class_counts] class_weights torch.tensor(weights, dtypetorch.float32).to(device) criterion nn.CrossEntropyLoss(weightclass_weights)权重计算方式是总样本数 / (类别数 * 该类样本数)样本少的类别权重自动变大。加权重后整体准确率可能略降但每类平均 F1 会更好看课程设计答辩时这个点是加分项。3. 迁移学习选型与模型训练预训练权重怎么用才合理3.1 为什么选 ResNet18 而不是自己搭 CNN垃圾分类系统如果从零训练一个深度 CNN在课程设计的数据量下很难收敛。垃圾图像的类间差异并不小纸箱和易拉罐从形状到颜色都有明显区别但类内差异也很大外卖盒有塑料的、纸质的、泡沫的需要模型有一定的特征提取能力。这时候迁移学习的价值就体现出来了——ImageNet 上预训练的模型已经把边缘、纹理、颜色等底层特征的提取器练好了我们只需要微调高层特征和分类头。ResNet18 是课程设计场景下的最优解之一参数量只有 1100 万左右CPU 上训练一个 epoch 耗时可控结构简单、残差连接的原理好讲清楚答辩时不会被追问到难以招架如果换成 ResNet50 或 EfficientNet-B4准确率提升有限但训练时间和显存占用会明显增加。如果你用的是自带 GPU 的机器也可以考虑 EfficientNet-B0 配合混合精度训练但在代码交付时要留好 CPU 回退路径。import torchvision.models as models import torch.nn as nn model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) num_features model.fc.in_features model.fc nn.Linear(num_features, num_classes) for param in model.parameters(): param.requires_grad False for param in model.fc.parameters(): param.requires_grad True这段代码先加载 ImageNet 预训练权重再把最后一层全连接替换成输出维度为num_classes的新分类头。requires_grad的设置表示冻结主干网络、只训练分类头——这种做法在数据量只有几千张时最稳妥主干网络的底层特征已经足够通用微调全部参数反而容易过拟合。如果你的数据量超过 5000 张或者发现只训练分类头时验证准确率卡在某个值上不去可以考虑把最后一个残差块model.layer4也解冻参与训练学习率设为主干的 1/10。具体做法是把layer4参数的requires_grad改为True。3.2 训练超参配置与完整训练循环训练阶段的参数设置直接决定最终效果下面给出一套经过验证的基线配置。参数推荐值说明batch_size32CPU 训练可调至 16 降低内存压力初始学习率0.001微调分类头时常用区间优化器Adam收敛快、对学习率不敏感权重衰减1e-4抑制过拟合配合冻结主干效果更好训练轮数20 ~ 30轮数再多容易过拟合需配合早停学习率调整StepLR, step10, gamma0.1每 10 轮学习率降为原来的 1/10学习率是一个需要额外强调的参数。冻结主干、只训练分类头时0.001 是安全的起点如果解冻了layer4主干部分学习率要降到 0.0001否则预训练权重会被大步长更新破坏掉。你可以在优化器里给不同参数组设置不同的学习率optimizer torch.optim.Adam([ {params: model.fc.parameters(), lr: 0.001}, {params: model.layer4.parameters(), lr: 0.0001} ], weight_decay1e-4)完整训练循环的核心部分如下from torch.optim.lr_scheduler import StepLR scheduler StepLR(optimizer, step_size10, gamma0.1) best_acc 0.0 for epoch in range(30): model.train() running_loss 0.0 for inputs, labels in train_loader: inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * inputs.size(0) model.eval() val_correct 0 val_total 0 with torch.no_grad(): for inputs, labels in val_loader: inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) _, predicted torch.max(outputs, 1) val_correct (predicted labels).sum().item() val_total labels.size(0) val_acc val_correct / val_total scheduler.step() print(fEpoch {epoch1}/30, Loss: {running_loss/len(train_dataset):.4f}, Val Acc: {val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save({ model_state_dict: model.state_dict(), class_to_idx: class_to_idx, num_classes: num_classes, }, best_model.pth)这段代码中有几个细节值得注意。model.train()和model.eval()必须成对出现因为RandomResizedCrop这类增强只在训练模式下生效而 BatchNorm 在训练和推理时的行为本身就不一样训练用 batch 统计量推理用全局统计量。torch.no_grad()在验证时关闭梯度计算既省显存又提升推理速度。torch.max(outputs, 1)返回每行最大值的索引也就是预测类别这里不要用outputs.argmax(dim1)也完全可以但要注意 dim 参数不能写错。保存模型时用字典结构而不是只存state_dict因为class_to_idx和num_classes在推理阶段必须用到——你无法保证下次加载模型时类别顺序和一 致。如果不存这些信息推理时输出的类别索引就是无意义的数字。3.3 早停与模型选择策略训练过程中如果发现验证集准确率不再提升继续训练只会浪费时间并且可能过拟合。常见的做法是记录最佳验证准确率连续 N 个 epoch 没有超过就终止训练early_stop_patience 5 epochs_no_improve 0 for epoch in range(30): # 训练与验证代码省略val_acc 为当前验证准确率 if val_acc best_acc: best_acc val_acc epochs_no_improve 0 torch.save(model.state_dict(), best_model.pth) else: epochs_no_improve 1 if epochs_no_improve early_stop_patience: print(fEarly stopping at epoch {epoch1}) breakpatience 设 5 比较合理数据规模小的数据集验证准确率会有抖动1~2 个 epoch 没提升不代表真的到了瓶颈但 5 个以上没提升基本说明该停下来了。训练曲线最后要截图放进报告里这比任何文字描述都有说服力。4. 系统交互与推理流程把模型封装成可演示的桌面应用4.1 PyQt 还是 Flask课程设计怎么选垃圾分拣系统的交付形态决定了界面的选型。如果只需要本机演示、老师现场看效果推荐 PyQt5 桌面应用——不依赖网络、启动快、交互直观。如果是课程设计有 Web 展示需求用 Flask 起一个本地服务也完全够用。资料包里的源码如果默认是 PyQt 版本你只需要关注推理部分即可UI 代码不需要大改。4.2 预测流程预处理对齐是推理翻车的高发区推理阶段的坑比训练阶段多得多最典型的问题是训练时做了数据增强和标准化推理时忘了做同样的预处理导致输入图像分布和训练数据不一致准确率断崖式下跌。def predict(image_path, model, device, class_to_idx, top_k3): transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): outputs model(input_tensor) probabilities torch.softmax(outputs, dim1) top_probs, top_indices torch.topk(probabilities, ktop_k) idx_to_class {v: k for k, v in class_to_idx.items()} results [] for prob, idx in zip(top_probs[0], top_indices[0]): class_name idx_to_class[idx.item()] results.append((class_name, prob.item())) return results这段代码完整覆盖了一次推理的所有环节。unsqueeze(0)给单张图片增加一个 batch 维度因为 PyTorch 模型要求输入是四维张量(batch, channel, height, width)。torch.softmax把原始 logits 转成概率分布这样能同时输出 Top-3 预测结果及置信度——课程设计演示时把“前三个候选类别概率”显示在界面上比只显示一个结果信息量大很多。一个容易被忽略的点idx_to_class的反向映射依赖训练时保存的class_to_idx。如果你直接从文件夹名称推断类别忘记保存这个映射预测结果就会张冠李戴。上面 3.2 节保存模型时的字典结构就在这里派上用场。checkpoint torch.load(best_model.pth, map_locationdevice) model.load_state_dict(checkpoint[model_state_dict]) class_to_idx checkpoint[class_to_idx]map_locationdevice在 CPU 上加载 GPU 训练的模型时是必须的否则会报RuntimeError: Attempting to deserialize object on a CUDA device。4.3 环境依赖与启动配置课程设计要求源码能在他人的电脑上直接运行所以依赖清单要写得足够明确。一个经过验证的requirements.txt内容大致为torch1.10.0 torchvision0.11.0 Pillow8.0.0 numpy1.19.0 matplotlib3.3.0 scikit-learn0.24.0 PyQt55.15.0注意 PyTorch 的 CPU 版本和 GPU 版本安装方式不同建议在文档里注明只跑课程设计就直接安装 CPU 版本体积更小、不需要 CUDA 环境。如果你用的是 Flask 版本把PyQt5换成Flask2.0.0即可。5. 验证与调优把准确率从 85% 提到 92% 的三个实战技巧5.1 混淆矩阵定位错误集中区准确率不是唯一指标混淆矩阵能告诉你模型到底在哪些类别之间“打架”。课程设计答辩时能说出“我的模型有 30% 的有害垃圾被误判为厨余垃圾原因可能是该类训练样本中塑料袋占比过高”远比“准确率 90%”有说服力。from sklearn.metrics import confusion_matrix, classification_report import matplotlib.pyplot as plt import numpy as np all_preds, all_labels [], [] model.eval() with torch.no_grad(): for inputs, labels in val_loader: inputs inputs.to(device) outputs model(inputs) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_namesclasses)) # 可视化 plt.figure(figsize(8, 6)) plt.imshow(cm, interpolationnearest, cmapBlues) plt.colorbar() tick_marks np.arange(len(classes)) plt.xticks(tick_marks, classes, rotation45) plt.yticks(tick_marks, classes) plt.ylabel(True Label) plt.xlabel(Predicted Label) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150)运行这段代码会生成一个混淆矩阵图和一个分类报告。分类报告里的precision、recall、f1-score是逐类计算的如果你的某个类别 F1 明显低于其他类别优先检查该类别的样本数量和图像质量而不是急着换模型。5.2 测试时增强TTA提升推理稳定性训练时用数据增强推理时只用单次 CenterCrop这其实是“浪费”了模型对多尺度输入的鲁棒性。测试时增强就是推理时对同一张图做多次增强把多次预测结果取平均通常能带来 0.5~2 个百分点的准确率提升。tta_transforms [ transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), transforms.Compose([ transforms.Resize(256), transforms.RandomHorizontalFlip(p1.0), # 水平翻转 transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), ] def predict_with_tta(image_path, model, device, class_to_idx): image Image.open(image_path).convert(RGB) prob_sum None for transform in tta_transforms: input_tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) prob torch.softmax(output, dim1) prob_sum prob if prob_sum is None else prob_sum prob avg_prob prob_sum / len(tta_transforms) top_prob, top_idx torch.max(avg_prob, dim1) idx_to_class {v: k for k, v in class_to_idx.items()} return idx_to_class[top_idx.item()], top_prob.item()需要说明的是TTA 不是银弹。如果你的基线准确率只有 70% 出头说明模型本身欠拟合或数据有问题TTA 只能锦上添花不能雪中送炭。5.3 数据泄露课程设计报告中最容易被忽视的扣分项数据泄露在课程设计里最常见的形态是训练集和验证集来自同一个拍摄批次、同一场景下的连续帧导致验证集准确率虚高。比如用手机拍了一段视频然后抽帧做数据集相邻帧高度相似模型在训练时相当于“见过”验证集的内容。检查方法很简单随机挑选几张验证集图像人工对比训练集中是否存在同一物品、同一背景的相似图像。更系统的做法是检查图像文件的拍摄时间戳或文件名中的序列编码。如果发现泄露把数据按场景分组后重新划分而不是按文件随机乱分。另一个隐蔽的泄露点发生在预处理环节如果你在划分训练/验证集之前就对全量数据做了标准化计算全局均值和方差那验证集的信息就通过统计量流入了训练过程。正确的做法是先划分数据集再在训练集上计算均值方差验证集沿用训练集的统计值。上面所有代码示例中的Normalize参数用的是 ImageNet 统计量就是为了规避这个问题的通用做法。最后一个实质性建议课程设计源码拿到手后不要满眼只盯着train.py里的模型结构和 loss先跑通数据加载、确认类别映射和预处理链路再动训练参数。这套顺序能帮你省掉大量“训练过程报错但不知道为什么”的排查时间让课程设计从“跑通源码”变成“跑通且讲得明白”。本文还有配套的精品资源点击获取
返回列表