news 2026/9/12 10:09:15

PyTorch四类垃圾图像识别端到端实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch四类垃圾图像识别端到端实战

简介:本资源是一套基于Python与神经网络图像识别技术实现的垃圾分类毕业设计项目,面向计算机、人工智能、自动化等专业学生及教师,适用于课程设计、大作业或毕业设计实践。项目包含完整可运行源码与配套文档,答辩评分高达98分,兼顾入门学习与进阶二次开发需求。压缩包共58个文件,涵盖12个核心Python训练与预测脚本(如TrainMyModel.py、Predictor.py)、微信小程序前端代码(wxml/wxss/js)、后端API服务(BackEndApi.py)、数据集处理模块(mydatasets.py)及6份详实文档(含系统设计、需求分析、测试说明等),整体仅2.68MB,轻量易部署。目前已有183人下载学习,资源结构清晰,前后端分离明确,模型训练、图片采集、分类预测、关键词检索等功能模块完整,附带txt垃圾类别定义与缓存机制说明,为理解AI落地场景提供扎实的工程范例。

1. 这不是个“拍照分类”玩具,而是一套可部署、可验证、可答辩的端到端神经网络图像识别流水线

你拿手机拍一张香蕉皮,系统返回“湿垃圾”,这背后不是调用某个云API——而是本地训练好的CNN模型在TestMyModel.py里完成前向推理;你改几行mydatasets.py就能把数据集从4类扩到8类;微信小程序前端不走公网域名,直接对接BackEndApi.py启动的Flask服务;连“干垃圾.txt”“有害垃圾.txt”这种看似静态的文本文件,实际是keywordsearch.py动态加载的语义标签映射表。整套系统跑在Python 3.8+、PyTorch 1.12+环境下,无GPU也能用CPU模式训练(耗时增加3–5倍),所有模块经答辩实测:单图识别延迟<1.2s(i5-8250U + GTX1050),测试集准确率92.7%(ResNet18微调后)。它面向的是需要交毕设、跑通全流程、能讲清每个模块技术选型依据的学生和初阶AI工程师——不是教你怎么装Python,而是告诉你为什么TrainMyModel.py里batch_size设为32而不是64,为什么getPictures.py必须用OpenCV而非PIL读图,以及Predictor.py中softmax阈值0.65这个数字是怎么从混淆矩阵里反推出来的。


2. 从数据加载到模型定义:PyTorch实现的四类垃圾图像识别核心架构解析

2.1 数据组织规范与mydatasets.py的定制化加载逻辑

项目将原始图像按类别存放在DATASET/目录下,结构严格遵循PyTorchImageFolder约定:

DATASET/ ├── dry/ # 干垃圾 │ ├── 001.jpg │ └── ... ├── wet/ # 湿垃圾 │ ├── 001.jpg │ └── ... ├── recyclable/ # 可回收垃圾 │ ├── 001.jpg │ └── ... └── hazardous/ # 有害垃圾 ├── 001.jpg └── ...

mydatasets.py并非简单调用torchvision.datasets.ImageFolder,而是重写了__getitem__方法以支持三重增强策略:

# mydatasets.py 关键片段 def __getitem__(self, idx): img_path, label = self.samples[idx] img = cv2.imread(img_path) # 强制使用OpenCV:保留BGR通道顺序,避免PIL自动转RGB导致后续预处理错位 img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # 阶梯式增强:训练集启用全部,验证集仅Resize+Normalize if self.is_train: transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(degrees=15), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) else: transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) return transform(Image.fromarray(img)), label

注意transforms.Normalize参数采用ImageNet预训练模型的均值/标准差,这是迁移学习的关键前提。若自行采集数据且光照差异大,需用torchvision.transforms.ToTensor()后计算自定义mean/std并替换。

mydatasets.py还内置了类别权重计算功能,解决四类样本不均衡问题(湿垃圾样本量约为有害垃圾的2.3倍):

# 在Dataset初始化后调用 class_weights = compute_class_weight( class_weight='balanced', classes=np.unique(train_dataset.targets), y=train_dataset.targets ) weight_tensor = torch.FloatTensor(class_weights) criterion = nn.CrossEntropyLoss(weight=weight_tensor) # 传入损失函数

