导读:本期聚焦于会飞的猪创作的《如何用Python+深度学习构建一个垃圾分类识别Agent?从数据集到部署完整实战》,敬请观看详情。垃圾桶前拿着外卖盒犹豫该扔哪个口?这正是垃圾分类识别Agent要解决的问题。本文以一个完整的实战案例为主线,讲解如何利用卷积神经网络对垃圾图片进行自动分类,涵盖数据集获取与清洗、模型训练与调优、以及将模型封装成可交互Agent的全过程。文中会给出基于PyTorch的完整训练代码、模型评估方法,并介绍如何用Flask把训练好的模型做成API服务,再接入对话逻辑让它真正具备回答能力。整个过程不依赖昂贵的GPU集群,普通电脑即可跑通,适合想在图像分类和智能体开发方向练手的开发者参考。

垃圾分类政策落地之后,可回收物、厨余垃圾、有害垃圾、其他垃圾这四个分类让不少人犯了难。一个能拍照识别垃圾类别并给出投放建议的智能Agent,就成了非常典型的AI落地场景。这篇文章就把这个项目完整拆解一遍:从数据准备、模型训练,到把模型包装成一个可以对话交互的Agent服务,每一步都给出可运行的代码。

如何用Python+深度学习构建一个垃圾分类识别Agent?从数据集到部署完整实战

一、项目整体设计思路

先明确这个Agent要做的事情:用户上传一张垃圾图片(或者描述一个物品名称),Agent识别出它属于哪个垃圾类别,并返回投放建议。整个系统可以拆成三层:最底层是图像分类模型,负责判断图片内容;中间层是推理服务,把模型封装成API;最上层是交互层,负责接收用户输入并组织回复话术。

为什么要分层?因为模型训练和服务部署的生命周期完全不同。模型可能每隔几个月要重训一次,而服务要长期在线。如果把训练代码和推理代码揉在一起,后期维护会非常痛苦。分层之后,模型以文件形式存盘,服务启动时加载,互不干扰。

在模型选型上,考虑到垃圾图片识别是典型的多分类任务,且类别数不多(通常4到40类不等),我们采用迁移学习方案:拿在ImageNet上预训练好的ResNet18做骨干网络,替换最后的全连接层。这样即使只有几千张训练图片,也能达到不错的准确率,普通CPU训练几个小时就能收敛。

二、数据集准备与预处理

公开可用的垃圾图片数据集不少,比如华为云垃圾分类竞赛数据集、TrashNet等。TrashNet包含约2500张图片,分为玻璃、纸、硬纸板、塑料、金属、其他六大类,非常适合入门。可以将这六类再映射到国标的四大分类里,比如玻璃、金属、纸归入可回收物。

拿到数据后第一件事是做数据清洗。真实数据集里经常混有模糊图、重复图、标注错误的图,这些脏数据对模型伤害很大。简单有效的做法是先随机抽样浏览几百张,把明显有问题的挑出来,再借助感知哈希去重。

数据增强是提升泛化能力的关键手段。垃圾图片的拍摄角度、光照条件千差万别,训练时加入随机裁剪、水平翻转、色彩抖动,可以显著降低过拟合风险。下面是完整的Dataset和数据加载代码:

import torch
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

# 训练集数据增强:随机裁剪、翻转、色彩抖动
train_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.RandomCrop(224),
    transforms.RandomHorizontalFlip(),
    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])
])

# 验证集只做基础缩放,保证评估结果稳定
val_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])
])

train_ds = datasets.ImageFolder('data/train', transform=train_transform)
val_ds = datasets.ImageFolder('data/val', transform=val_transform)

train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4)
val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4)

print(train_ds.classes)  # 查看类别顺序,后面推理时要保持一致

注意ImageFolder要求每个类别的图片放在以类别名命名的子目录里,目录名的字母序就是类别索引。这个类别顺序一定要记下来,部署时如果顺序对不上,预测结果会张冠李戴,这是新手最容易踩的坑之一。

三、模型训练与评估

