垃圾分类政策落地之后,可回收物、厨余垃圾、有害垃圾、其他垃圾这四个分类让不少人犯了难。一个能拍照识别垃圾类别并给出投放建议的智能Agent,就成了非常典型的AI落地场景。这篇文章就把这个项目完整拆解一遍:从数据准备、模型训练,到把模型包装成一个可以对话交互的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做模型加速以提升并发能力,或者做一个微信小程序作为前端。整个项目跑通之后,你会对图像分类从数据到服务的完整链路有一个非常扎实的理解。