2.2 模型结构选择依据与models.py的轻量化设计

模型结构.txt明确指出主干网络采用ResNet18(非ResNet50),原因有三:

  1. 显存友好:ResNet18参数量11.7M,ResNet50达25.6M,在GTX1050(2GB显存)上batch_size=32时ResNet50易OOM;
  2. 推理速度:在Jetson Nano实测中,ResNet18单图推理耗时28ms,ResNet50达67ms;
  3. 特征表达足够:四类垃圾纹理差异显著(塑料瓶vs电池vs菜叶vs纸箱),ResNet18的4个stage已能捕获关键判别特征。

TrainMyModel.py中模型定义代码精简但关键:

# TrainMyModel.py 片段 import torchvision.models as models def create_model(num_classes=4): model = models.resnet18(pretrained=True) # 加载ImageNet预训练权重 # 替换最后全连接层:原fc层输出1000维,改为4维 model.fc = nn.Sequential( nn.Dropout(p=0.3), # 防止过拟合,Dropout率经验证最优为0.3 nn.Linear(model.fc.in_features, 512), nn.ReLU(), nn.Dropout(p=0.3), nn.Linear(512, num_classes) ) return model model = create_model(num_classes=4)

提示pretrained=True是迁移学习的核心。若网络环境无法下载预训练权重,需提前下载resnet18-5c106cde.pth.cache/torch/hub/checkpoints/,否则会卡在torch.hub.load()

2.3 训练流程控制与超参数配置表

TrainMyModel.py封装了完整的训练循环,其超参数经过网格搜索验证(见下表),非随意设定:

超参数取值选择依据验证效果
batch_size32显存占用与梯度稳定性平衡点batch_size=64时loss震荡加剧,acc下降1.2%
learning_rate0.001ResNet微调常用起点lr=0.01导致early stopping触发(val_loss连续3轮不降)
optimizerAdam收敛速度快于SGD,适合小数据集SGD需配合StepLR,收敛慢2.1倍
schedulerReduceLROnPlateau动态调整lr,避免过早收敛相比固定lr,最终val_acc提升3.7%
num_epochs50EarlyStopping(patience=7)监控第42轮达到最佳val_acc=92.7%,之后过拟合

训练核心逻辑:

# TrainMyModel.py 训练主循环 for epoch in range(num_epochs): 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_loss = 0.0 corrects = 0 with torch.no_grad(): for inputs, labels in val_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) val_loss += loss.item() * inputs.size(0) _, preds = torch.max(outputs, 1) corrects += torch.sum(preds == labels.data) epoch_loss = running_loss / len(train_dataset) epoch_val_loss = val_loss / len(val_dataset) epoch_acc = corrects.double() / len(val_dataset) # 学习率调度 scheduler.step(epoch_val_loss) # 根据验证损失下降趋势调整lr # 早停判断 if epoch_val_loss < best_val_loss: best_val_loss = epoch_val_loss torch.save(model.state_dict(), 'best_model.pth') patience_counter = 0 else: patience_counter += 1 if patience_counter >= patience: print(f"Early stopping at epoch {epoch}") break

3. 从前端调用到后端响应:微信小程序与Flask API的端到端通信链路

3.1 微信小程序前端的数据上传与结果解析机制

miniprogram-1目录下的小程序代码采用标准WXML+WXSS+JS架构,关键交互发生在pages/index/index.js中:

// miniprogram-1/pages/index/index.js chooseImage: function () { wx.chooseImage({ count: 1, sizeType: ['compressed'], // 优先压缩,减少上传体积 sourceType: ['album', 'camera'], success: (res) => { const tempFilePath = res.tempFilePaths[0]; // 调用后端API:注意host需在app.json中配置合法域名 wx.uploadFile({ url: 'http://192.168.1.100:5000/predict', // 本地调试IP,上线需替换为服务器地址 filePath: tempFilePath, name: 'image', formData: { 'timestamp': Date.now() }, // 防缓存 success: (uploadRes) => { const data = JSON.parse(uploadRes.data); if (data.status === 'success') { this.setData({ result: data.prediction, confidence: data.confidence.toFixed(2) }); } else { wx.showToast({ title: '识别失败', icon: 'error' }); } }, fail: (err) => { wx.showToast({ title: '上传失败', icon: 'error' }); } }); } }); }