迁移学习的核心思路是:冻结(或以较小学习率微调)预训练的卷积层,只重点训练新加的分类头。ResNet18在ImageNet上学到的边缘、纹理、形状特征,对垃圾图片同样适用,我们只需要让模型学会把这些特征映射到垃圾类别上。

import torch.nn as nn
from torchvision import models

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# 加载预训练模型并替换分类头
model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
model.fc = nn.Linear(model.fc.in_features, len(train_ds.classes))
model = model.to(device)

criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)

def train_one_epoch():
    model.train()
    total_loss = 0
    for x, y in train_loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()
        loss = criterion(model(x), y)
        loss.backward()
        optimizer.step()
        total_loss += loss.item() * x.size(0)
    return total_loss / len(train_ds)

@torch.no_grad()
def evaluate():
    model.eval()
    correct = 0
    for x, y in val_loader:
        x, y = x.to(device), y.to(device)
        pred = model(x).argmax(dim=1)
        correct += (pred == y).sum().item()
    return correct / len(val_ds)

best_acc = 0
for epoch in range(15):
    loss = train_one_epoch()
    acc = evaluate()
    print(f'epoch {epoch}, loss={loss:.4f}, val_acc={acc:.4f}')
    if acc > best_acc:
        best_acc = acc
        torch.save(model.state_dict(), 'garbage_model.pth')

训练几个epoch后,验证集准确率通常能到85%以上。如果发现训练准确率很高而验证准确率上不去,说明过拟合了,可以加大数据增强力度、加dropout,或者冻结更多卷积层。反过来如果两者都低,可能是数据量不够或学习率设置不合理。

除了准确率,建议再看一下混淆矩阵。垃圾类别中塑料和金属、纸和硬纸板很容易互相混淆,混淆矩阵能直观暴露模型的薄弱环节,针对性地补充这两类的训练样本往往比盲目加数据更有效。

四、把模型封装成可交互的Agent

模型只是Agent的大脑,用户真正接触的是交互界面。我们用Flask搭一个简单的API服务,接收图片后返回分类结果和投放建议,同时支持文字提问模式,让它像个真正的助手。

from flask import Flask, request, jsonify
from PIL import Image
import io, torch

app = Flask(__name__)

# 类别到投放建议的映射
ADVICE = {
    'glass': '可回收物:请清空内容物后投入可回收物桶',
    'paper': '可回收物:请展平后投入可回收物桶',
    'plastic': '可回收物:请清洗晾干后投入可回收物桶',
    'metal': '可回收物:请压扁后投入可回收物桶',
    'cardboard': '可回收物:请拆除胶带后投入可回收物桶',
    'trash': '其他垃圾:请投入其他垃圾桶'
}

model.load_state_dict(torch.load('garbage_model.pth', map_location=device))
model.eval()

@app.route('/recognize', methods=['POST'])
def recognize():
    file = request.files.get('image')
    if not file:
        return jsonify({'error': '请上传图片'}), 400
    img = Image.open(io.BytesIO(file.read())).convert('RGB')
    x = val_transform(img).unsqueeze(0).to(device)
    with torch.no_grad():
        probs = torch.softmax(model(x), dim=1)[0]
    idx = probs.argmax().item()
    cls = train_ds.classes[idx]
    return jsonify({
        'category': cls,
        'confidence': round(probs[idx].item(), 4),
        'advice': ADVICE.get(cls, '暂无建议')
    })

if __name__ == '__main__':
    app.run(host='0.0.0.0', port=5000)

这里有个细节值得强调:返回结果一定要带上置信度。当置信度低于某个阈值(比如0.6)时,Agent应该坦白说识别不确定,建议用户换个角度重拍,而不是硬给一个可能错误的答案。诚实的Agent比看起来无所不能的Agent更值得信任。

进一步提升的方向也有很多:接入大语言模型处理文字咨询类的提问(比如“过期药品是什么垃圾”),用ONNX或TorchScript做模型加速以提升并发能力,或者做一个微信小程序作为前端。整个项目跑通之后,你会对图像分类从数据到服务的完整链路有一个非常扎实的理解。

垃圾分类识别深度学习图像分类PythonAgent修改时间:2026-09-11 11:40:44

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/0911/54644.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。