ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

MobileNetV2垃圾分类识别实战:数据集划分、模型微调与部署全攻略

MobileNetV2垃圾分类识别实战:数据集划分、模型微调与部署全攻略 简介基于MobileNetV2的垃圾识别分类源码包面向计算机视觉方向的学生与开发者可直接服务于毕业设计、课程设计以及小型垃圾分类项目。资源包共包含191个文件整体大小21.76MB内容涵盖6种常见垃圾类别的176张jpg图像样本、5个Python源码文件、2个pth训练好的模型权重以及xml、json等标注与配置信息同时附带1个iml项目文件方便导入开发环境。代码结构清晰模型训练和推理流程可直接运行根据资源描述其识别准确率可达到99%能快速构建垃圾类别检测应用。当前已有894人学习下载适合需要快速掌握MobileNetV2图像分类实战、或希望基于已有数据集进行迁移学习与调优的读者。使用者可同时获得数据集、训练好的模型和源码省去自行收集标注数据与训练模型的时间直接用于课题展示或功能原型验证。1. 垃圾识别分类为什么绕不开 MobileNetV2很多人拿到“基于Mobilenetv2网络的垃圾识别分类源码6种垃圾数据集训练好的模型.zip”这个包第一反应是先找 predict.py想立刻对一张照片跑出结果。但真正的门槛从来不在“跑一下”而在数据目录怎么摆、预训练权重怎么接、训练时哪些超参决定收敛、模型文件加载后为什么总差一口气。MobileNetV2 被选作这个任务的主干网络不是巧合它用深度可分离卷积把参数量压到 AlexNet 的十几分之一在 CPU 上对一张 224 分辨率图片做前向只要几毫秒分类精度又能靠 ImageNet 预训练权重和微调补上来。这个标题浓缩的是一条完整链路6 类垃圾样本的组织方式、训练脚本的参数选择、模型导出与部署验证。本文按照做这类项目最常见的工程顺序展开新手可以照步骤复现熟手可以直接跳到第 2.3、4.2 和 5.1 看参数边界。2. 6 种垃圾数据集的目录划分、加载方式与增强参数2.1 数据集目录结构决定了 ImageFolder 能不能直接用垃圾分类数据集最常见的组织方式是“每个类别一个目录”这也是 torchvision 里ImageFolder约定俗成的格式。6 类垃圾通常指纸板、玻璃、金属、纸张、塑料和其余垃圾目录名建议用英文小写加下划线避免 Windows 和 Linux 跨平台解压后中文路径出问题。常见目录结构如下dataset/ ├── train/ │ ├── cardboard/ # 纸板 │ ├── glass/ # 玻璃 │ ├── metal/ # 金属 │ ├── paper/ # 纸张 │ ├── plastic/ # 塑料 │ └── trash/ # 其他垃圾 ├── val/ │ └── ... # 与 train 结构一致 └── label_map.jsonlabel_map.json内容很简单但强烈建议在训练前就生成{cardboard: 0, glass: 1, metal: 2, paper: 3, plastic: 4, trash: 5}这个文件必须和ImageFolder.class_to_idx对照检查。因为ImageFolder是按目录名的 ASCII 顺序分配索引的如果目录名是01-cardboard这类带前缀的记法索引顺序会变。训练脚本里classes dataset.classes的打印结果要保留下来后续推理时类别映射全靠它。2.2 手动划分训练集与验证集的最小代码如果 zip 里给的是一个大目录而不是已经分好 train/val需要用脚本按比例拆分。我一般用随机数种子保证划分稳定import random import shutil from pathlib import Path random.seed(42) split_ratio 0.85 src_root Path(dataset/raw_images) train_root Path(dataset/train) val_root Path(dataset/val) for cls_dir in src_root.iterdir(): if not cls_dir.is_dir(): continue images list(cls_dir.glob(*)) random.shuffle(images) split_idx int(len(images) * split_ratio) for img in images[:split_idx]: target train_root / cls_dir.name / img.name target.parent.mkdir(parentsTrue, exist_okTrue) shutil.copy(str(img), str(target)) for img in images[split_idx:]: target val_root / cls_dir.name / img.name target.parent.mkdir(parentsTrue, exist_okTrue) shutil.copy(str(img), str(target))参数说明split_ratio 0.85表示 85% 进训练集。如果你手里的垃圾图片总量只有一两千张可以提高到 0.9超过五千张再考虑 0.8。随机种子固定是必须的否则两次划分的数据不一致后面做模型对比时没法证明精度提升来自训练还是来自数据变化。2.3 数据增强参数哪些值得加哪些加了反而掉点垃圾识别和其他图像分类不同的地方在于同一类垃圾的外观差异极大但同一张图片里通常只有一个主体。所以增强策略不是越猛越好。下表是一组在多个垃圾集上表现稳定的参数范围增强操作推荐值作用与风险Resize256 后 CenterCrop 224保留更多边缘细节避免物体被切掉RandomResizedCropscale(0.6, 1.0)模拟手机拍摄距离远近scale 低于 0.5 会让塑料袋这类软垃圾变形严重RandomHorizontalFlipp0.5水平翻转不改变“瓶子可回收”的语义垃圾分类基本都安全RandomRotation15 度模拟流水线上传送带的任意角度超过 20 度会让金属拉罐的反射光照语义混乱ColorJitterbrightness0.2, contrast0.2适应室内不同灯光饱和度和色相不要动玻璃和透明塑料的区别会被色调扰动破坏验证集做同样的 CenterCrop 但不要随机增强。训练集增强后每个 epoch 看到的图片都不同验证集恒定才能让 loss 曲线有可比性。2.4 类别不均衡不能只看总量6 类垃圾里“其余垃圾”这类往往样本特别多而金属和纸张偏少。如果不处理模型会倾向把所有不确定的物体都划到样本多的类别。最简单的缓解办法是WeightedRandomSampler每个 batch 按类别权重的反比抽样from torch.utils.data import WeightedRandomSampler class_counts [len(list((train_root / c).glob(*))) for c in class_names] total sum(class_counts) weights [total / (len(class_counts) * c) for c in class_counts] sample_weights [] for idx, (_, label) in enumerate(train_dataset.samples): sample_weights.append(weights[label]) sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue)参数说明replacementTrue表示同一张图可能在同一个 epoch 里被抽到多次这是WeightedRandomSampler的常规用法目的是把样本少的类别补上来。num_samples保持和数据集长度一致让每个 epoch 处理的图片总量不变loss 曲线不会因为 epoch 之间 iteration 数不同而抖动。用了 sampler 之后就不要再给DataLoader传shuffleTrue两者冲突且后者会覆盖前者的抽样逻辑。3. MobileNetV2 微调全流程冻结骨干、训练参数与早停3.1 冻结前几层还是全量微调取决于你手上的数据量MobileNetV2 的 ImageNet 预训练权重已经在 1000 类自然图像上学到了通用边缘、纹理和形状特征。垃圾图片虽然和 ImageNet 的猫狗汽车差异很大但低层特征依然有效。常见判断标准是样本总量小于 3000 张冻结骨干只训练分类头样本超过 5000 张全量微调效果更好。冻结骨干的另一个好处是显存占用大幅降低CPU 训练也能跑得动。冻结骨干的代码要注意两个易错点只对requires_grad置 False 还不够优化器里必须只传需要梯度的参数另外 BatchNorm 统计量在冻结模式下要用model.eval()控制否则训练时 BN 层仍会被当前 batch 更新。import torch import torchvision.models as models model models.mobilenet_v2(weightsmodels.MobileNet_V2_Weights.IMAGENET1K_V1) for name, param in model.features.named_parameters(): if int(name.split(.)[0]) 14: # 冻结前 14 个倒残差块 param.requires_grad False model.classifier[1] torch.nn.Linear(model.classifier[1].in_features, 6)参数说明MobileNet_V2_Weights.IMAGENET1K_V1是较新版 torchvision 的写法旧版用pretrainedTrue后者在新版本里已经标记为弃用。features一共 18 层这里冻结前 14 层只微调最后 4 层和新增的分类头是“半冻结”策略比纯训分类头更稳也比全量微调慢不了太多。3.2 训练脚本里的 4 个关键参数optimizer、lr、scheduler、epochs参数推荐值说明optimizerAdamWweight_decay1e-4Adam 加权重衰减对 BN 层参数也生效SGDmomentum 效果类似但收敛慢初始 lr冻结骨干 1e-3全量微调 1e-4全量微调用大 lr 容易把预训练权重洗掉schedulerReduceLROnPlateaufactor0.3patience3监控 val_loss连续 3 个 epoch 不降就降 lrepochs30 起步配合早停小数据集 20 个 epoch 左右就能稳定超过 50 基本过拟合训练循环主体代码如下这里只看核心部分criterion torch.nn.CrossEntropyLoss() optimizer torch.optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, factor0.3, patience3) best_val_loss float(inf) patience_counter 0 for epoch in range(30): model.train() train_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() train_loss loss.item() * images.size(0) val_loss, val_acc evaluate(model, val_loader, criterion, device) scheduler.step(val_loss) if val_loss best_val_loss: best_val_loss val_loss torch.save({ state_dict: model.state_dict(), class_names: class_names, epoch: epoch, }, best_model.pth) patience_counter 0 else: patience_counter 1 if patience_counter 5: print(fEarly stop at epoch {epoch}) break逻辑说明filter(lambda p: p.requires_grad, model.parameters())是配合前面冻结骨干的关键它保证优化器只更新可训练参数。保存 checkpoint 时我习惯存 dict 而不是裸的model.state_dict()因为class_names和epoch必须一起封存否则模型在一个月后加载时类名顺序已经忘了。早停的阈值设成验证 loss 连续 5 个 epoch 不降低比光看准确率更稳。3.3 用混淆矩阵而不是只看准确率6 类垃圾里玻璃和透明塑料在单张 RGB 图上本来就难分准确率 90% 和 92% 的差距往往集中在某一类上。评估时把混淆矩阵打印出来才能知道模型真正混淆的是哪一对。简化版代码如下import numpy as np from sklearn.metrics import confusion_matrix, classification_report model.eval() all_preds, all_labels [], [] with torch.no_grad(): for images, labels in val_loader: outputs model(images.to(device)) preds outputs.argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_namesclass_names, digits3))观察重点对角线以外数值最大的位置就是模型最容易搞混的类别对。我见过最常见的是 metal 和 glass 互相串因为金属拉罐和高脚杯都有强烈反光。如果这种情况出现在你的验证集里优先检查增强里有没有开ColorJitter色相扰动其次再考虑加数据。4. 从 checkpoint 到本地推理加载模型、预处理对齐与置信度阈值4.1 解压 zip 后先做三件事拿到“训练好的模型.zip”别急着跑。先看里面有没有这几个文件模型权重、类别映射、训练参数记录。常见打包布局如下unzip 基于Mobilenetv2网络的垃圾识别分类.zip -d waste_model cd waste_model ls -Rwaste_model/ ├── best_model.pth ├── label_map.json └── inference.py如果没有label_map.json可以从best_model.pth里的class_names字段恢复。加载模型的关键坑在于原始训练脚本里分类头的输出维度是 6加载时也必须先把model.classifier[1]换成对应维度的 Linear再load_state_dict否则会出现尺寸不匹配报错。import torch import torchvision.models as models ckpt torch.load(best_model.pth, map_locationcpu) model models.mobilenet_v2(weightsNone) model.classifier[1] torch.nn.Linear(model.classifier[1].in_features, 6) state_dict ckpt.get(state_dict, ckpt) model.load_state_dict(state_dict, strictTrue) model.eval()这里strictTrue是推荐的任何 key 对不上都说明网络结构没重建对。新手最常见的两个错误是把整个 checkpoint dict 直接load_state_dict或者忘了替换分类头。注意torch.load时map_locationcpu可以避免 CUDA 环境不一致的问题。4.2 单张图片推理的预处理必须和训练严格一致推理代码出问题最多的地方不是模型而是预处理。训练时用Resize(256)CenterCrop(224)推理时就必须一字不差地复现。很多人平时推理用Resize(224)结果玻璃和塑料的边界直接被拉伸变形精度肉眼可见掉 5 个百分点。完整推理代码from PIL import Image from torchvision import transforms 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]), ]) def predict(image_path, model, label_map, threshold0.5): image Image.open(image_path).convert(RGB) tensor transform(image).unsqueeze(0) with torch.no_grad(): logits model(tensor) probs torch.softmax(logits, dim1).squeeze(0) max_prob, max_idx torch.max(probs, dim0) max_idx max_idx.item() if max_prob.item() threshold: return unknown, max_prob.item() label [k for k, v in label_map.items() if v max_idx] return label[0] if label else unknown, max_prob.item()参数说明unsqueeze(0)把单张图片变成 batch 维度为 1 的输入threshold表示置信度低于该值就返回 unknown。这里用 0.5 是保守做法不同类别代价不一样。有害垃圾如果误判成玻璃后果比玻璃误判成有害更严重所以阈值可以按类别拉开见下面表格。4.3 不同类别的置信度阈值应该分开设垃圾识别不能用一个统一阈值覆盖所有类别。下表是我在多个项目里常用的阈值方案类别阈值理由glass0.6很容易和塑料、金属互相混淆低置信度宁可拒识metal0.5金属特征明显误判风险不高paper0.4纸张纹理清楚低置信度也可以接受plastic0.6透明塑料和玻璃难区分需要更严格cardboard0.5纸板颜色特征强烈trash0.7负类任何不确定的都倾向归到这里必须拉高推理时按预测类别取对应阈值而不是全局比较。实现上把上面predict函数的threshold参数换成label_threshold {glass: 0.6, ...}预测出类别后再查表判断是否接受。4.4 常见加载与推理报错排查报错现象原因处理方式Missing key(s) in state_dict分类头输出维度不是 6或网络结构不匹配先把model.classifier[1]换成 in_features 到 6 的 Linear图片全预测成一类预处理未归一化或数据集中该类别占比过高检查 mean/std再检查训练时是否用了类别均衡推理速度比预期慢模型没有eval()BN 和 Dropout 仍处于训练模式推理前调用model.eval()并包在torch.no_grad()里单张图片预测消耗巨大显存batch 维度没有 torch.no_grad保留了大量计算图用with torch.no_grad():包裹前向核心理念加载模型时别只看报错报错信息里“size mismatch”后的数字会直接告诉你训练时的输出维度和当前模型输出维度差在哪。5. 进一步部署TorchScript 导出与 CPU 推理验证技巧垃圾识别项目做到“模型能跑”只是第一步真正要落地到边缘设备通常需要把模型导出成不带预训练权重依赖的文件。PyTorch 官方的torch.jit.trace是最简单的方式它对 MobileNetV2 这种结构固定的 CNN 很友好。import torch model.eval() example_input torch.randn(1, 3, 224, 224) traced_model torch.jit.trace(model, example_input) traced_model.save(mobile_net_v2_scripted.pt)导出的调用方式与普通模型完全一致但加载时不依赖 torchvision 的网络结构定义部署环境不需要安装 torchvision只需要 torch。注意 trace 必须在eval()模式下执行且输入尺寸必须和训练一致否则 trace 会把动态 shape 的错误假设固化进图里。导出的文件通常比原始 pth 小 2 到 3MB。如果你手里的边缘设备是 CPU 为主可以再试一下动态量化只量化 Linear 层quantized_model torch.quantization.quantize_dynamic( traced_model, {torch.nn.Linear}, dtypetorch.qint8 ) torch.jit.save(quantized_model, mobile_net_v2_quantized.pt)要注意 MobileNetV2 的主体是卷积层动态量化对 Linear 层之外的部分不生效实测收益有限。更有效的做法是静态量化用训练集抽样 200 到 400 张图作为校准数据那部分代码会多出不少但精度损失通常能控制在 2% 以内。量化前后的对比验证不要只跑一两张图看感觉用脚本批量跑验证集time python inference.py --image test_glass_001.jpg --model mobile_net_v2_quantized.pt time python inference.py --image test_glass_002.jpg --model mobile_net_v2_scripted.pt至少跑 30 次取中位数首轮推理包含模型加载和算子预热不计入对比。同时对比两种情况下的混淆矩阵差异重点关注 glass 和 plastic 之间是否出现新的串扰。如果量化让某一类掉点特别多可以把那个类别单独设置更高的阈值而不是整体抬升。最后养成一个习惯在模型文件夹里配一个model_card.txt写清楚训练时用哪张 label_map、图片尺寸、Normalize 参数和阈值表。半年后你重新翻这个 zip能救命的一定不是代码本身而是这份配置记录。本文还有配套的精品资源点击获取
返回列表