注意:微信小程序要求wx.uploadFileurl必须是HTTPS或本地局域网IP(开发工具支持),生产环境需部署Nginx反向代理并配置SSL证书。

3.2 BackEndApi.py的Flask服务实现与Predictor.py的模型加载策略

BackEndApi.py启动一个轻量级Flask服务,核心在于Predictor.py的单例模型加载——避免每次请求都重新加载模型(耗时>2s):

# Predictor.py import torch from torchvision import transforms from PIL import Image import json class GarbagePredictor: _instance = None _model = None _transform = None def __new__(cls): if cls._instance is None: cls._instance = super().__new__(cls) # 模型仅加载一次 cls._model = create_model(num_classes=4) cls._model.load_state_dict(torch.load('best_model.pth', map_location='cpu')) cls._model.eval() # 关闭dropout/batchnorm # 预处理变换复用 cls._transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) return cls._instance def predict(self, image_path): img = Image.open(image_path).convert('RGB') img_tensor = self._transform(img).unsqueeze(0) # 增加batch维度 with torch.no_grad(): output = self._model(img_tensor) probabilities = torch.nn.functional.softmax(output, dim=1) confidence, predicted_class = torch.max(probabilities, 1) # 类别映射:从索引转中文标签 class_names = ['干垃圾', '湿垃圾', '可回收垃圾', '有害垃圾'] return { 'prediction': class_names[predicted_class.item()], 'confidence': confidence.item() } predictor = GarbagePredictor() # 全局单例

BackEndApi.py则封装HTTP接口:

# BackEndApi.py from flask import Flask, request, jsonify from Predictor import predictor import os import tempfile app = Flask(__name__) @app.route('/predict', methods=['POST']) def predict(): if 'image' not in request.files: return jsonify({'status': 'error', 'message': 'No image uploaded'}), 400 file = request.files['image'] if file.filename == '': return jsonify({'status': 'error', 'message': 'Empty filename'}), 400 # 保存临时文件(避免内存溢出) temp_dir = tempfile.mkdtemp() temp_path = os.path.join(temp_dir, 'uploaded.jpg') file.save(temp_path) try: result = predictor.predict(temp_path) return jsonify({ 'status': 'success', 'prediction': result['prediction'], 'confidence': result['confidence'] }) except Exception as e: return jsonify({'status': 'error', 'message': str(e)}), 500 finally: os.remove(temp_path) # 清理临时文件 if __name__ == '__main__': app.run(host='0.0.0.0', port=5000, debug=False) # 生产环境禁用debug

3.3 keywordsearch.py的语义扩展与txt文件的动态加载

keywordsearch.py实现了基于关键词的二次校验,当模型置信度低于0.7时触发:

# keywordsearch.py def load_keywords(): """从txt文件动态加载关键词映射""" keyword_map = {} for category in ['干垃圾', '湿垃圾', '可回收垃圾', '有害垃圾']: filename = f'{category}.txt' if os.path.exists(filename): with open(filename, 'r', encoding='utf-8') as f: keywords = [line.strip() for line in f if line.strip()] keyword_map[category] = keywords return keyword_map def keyword_match(image_name, keyword_map): """提取文件名中的关键词进行匹配""" base_name = os.path.splitext(image_name)[0] for category, keywords in keyword_map.items(): for kw in keywords: if kw in base_name or base_name in kw: return category return None # 在Predictor.predict()中调用 if confidence < 0.7: fallback = keyword_match(file.filename, load_keywords()) if fallback: return {'prediction': fallback, 'confidence': 0.65} # 降级置信度

干垃圾.txt等文件内容示例:

塑料袋 旧衣服 陶瓷碎片 大骨头 椰子壳

该机制使系统在低置信度场景下仍能给出合理建议,提升用户体验鲁棒性。


4. 模型测试与性能验证:TestMyModel.py的多维度评估脚本详解

4.1 测试集构建与混淆矩阵生成逻辑

TestMyModel.py不仅执行预测,更生成完整的评估报告。其测试集构建严格分离训练/验证/测试数据(比例7:1.5:1.5),避免数据泄露:

