ARTICLE DETAIL

资讯详情

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

基于CNN的智能垃圾分类系统:从数据清洗到Gradio部署全流程

基于CNN的智能垃圾分类系统:从数据清洗到Gradio部署全流程 简介这是一套面向计算机、人工智能、自动化等专业学生与教师的智能垃圾分类系统毕业设计源码基于深度学习卷积神经网络实现可用于课程设计、大作业或毕设参考。项目代码经过调试测试答辩评审分达98分适合小白学习与进阶基础较好者也可在此基础上修改扩展功能。资源包共216个文件约17.29MB包含Python训练脚本、checkpoint与data权重文件、Java与xml配置、gradle构建脚本以及png、jpg图片素材和md说明文档覆盖模型训练、界面配置与项目说明等模块。目前已有182人学习下载。整体目录结构清晰读者可获取完整的垃圾分类识别方案、模型权重与运行说明便于快速复现实验、理解卷积神经网络在图像分类中的落地流程并对照代码梳理训练与部署思路。1. 智能垃圾分类系统从一张照片到正确桶位的完整链路你拍一张外卖餐盒的照片系统要在几百毫秒内告诉你它属于干垃圾、湿垃圾、可回收物还是有害垃圾——这件事听起来像是一个简单的图像分类任务但真正动手做的时候你会发现坑远比想象中多。光照变化、拍摄角度、物体遮挡、类别边界模糊比如沾了油渍的纸盒到底算干垃圾还是可回收每一个因素都能让模型精度掉一大截。基于深度学习卷积神经网络的智能垃圾分类系统核心就是用 CNN 提取图像特征再通过分类头输出类别概率。这类系统在 Python 生态下已经有非常成熟的实现路径从数据采集、模型选型、训练调参到推理部署整条链路都可以在一台带显卡的普通电脑上跑通。这篇文章面向正在做毕业设计、想找一个能落地、能讲清楚技术细节的深度学习项目的同学也面向想快速搭一个垃圾分类 demo 的工程师。我会把选型理由、代码实现、参数设置和踩过的坑都摊开讲让你看完能直接动手复现。2. 数据集构建与 CNN 模型选型为什么不是随便找个网络就能用2.1 垃圾分类数据集从哪来、怎么清洗公开的垃圾分类数据集常见的有 TrashNet、Kaggle 上的垃圾分类竞赛数据以及国内一些高校整理的中文垃圾分类数据集。TrashNet 只有 6 个类别、约 2500 张图片规模偏小直接拿来训练容易过拟合。我的做法是以 TrashNet 为基础再补充自己拍摄的图片和网络爬取的图片把类别扩展到 4 大类可回收物、有害垃圾、湿垃圾、干垃圾每个类别至少保证 1500 张有效图片。数据清洗这一步不能省。常见的问题是同一张图重复出现、标签错误、图片尺寸差异过大、存在大量纯色背景的“摆拍图”。我一般会写一个脚本做去重和尺寸统计import os import hashlib from PIL import Image from collections import Counter def file_hash(filepath): 计算文件的 MD5 值用于去重 with open(filepath, rb) as f: return hashlib.md5(f.read()).hexdigest() def clean_dataset(root_dir): 遍历数据集输出去重和尺寸统计结果 hashes {} sizes [] duplicates [] for cls in os.listdir(root_dir): cls_dir os.path.join(root_dir, cls) if not os.path.isdir(cls_dir): continue for img_name in os.listdir(cls_dir): img_path os.path.join(cls_dir, img_name) h file_hash(img_path) if h in hashes: duplicates.append((img_path, hashes[h])) else: hashes[h] img_path # 统计尺寸 try: with Image.open(img_path) as im: sizes.append(im.size) except Exception as e: print(f无法读取: {img_path}, 错误: {e}) print(f重复图片数量: {len(duplicates)}) print(f尺寸分布: {Counter(sizes).most_common(5)}) return duplicates # 调用示例 dups clean_dataset(./dataset/train)这段代码的逻辑很直接用 MD5 做文件级去重同时统计图片尺寸分布。参数方面root_dir指向你的训练集根目录目录结构应该是train/类别名/图片文件。去重之后对于重复图片保留第一张即可。尺寸统计的目的是决定后续统一缩放到多大——如果大部分图片长边在 500 到 800 之间那统一缩放到 224×224 或 256×256 都合理如果原图普遍很小比如 100×100强行放大到 224 只会引入插值噪声这时候要考虑是否补充更多高质量数据。清洗完之后还需要做类别平衡。如果某个类别图片数量明显偏少可以用数据增强随机裁剪、旋转、颜色抖动来扩充或者用WeightedRandomSampler在训练时给少数类更高采样权重。2.2 选 ResNet 还是 MobileNet精度和速度的权衡CNN 模型选型这件事没有“最好”只有“最合适”。毕业设计场景下你需要在精度、训练速度、推理速度和代码复杂度之间做取舍。我对比过几个常见骨干网络在垃圾分类任务上的表现模型参数量输入尺寸训练速度单卡适合场景LeNet-5约 6 万32×32极快教学演示精度低VGG16约 1.38 亿224×224慢精度尚可但太重ResNet50约 2500 万224×224中等精度和速度平衡MobileNetV3约 540 万224×224快移动端/嵌入式部署EfficientNet-B0约 530 万224×224中等偏快精度优先且资源有限如果你是做毕业设计我建议从 ResNet50 或 EfficientNet-B0 入手。ResNet 的残差结构解决了深层网络梯度消失的问题代码实现成熟预训练权重容易获取。EfficientNet 系列通过复合缩放策略在同等参数量下通常有更高精度但代码稍复杂一些。MobileNetV3 适合你后续想把模型部署到手机或树莓派上的场景。选型时还要考虑你的显卡显存。ResNet50 在 batch size 为 32、输入 224×224 时大约需要 6GB 显存。如果显存不够要么减小 batch size要么换更轻的模型要么用梯度累积来模拟大 batch。2.3 用迁移学习把预训练权重用起来从零训练一个 CNN 在垃圾分类这种中等规模数据集上很难收敛到理想精度。迁移学习是标准做法加载在 ImageNet 上预训练的权重替换最后的全连接层然后分阶段微调。import torch import torch.nn as nn from torchvision import models def build_model(num_classes4, model_nameresnet50, pretrainedTrue): 构建迁移学习模型 if model_name resnet50: model models.resnet50(pretrainedpretrained) # 替换最后的全连接层 in_features model.fc.in_features model.fc nn.Sequential( nn.Dropout(0.3), # 防止过拟合 nn.Linear(in_features, 256), nn.ReLU(inplaceTrue), nn.Dropout(0.2), nn.Linear(256, num_classes) ) elif model_name mobilenet_v3: model models.mobilenet_v3_large(pretrainedpretrained) in_features model.classifier[-1].in_features model.classifier[-1] nn.Linear(in_features, num_classes) return model # 构建模型 model build_model(num_classes4, model_nameresnet50) print(f模型参数量: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M)这段代码的关键点在于pretrainedTrue会下载 ImageNet 预训练权重需要网络连接然后把 ResNet50 原本的 1000 类输出层替换成适合你类别数的自定义分类头。我加了两层 Dropout 和一层中间全连接层目的是在小数据集上增强泛化能力。num_classes根据你的实际类别数设置垃圾分类通常是 4 类或 6 类。Dropout(0.3)和Dropout(0.2)是经验值如果你的训练集很大每类超过 5000 张可以适当降低 dropout 比例。训练策略上我一般分两步先冻结骨干网络只训练新加的分类头 5 到 10 个 epoch然后解冻全部参数用更小的学习率比如 1e-4做全局微调。这样能避免随机初始化的分类头在训练初期产生大梯度破坏预训练权重。3. 训练流程与参数配置把模型真正跑起来3.1 数据加载与增强的工程实现数据管道是训练流程里最容易被忽视但影响很大的环节。PyTorch 的DataLoader配合torchvision.transforms可以完成大部分工作但有几个细节需要注意训练集和验证集的增强策略必须不同验证集不能做随机增强num_workers设置要合理设太大反而会因为进程切换拖慢速度。from torchvision import transforms from torch.utils.data import DataLoader, WeightedRandomSampler from torchvision.datasets import ImageFolder import numpy as np # 训练集增强随机裁剪、翻转、颜色抖动 train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), # 随机裁剪并缩放 transforms.RandomHorizontalFlip(p0.5), # 水平翻转 transforms.RandomRotation(15), # 随机旋转 transforms.ColorJitter(brightness0.2, contrast0.2, # 颜色抖动 saturation0.2, hue0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], # ImageNet 均值 std[0.229, 0.224, 0.225]) # ImageNet 标准差 ]) # 验证集只做缩放和归一化 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]) ]) # 加载数据集 train_dataset ImageFolder(./dataset/train, transformtrain_transform) val_dataset ImageFolder(./dataset/val, transformval_transform) # 处理类别不平衡计算每个类别的采样权重 targets [s[1] for s in train_dataset.samples] class_counts np.bincount(targets) class_weights 1.0 / class_counts sample_weights [class_weights[t] for t in targets] sampler WeightedRandomSampler(sample_weights, len(sample_weights)) train_loader DataLoader(train_dataset, batch_size32, samplersampler, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)RandomResizedCrop(224, scale(0.7, 1.0))表示随机裁剪原图 70% 到 100% 的区域再缩放到 224×224这个参数对垃圾分类任务比较合适因为垃圾物体通常占据画面主体。ColorJitter的四个参数分别控制亮度、对比度、饱和度和色调的扰动幅度垃圾分类场景下光照变化大适当增强颜色鲁棒性有帮助。Normalize用的均值和标准差是 ImageNet 的统计值因为骨干网络是在 ImageNet 上预训练的保持一致的归一化方式很重要。WeightedRandomSampler解决了类别不平衡问题图片少的类别会被更频繁地采样。注意用了sampler之后shuffle必须设为False否则会报错。3.2 损失函数、优化器和学习率调度垃圾分类是一个单标签多分类任务损失函数用交叉熵就够了。但如果你发现某些类别之间容易混淆比如干垃圾和可回收物边界模糊可以考虑用标签平滑Label Smoothing来缓解过拟合。import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) # 损失函数带标签平滑的交叉熵 criterion nn.CrossEntropyLoss(label_smoothing0.1) # 优化器AdamW 比 Adam 更适合带权重衰减的场景 optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) # 学习率调度余弦退火 scheduler CosineAnnealingLR(optimizer, T_max30, eta_min1e-6) # 训练循环 def train_one_epoch(model, loader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) _, predicted outputs.max(1) correct predicted.eq(labels).sum().item() total labels.size(0) epoch_loss running_loss / total epoch_acc correct / total return epoch_loss, epoch_acclabel_smoothing0.1的作用是把硬标签如 [0, 1, 0, 0]变成软标签如 [0.025, 0.925, 0.025, 0.025]让模型不要对某一类过度自信通常能提升 1 到 2 个百分点的验证集精度。AdamW相比Adam把权重衰减从梯度更新中解耦出来在微调预训练模型时更稳定。学习率用1e-3是针对新加的分类头如果你解冻了全部参数做全局微调要降到1e-4甚至更低。CosineAnnealingLR让学习率按余弦曲线从初始值降到接近零T_max是半个周期通常设为总 epoch 数。训练过程中要监控训练损失、验证损失和验证精度。如果训练损失持续下降但验证损失开始上升说明过拟合了需要增加数据增强、提高 dropout 或加早停。如果两者都下不去可能是学习率太大或模型容量不够。3.3 训练过程监控与模型保存训练不是跑完就完事你需要知道每一轮发生了什么。我习惯在每个 epoch 结束后打印指标并保存验证集上表现最好的模型权重。best_acc 0.0 patience 7 # 早停耐心值 counter 0 # 验证精度未提升的轮数 for epoch in range(30): train_loss, train_acc train_one_epoch(model, train_loader, criterion, optimizer, device) # 验证阶段 model.eval() val_loss 0.0 val_correct 0 val_total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() * images.size(0) _, predicted outputs.max(1) val_correct predicted.eq(labels).sum().item() val_total labels.size(0) val_loss val_loss / val_total val_acc val_correct / val_total scheduler.step() print(fEpoch {epoch1:02d} | Train Loss: {train_loss:.4f} Acc: {train_acc:.4f} | fVal Loss: {val_loss:.4f} Acc: {val_acc:.4f} | LR: {scheduler.get_last_lr()[0]:.6f}) # 保存最佳模型 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) counter 0 else: counter 1 if counter patience: print(f验证精度连续 {patience} 轮未提升提前停止训练) break早停机制是防止过拟合的有效手段。patience7表示验证精度连续 7 轮没有超过历史最佳就停止训练。保存模型时用state_dict()而不是整个模型对象这样加载时更灵活文件也更小。scheduler.step()放在验证之后调用确保学习率调整基于当前 epoch 的完整信息。4. 推理部署与界面集成让系统真正能用4.1 单张图片推理与批量测试训练完之后你需要一个干净的推理脚本能加载模型权重、处理输入图片、输出预测类别和置信度。from PIL import Image import torch.nn.functional as F # 类别名称映射根据你的数据集调整 class_names [可回收物, 有害垃圾, 湿垃圾, 干垃圾] def predict_image(model, image_path, transform, device): 对单张图片进行预测 model.eval() image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0).to(device) # 增加 batch 维度 with torch.no_grad(): outputs model(input_tensor) probabilities F.softmax(outputs, dim1) confidence, predicted probabilities.max(1) pred_class class_names[predicted.item()] conf confidence.item() # 输出所有类别的概率 print(f预测结果: {pred_class} (置信度: {conf:.4f})) for i, name in enumerate(class_names): print(f {name}: {probabilities[0][i].item():.4f}) return pred_class, conf # 加载模型 model build_model(num_classes4, model_nameresnet50, pretrainedFalse) model.load_state_dict(torch.load(best_model.pth, map_locationdevice)) model model.to(device) # 推理 predict_image(model, ./test_images/sample.jpg, val_transform, device)推理时的transform必须和验证集一致不能加随机增强。unsqueeze(0)是给单张图片增加 batch 维度因为模型期望输入是[B, C, H, W]。F.softmax把 logits 转成概率分布方便看置信度。如果置信度低于某个阈值比如 0.6可以在业务逻辑里加一个“不确定”的兜底提示而不是强行分类。4.2 用 Gradio 快速搭一个可交互的演示界面毕业设计答辩时一个能实时演示的界面比一堆命令行输出有说服力得多。Gradio 是目前最省事的方案几行代码就能搭一个网页界面。import gradio as gr def classify_image(image): Gradio 回调函数接收 PIL 图片返回预测结果 if image is None: return 请上传一张图片 model.eval() input_tensor val_transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs model(input_tensor) probabilities F.softmax(outputs, dim1) confidence, predicted probabilities.max(1) pred_class class_names[predicted.item()] conf confidence.item() # 返回格式化的结果 result f分类结果: {pred_class}\n置信度: {conf:.2%}\n\n result 各类别概率:\n for i, name in enumerate(class_names): result f {name}: {probabilities[0][i].item():.2%}\n return result # 创建界面 interface gr.Interface( fnclassify_image, inputsgr.Image(typepil, label上传垃圾图片), outputsgr.Textbox(label分类结果, lines8), title智能垃圾分类系统, description上传一张垃圾图片系统会自动识别其类别 ) interface.launch(server_name0.0.0.0, server_port7860)gr.Image(typepil)表示输入是 PIL 图像对象Gradio 会自动处理上传和格式转换。server_name0.0.0.0让局域网内其他设备也能访问方便答辩时用手机演示。server_port可以改成你喜欢的端口。这个界面虽然简单但足够展示核心功能。如果你想做得更完整可以加一个摄像头输入组件用gr.Image(sourcewebcam)实现实时拍摄分类。4.3 模型导出与轻量化部署思路如果你的毕业设计涉及嵌入式设备或移动端部署需要把 PyTorch 模型转成 ONNX 或 TorchScript 格式再用推理引擎加速。# 导出为 ONNX 格式 dummy_input torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, garbage_classifier.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}}, opset_version11 ) print(ONNX 模型导出完成) # 验证 ONNX 模型 import onnxruntime as ort import numpy as np ort_session ort.InferenceSession(garbage_classifier.onnx) ort_inputs {ort_session.get_inputs()[0].name: dummy_input.cpu().numpy()} ort_outputs ort_session.run(None, ort_inputs) print(fONNX 推理输出形状: {ort_outputs[0].shape})dynamic_axes参数让导出的模型支持动态 batch size这样推理时可以用单张也可以用批量。opset_version11是兼容性比较好的版本大部分推理引擎都支持。导出后用onnxruntime验证一遍确保输出形状和 PyTorch 一致。如果要在树莓派上部署可以用 ONNX Runtime 的 ARM 版本或者转成 TensorFlow Lite 格式配合 Coral USB 加速棒使用。5. 避坑与排查那些让我熬夜的翻车现场5.1 训练精度很高但实际用起来一塌糊涂现象验证集精度到了 95%但拿手机拍的照片去测经常分错尤其是湿垃圾和干垃圾混淆严重。原因训练集里的图片大多是白底摆拍图和真实场景差异太大。模型学到的是背景特征而不是物体特征。解决在数据增强里加入随机背景替换或者用 RandAugment、CutMix 这类更强的增强策略。更直接的办法是补充真实场景拍摄的数据哪怕每类只加 200 张效果也会明显改善。另外可以在验证集里专门划出一个“真实场景测试集”用来评估模型的泛化能力。5.2 训练损失不下降精度卡在 25%现象四分类任务训练了 10 个 epoch精度一直在 25% 左右相当于随机猜。原因最常见的是数据标签和类别目录不对应。ImageFolder是按文件夹名称排序来分配标签的如果你手动改了类别顺序但没改class_names列表预测结果就会全部错位。解决打印train_dataset.class_to_idx确认类别映射关系然后同步更新推理脚本里的class_names。另一个可能原因是学习率太大导致模型发散把学习率降到 1e-4 试试。5.3 显存溢出OOM导致训练中断现象训练到一半报CUDA out of memorybatch size 已经调到 8 了还是不行。原因除了 batch size输入分辨率、模型参数量、是否累积了梯度都会影响显存占用。另外 PyTorch 的缓存分配器不会立即释放显存有时候看起来还有余量但实际已经碎片化了。解决先把输入尺寸从 224 降到 192 或 160 试试用torch.cuda.empty_cache()清理缓存如果还不够用梯度累积模拟大 batchaccumulation_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): outputs model(images.to(device)) loss criterion(outputs, labels.to(device)) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这样等效于 batch size 扩大了 4 倍但显存占用不变。5.4 推理速度太慢单张图片要等好几秒现象Gradio 界面上传图片后要等 3 到 5 秒才出结果答辩演示时很尴尬。原因模型在 CPU 上推理或者没有用torch.no_grad()导致计算图被构建或者输入图片没有做尺寸限制。解决确保推理时用with torch.no_grad():包裹如果机器有显卡确认模型和数据都.to(device)了在val_transform里加transforms.Resize(256)限制输入尺寸避免超大图直接进网络。如果还是慢考虑换 MobileNetV3 或对模型做量化# 动态量化仅适用于 CPU 推理 quantized_model torch.quantization.quantize_dynamic( model, {nn.Linear}, dtypetorch.qint8 )量化后模型大小会缩小到原来的 1/4 左右CPU 推理速度提升 2 到 3 倍精度损失通常在 1 个百分点以内。5.5 类别不平衡导致少数类几乎不被识别现象有害垃圾只有 300 张图其他类别有 2000 张训练完后有害垃圾的召回率不到 40%。原因标准交叉熵损失对每个样本一视同仁多数类主导了梯度方向。解决除了前面提到的WeightedRandomSampler还可以在损失函数里给少数类更高权重class_weights torch.tensor([1.0, 3.0, 1.0, 1.0]).to(device) # 有害垃圾权重更高 criterion nn.CrossEntropyLoss(weightclass_weights, label_smoothing0.1)权重值根据类别频率的倒数来设置但不要设得过于极端否则模型会偏向少数类而牺牲整体精度。通常把权重控制在 1 到 5 之间比较稳妥。6. 把精度再往上推一推几个我常用的进阶技巧6.1 用测试时增强TTA榨出最后两个点模型训练完之后推理阶段还有免费的精度提升空间。测试时增强的思路是对同一张测试图片做多次不同的变换比如水平翻转、不同尺度裁剪分别推理后把概率平均。这样做几乎不增加训练成本但通常能提升 1 到 2 个百分点的精度。def predict_with_tta(model, image, device): 测试时增强推理 model.eval() # 定义多种变换 tta_transforms [ val_transform, # 原始 transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.RandomHorizontalFlip(p1.0), # 强制翻转 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), transforms.Compose([ transforms.Resize(280), # 不同尺度 transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), ] all_probs [] for t in tta_transforms: input_tensor t(image).unsqueeze(0).to(device) with torch.no_grad(): outputs model(input_tensor) probs F.softmax(outputs, dim1) all_probs.append(probs) # 平均概率 avg_probs torch.mean(torch.stack(all_probs), dim0) confidence, predicted avg_probs.max(1) return class_names[predicted.item()], confidence.item(), avg_probs[0]TTA 的代价是推理时间变成原来的 N 倍N 是变换数量。如果对实时性要求高可以只用两种变换原始 翻转精度提升大约 0.5 到 1 个百分点推理时间只增加一倍。6.2 用混淆矩阵找到模型的“死穴”精度数字看不出模型到底在哪些类别上犯错。混淆矩阵能直观展示类别之间的误判分布帮你定位问题。from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt def evaluate_model(model, loader, device): 在验证集上评估模型输出混淆矩阵 model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in loader: images images.to(device) outputs model(images) _, predicted outputs.max(1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.numpy()) # 打印分类报告 print(classification_report(all_labels, all_preds, target_namesclass_names)) # 绘制混淆矩阵 cm confusion_matrix(all_labels, all_preds) plt.figure(figsize(8, 6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclass_names, yticklabelsclass_names) plt.xlabel(预测类别) plt.ylabel(真实类别) plt.title(混淆矩阵) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150) plt.show() return cm cm evaluate_model(model, val_loader, device)拿到混淆矩阵后重点看非对角线上的大数字。如果“湿垃圾”和“干垃圾”之间的误判特别多说明这两个类别的特征区分度不够。解决办法有两个方向一是补充更多这两类的边界样本比如沾了油的纸盒、带包装的食物残渣二是考虑用多标签分类或层次分类先粗分再细分。6.3 模型集成简单但有效的最后手段如果你训练了多个不同骨干网络的模型比如 ResNet50 和 EfficientNet-B0 各一个可以把它们的预测概率平均通常能再提升 1 到 3 个百分点。代价是推理时需要加载多个模型显存和计算量都翻倍。def ensemble_predict(models, image, transform, device): 多模型集成推理 all_probs [] for model in models: model.eval() input_tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs model(input_tensor) probs F.softmax(outputs, dim1) all_probs.append(probs) avg_probs torch.mean(torch.stack(all_probs), dim0) confidence, predicted avg_probs.max(1) return class_names[predicted.item()], confidence.item() # 加载多个模型 model1 build_model(4, resnet50, pretrainedFalse) model1.load_state_dict(torch.load(best_resnet50.pth)) model1 model1.to(device) model2 build_model(4, mobilenet_v3, pretrainedFalse) model2.load_state_dict(torch.load(best_mobilenet.pth)) model2 model2.to(device) result ensemble_predict([model1, model2], Image.open(./test.jpg), val_transform, device) print(f集成预测: {result[0]}, 置信度: {result[1]:.4f})集成方法在毕业设计答辩时是一个很好的加分项因为它体现了你对模型泛化能力的理解。但要注意集成的前提是各个模型之间有足够的差异性——如果两个模型都是 ResNet50 只是随机种子不同集成效果很有限。用不同骨干网络、不同输入尺寸、甚至不同数据增强策略训练出来的模型集成收益才明显。我自己做这类项目最大的教训是不要一上来就追求 SOTA 模型先把数据质量、训练流程和推理链路跑通再逐步优化。很多时候把数据清洗干净、把类别平衡做好比换一个更复杂的网络带来的提升大得多。另外答辩时老师更看重你对整个系统的理解深度而不是单纯的精度数字——能说清楚为什么选这个模型、每个参数为什么这么设、遇到问题怎么排查比一个 99% 的精度更有说服力。希望帮到你。本文还有配套的精品资源点击获取
返回列表