拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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.2si5-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) # 阶梯式增强训练集启用全部验证集仅ResizeNormalize if self.is_train: transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.2, contrast0.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_weightbalanced, classesnp.unique(train_dataset.targets), ytrain_dataset.targets ) weight_tensor torch.FloatTensor(class_weights) criterion nn.CrossEntropyLoss(weightweight_tensor) # 传入损失函数2.2 模型结构选择依据与models.py的轻量化设计模型结构.txt明确指出主干网络采用ResNet18非ResNet50原因有三显存友好ResNet18参数量11.7MResNet50达25.6M在GTX10502GB显存上batch_size32时ResNet50易OOM推理速度在Jetson Nano实测中ResNet18单图推理耗时28msResNet50达67ms特征表达足够四类垃圾纹理差异显著塑料瓶vs电池vs菜叶vs纸箱ResNet18的4个stage已能捕获关键判别特征。TrainMyModel.py中模型定义代码精简但关键# TrainMyModel.py 片段 import torchvision.models as models def create_model(num_classes4): model models.resnet18(pretrainedTrue) # 加载ImageNet预训练权重 # 替换最后全连接层原fc层输出1000维改为4维 model.fc nn.Sequential( nn.Dropout(p0.3), # 防止过拟合Dropout率经验证最优为0.3 nn.Linear(model.fc.in_features, 512), nn.ReLU(), nn.Dropout(p0.3), nn.Linear(512, num_classes) ) return model model create_model(num_classes4)提示pretrainedTrue是迁移学习的核心。若网络环境无法下载预训练权重需提前下载resnet18-5c106cde.pth至.cache/torch/hub/checkpoints/否则会卡在torch.hub.load()。2.3 训练流程控制与超参数配置表TrainMyModel.py封装了完整的训练循环其超参数经过网格搜索验证见下表非随意设定超参数取值选择依据验证效果batch_size32显存占用与梯度稳定性平衡点batch_size64时loss震荡加剧acc下降1.2%learning_rate0.001ResNet微调常用起点lr0.01导致early stopping触发val_loss连续3轮不降optimizerAdam收敛速度快于SGD适合小数据集SGD需配合StepLR收敛慢2.1倍schedulerReduceLROnPlateau动态调整lr避免过早收敛相比固定lr最终val_acc提升3.7%num_epochs50EarlyStopping(patience7)监控第42轮达到最佳val_acc92.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(fEarly stopping at epoch {epoch}) break3. 从前端调用到后端响应微信小程序与Flask API的端到端通信链路3.1 微信小程序前端的数据上传与结果解析机制miniprogram-1目录下的小程序代码采用标准WXMLWXSSJS架构关键交互发生在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.uploadFile的url必须是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_classes4) cls._model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) 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, dim1) 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(host0.0.0.0, port5000, debugFalse) # 生产环境禁用debug3.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, encodingutf-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, annotTrue, fmtd, cmapBlues, xticklabelsclass_names, yticklabelsclass_names) plt.title(Confusion Matrix) plt.ylabel(True Label) plt.xlabel(Predicted Label) plt.savefig(confusion_matrix.png, dpi300, bbox_inchestight) # 打印详细分类报告 print(classification_report(all_labels, all_preds, target_namesclass_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 10004.2 单图推理性能压测与CPU/GPU模式切换PredictorTest.py提供两种推理模式切换开关便于在不同硬件环境验证# PredictorTest.py def benchmark_inference(model, image_path, devicecpu, num_runs100): 压测单图推理延迟 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, devicedevice)实测数据i5-8250U GTX1050设备平均延迟内存占用备注CPU1120ms1.2GB RAM启用torch.set_num_threads(4)后优化至980msGPU28ms1.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, alpha0.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%图像模糊手机拍摄抖动导致PSNR25dB → 在mydatasets.py中加入transforms.GaussianBlur(kernel_size3)后提升至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中追加RandomPerspective和RandomAffine# 原transform增加以下两项 transforms.RandomPerspective(distortion_scale0.2, p0.5), transforms.RandomAffine(degrees0, translate(0.1, 0.1), scale(0.9, 1.1)),理由模拟手机拍摄角度倾斜与距离变化使模型对摆放姿态鲁棒。验证在TestMyModel.py中新增姿态扰动测试集旋转±30°、平移±10%准确率从92.7%→94.3%。步骤2损失函数升级提升1.2%替换CrossEntropyLoss为LabelSmoothing# TrainMyModel.py criterion nn.CrossEntropyLoss(label_smoothing0.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), dim1) 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个模型可并行加载到GPUtorch.stack().mean()计算开销0.5ms。模型迭代后系统设计文档.docx中需更新“性能对比表”测试与使用说明.docx补充新参数配置说明确保答辩材料与代码完全一致。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门