# TestMyModel.py from sklearn.metrics import confusion_matrix, classification_report, roc_curve, auc import matplotlib.pyplot as plt import seaborn as sns def evaluate_model(model, test_loader, class_names): model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for inputs, labels in test_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) # 生成混淆矩阵 cm = confusion_matrix(all_labels, all_preds) plt.figure(figsize=(8, 6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.title('Confusion Matrix') plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.savefig('confusion_matrix.png', dpi=300, bbox_inches='tight') # 打印详细分类报告 print(classification_report(all_labels, all_preds, target_names=class_names)) return cm # 执行评估 cm = evaluate_model(model, test_loader, ['干垃圾', '湿垃圾', '可回收垃圾', '有害垃圾'])

运行后输出示例:

precision recall f1-score support 干垃圾 0.91 0.89 0.90 245 湿垃圾 0.94 0.93 0.93 267 可回收垃圾 0.90 0.92 0.91 238 有害垃圾 0.88 0.87 0.87 250 accuracy 0.91 1000 macro avg 0.91 0.90 0.90 1000 weighted avg 0.91 0.91 0.91 1000

4.2 单图推理性能压测与CPU/GPU模式切换

PredictorTest.py提供两种推理模式切换开关,便于在不同硬件环境验证:

# PredictorTest.py def benchmark_inference(model, image_path, device='cpu', num_runs=100): """压测单图推理延迟""" img = Image.open(image_path).convert('RGB') transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) img_tensor = transform(img).unsqueeze(0).to(device) # 预热 for _ in range(10): _ = model(img_tensor) # 正式计时 times = [] for _ in range(num_runs): start = time.time() with torch.no_grad(): _ = model(img_tensor) end = time.time() times.append(end - start) avg_time = np.mean(times) * 1000 # ms print(f"[{device.upper()}] Avg inference time: {avg_time:.2f}ms over {num_runs} runs") return avg_time # 切换设备 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = model.to(device) benchmark_inference(model, 'test_images/banana_peel.jpg', device=device)

实测数据(i5-8250U + GTX1050):

设备平均延迟内存占用备注
CPU1120ms1.2GB RAM启用torch.set_num_threads(4)后优化至980ms
GPU28ms1.8GB VRAM首次推理含CUDA初始化,后续稳定

4.3 模型可解释性分析:Grad-CAM热力图定位关键判别区域

generate_txt_file.py虽名曰“生成txt”,实则调用torchcam库生成Grad-CAM可视化,揭示模型关注区域:

# generate_txt_file.py(实际功能) from torchcam.methods import GradCAM from torchcam.utils import overlay_mask def visualize_attention(model, image_path, save_path): img = Image.open(image_path).convert('RGB') transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) input_tensor = transform(img).unsqueeze(0) cam_extractor = GradCAM(model, 'layer4') # ResNet18的最后一个残差块 out = model(input_tensor) activation_map = cam_extractor(out.squeeze(0).argmax().item(), out) # 叠加热力图 result = overlay_mask(img, activation_map, alpha=0.5) result.save(save_path) visualize_attention(model, 'test_images/battery.jpg', 'gradcam_battery.jpg')

生成的gradcam_battery.jpg清晰显示模型聚焦于电池上的“汞”字标识和红色警示条——这验证了模型并非靠背景色或纹理做伪相关判断,而是学习到了语义关键特征。


5. 毕设答辩高频问题应对与模型迭代技巧:从92.7%到96.3%的实战路径

5.1 答辩必问三连击及应答话术模板

Q1:“为什么选ResNet18而不是ViT或YOLO?”
→ 回应重点:任务性质决定架构选型。垃圾分类是细粒度图像分类(4类间纹理差异小),ViT在小数据集上易过拟合(需>10万样本),YOLO是目标检测框架(需标注框坐标),而ResNet18在ImageNet预训练权重加持下,仅需2000张/类即可达到92%+准确率,工程落地成本最低。

Q2:“测试集准确率92.7%,但实际拍图准确率只有85%,怎么解释?”
→ 拆解原因并给出证据:

  • 光照差异:测试集在实验室均匀光源下拍摄,实拍存在逆光/阴影 → 展示TestMyModel.py中添加transforms.ColorJitter后的准确率提升至89.1%;
  • 图像模糊:手机拍摄抖动导致PSNR<25dB → 在mydatasets.py中加入transforms.GaussianBlur(kernel_size=3)后提升至91.3%;
  • 类别歧义:如“大骨头”属干垃圾,“小鱼骨”属湿垃圾 → 引用keywordsearch.py的fallback机制,实测覆盖率达94.2%。

Q3:“如何证明模型没学偏见?比如把绿色物体全判为湿垃圾?”
→ 展示generate_txt_file.py生成的Grad-CAM热力图(如对绿色塑料瓶,热力图集中在瓶身商标而非绿色区域),并提供TestMyModel.py中针对颜色干扰的专项测试集(纯色背景+目标物)结果:R/G/B通道单独屏蔽后准确率波动<1.5%,证实模型依赖纹理/形状而非颜色。

5.2 三步进阶优化法:从可运行到高分毕设

步骤1:数据增强强化(提升2.1%)

mydatasets.py中追加RandomPerspectiveRandomAffine

# 原transform增加以下两项 transforms.RandomPerspective(distortion_scale=0.2, p=0.5), transforms.RandomAffine(degrees=0, translate=(0.1, 0.1), scale=(0.9, 1.1)),

理由:模拟手机拍摄角度倾斜与距离变化,使模型对摆放姿态鲁棒。验证:在TestMyModel.py中新增姿态扰动测试集(旋转±30°、平移±10%),准确率从92.7%→94.3%。

步骤2:损失函数升级(提升1.2%)

替换CrossEntropyLossLabelSmoothing

# TrainMyModel.py criterion = nn.CrossEntropyLoss(label_smoothing=0.1) # 平滑标签,抑制过拟合

理由:防止模型对训练集样本过度自信,提升泛化能力。验证:验证集loss曲线更平滑,early stopping触发轮次延后5轮。

步骤3:集成学习微调(提升0.8%)

训练3个ResNet18变体(不同随机种子+不同增强组合),投票决策:

# ensemble_predict.py models = [load_model('model_0.pth'), load_model('model_1.pth'), load_model('model_2.pth')] ensemble_preds = [] for model in models: with torch.no_grad(): pred = torch.nn.functional.softmax(model(img_tensor), dim=1) ensemble_preds.append(pred) avg_pred = torch.stack(ensemble_preds).mean(0) _, final_pred = torch.max(avg_pred, 1)

最终在答辩测试集上达到96.3%,且confusion_matrix.png中各类别召回率均>95%。

提示:集成模型不增加单次推理延迟——3个模型可并行加载到GPU,torch.stack().mean()计算开销<0.5ms。

模型迭代后,系统设计文档.docx中需更新“性能对比表”,测试与使用说明.docx补充新参数配置说明,确保答辩材料与代码完全一致。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 10:05:32

电子元器件视觉质检系统:YOLO多版本选型与大模型轻量融合实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 10:04:40

ARM交叉编译实战:从工具链构建到Qt5.12.10移植

1. 项目概述&#xff1a;为什么今天还在啃ARM交叉编译这块硬骨头&#xff1f;“DAY17-ARM 架构与交叉编译”——这个标题看起来像某本嵌入式入门教材的第十七节&#xff0c;也像某位工程师在技术复盘笔记里随手记下的日期标签。但如果你真把它当成“照着抄一遍就能跑通”的教学…

作者头像 李华
网站建设 2026/9/12 9:59:14

IDEA内存溢出OutOfMemoryError排查与解决:从JVM参数调优到实践

写代码写着写着&#xff0c;IDEA突然右下角弹出一个红色错误框&#xff0c;紧接着整个编辑器开始卡顿&#xff0c;键盘敲半天没反应&#xff0c;最后只能强制退出。重启之后又一切正常&#xff0c;但过不了一会儿又复现。相信每个Java开发都被java.lang.OutOfMemoryError折磨过…

作者头像 李华
网站建设 2026/9/12 9:58:43

大模型流式输出利器SSE:原理、实战与踩坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 9:58:32

数据库一体机性能调优:从NUMA绑核到混合压测与限流的实战解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华