简介:面向垃圾分类入门与课题实践的图像识别资源,覆盖硬纸板、纸、塑料瓶、玻璃瓶、铜制品与不可回收垃圾六类常见样本,适合学习卷积神经网络训练流程或构建简易分类系统的学生与开发者。资源共7个文件,以5个Python脚本为主,分别承担模型训练、类别预测、图像预处理与结果可视化等环节,另附1个已训练好的h5权重文件和1个数据集压缩包,整体约161.27MB。代码基于TensorFlow、Keras与OpenCV,训练脚本可按自定义数据扩展类别,预测脚本只需输入图片路径即可输出分类结果,使用门槛较低。虽然原始训练数据量庞大未随包上传,但权重文件可支撑开箱即用的预测演示,便于快速验证模型效果。资源结构清晰,目录层级简洁,适合作为课程设计、毕业设计或小型环保应用开发的参考模板。目前已有30634人学习,具备较高的参考热度与实用价值。
1. 垃圾分类数据集与代码:看起来是个分类题,其实是一整条数据流水线
“垃圾分类数据集与代码”听起来像是一个标准的入门项目:下载一批图片,跑一个分类网络,拿到准确率就收工。但真上手的人大多会卡在同一步——验证集准确率死活上不去75%,换模型、加训练轮数都没用。问题通常不在模型,而在数据本身:重复图混进训练集和验证集、类别标签口径不一致、某一个类样本少得可怜,这些坑靠调参绕不过去。
这篇文章想讲的不是“从哪下载数据集然后一键跑通”,而是一条从原始图片到可复现训练代码的完整路径:任务选型与数据来源、脏数据清洗与划分、训练脚本与 checkpoint 组织,以及分类方案跑通之后怎么扩展成检测模型。适合赶毕设、做比赛,或者第一次用深度学习做图像分类的开发者。
2. 先定任务再定数据:分类方案与数据集来源怎么搭
想动手做垃圾分类,第一件事不是找数据集下载,而是先想清楚任务边界。垃圾分类在 CV 里通常有两个解法:图像分类和目标检测,两者对数据的要求完全不同。这步选错,后面清洗、标注、训练的功夫都白费。
2.1 四分类、多分类还是检测:任务粒度决定数据集长什么样
分类任务回答的是“这张图里最主要的东西属于哪一类”,检测任务回答的是“图里每个垃圾在哪、分别是什么”。如果你做的是智能垃圾桶投递口,摄像头对准投递口,一次拍摄只出现一件垃圾,分类就够用;如果要做手机随手拍的场景,画面里往往同时出现饮料瓶、纸巾和果皮,你就需要检测模型,而检测模型的数据集标签就不是“文件夹名”而是“每个目标的框坐标”。
这些细节在《垃圾分类数据集及代码》标题里没写出来,但数据集的形态完全由它决定。常见做法是这样权衡:
| 对比维度 | 图像分类 | 目标检测 |
|---|---|---|
| 输出结果 | 整张图的类别 | 每个目标的类别和位置框 |
| 标签形式 | 文件夹名 / CSV 一列 | XML / TXT 框坐标 |
| 数据量需求 | 每类几百张可起步 | 每类至少上千个实例框 |
| 适合场景 | 单目标、固定机位、识别速度快 | 多目标、场景杂、要联动机械臂或计数 |
如果你在题目里看到“定位”“抓取”“多个目标”,那就老老实实走检测路线;如果只是判断“这一袋属于什么垃圾”,分类是性价比更高的选择。iris 那种几十条样本的玩具数据集和图片分类根本不是一回事,千万别拿那个经验来套真实图片。
2.2 公开数据集与自建数据:从哪找、找完怎么验货
确定做分类后,数据来源一般两条路:公开数据集和自建数据。公开数据集的好处是省事,坏处是类别定义往往和你手上的标注规范对不上。比如“可回收垃圾”在不同数据包里可能拆成“塑料瓶/纸箱/易拉罐”,也可能合并成一个大类;有的数据包里“其他垃圾”一栏放的基本都是黑色的垃圾袋特写。下载完第一件事不是解压后直接写训练脚本,而是先翻一遍图片,看看每类的实际内容是否和目录名一致。
从 Hugging Face 上的一些数据集卡片或者 GitHub 的 release 附件里找公开数据,都是常见做法。下载前我会先看三个东西:类别列表、每类图片数、以及是否已经划分好 train/val。很多包装成“完整数据集”的包,其实只是一堆没有划分的原始图片,划分工作还得自己来。最后一步是校验文件总数和 README 是否对得上,防止有人传漏了一部分。CUB 这种鸟类分类数据集的维护者会强调图片边界和标签来源,垃圾数据集的维护者未必做得到,所以验货环节不能省。
自建数据这块,常见做法是手机拍摄加网络搜索组合。手机拍要覆盖不同光线、角度和距离——实际部署时没人会像拍照一样把垃圾摆正。网络搜索来的图要特别注意:缩略图、水印图、表情包、漫画插图经常混进来,这些都属于脏数据,后面清洗阶段会花掉不少时间。
2.3 数据集文件结构与标签规范:classes.txt 的顺序就是模型输出的顺序
我一般会把数据组织成这个结构:
dataset/ ├── images/ │ ├── train/ │ │ ├── kitchen_waste/ │ │ │ ├── 000001.jpg │ │ │ ├── 000002.jpg │ │ │ └── ... │ │ ├── recyclable/ │ │ └── ... │ ├── val/ │ │ ├── kitchen_waste/ │ │ └── ... │ └── test/ │ └── ... ├── labels/ │ └── classes.txt └── meta/ └── dataset_config.json这个文件结构有两个好处。第一,torchvision 的 ImageFolder 可以直接按照目录名读取标签,省掉自己写 CSV 映射的步骤;第二,人工检查时顺着目录一层层看,哪类图不对一眼就能发现。classes.txt 里的每行就是一个类别名,顺序不能乱,因为后面模型全连接层的输出就按这个顺序对齐。dataset_config.json 用来记录图片尺寸、归一化的 mean/std、类别数量和一次全量清洗的时间,这样哪怕三个月后回来续训,也能知道当前数据是哪个版本。
3. 把原始图片做成能喂给模型的数据集:清洗、去重与划分脚本
不管数据来自下载还是自拍,先过一遍清洗脚本再谈训练。这一步最容易被跳过,也最影响结果。下面几段脚本解决的是“原始图片目录如何变成可训练数据集”的问题。
3.1 清洗脚本:坏图、重复图和网络缩略图怎么扫地出门
以下脚本遍历 images 目录,做三件事:打开失败视为坏图、按 MD5 找完全重复的图、统计每类数量。
import hashlib from pathlib import Path from PIL import Image root = Path("dataset/images") suspected = Path("dataset/suspected") # 隔离区,不直接删 suspected.mkdir(exist_ok=True) seen_hashes = {} bad_count = 0 for img_path in sorted(root.rglob("*.jpg")): # png/jpeg 同理,可按需补全后缀 try: with Image.open(img_path) as im: im = im.convert("RGB") # 强制转 RGB,防灰度图混入 except Exception as exc: print(f"[坏图] {img_path}: {exc}") new_path = suspected / f"bad_{bad_count}_{img_path.name}" img_path.rename(new_path) bad_count += 1 continue md5 = hashlib.md5(img_path.read_bytes()).hexdigest() if md5 in seen_hashes: print(f"[重复] {img_path} 与 {seen_hashes[md5]} 内容一致") new_path = suspected / f"dup_{hashlib.md5(img_path.read_bytes()).hexdigest()[:8]}_{img_path.name}" img_path.rename(new_path) else: seen_hashes[md5] = img_path print(f"清洗完成,坏图和重复图共隔离 {bad_count + len(seen_hashes) - len(list(root.rglob('*.jpg')))} 张")坏图用 PIL 打开识别,文件后缀是 .jpg 但内容损坏的情况非常多见;重复图用 MD5 比对文件内容而不是比较文件名或大小,因为同一张图从不同渠道下载后文件名完全不同,但二进制内容相同。所有被标记的文件先移动到隔离目录而不是直接删除,相当于留一份后悔药,人工抽查确认没有问题再删除。
提示:隔离目录里如果只按原名移动,重复文件会重名覆盖。建议在移动时给目标文件名加 hash 前缀,避免相互覆盖。
参数说明:root.rglob("*.jpg")只匹配 jpg,如果数据里有 png、jpeg,需要改成对应后缀或直接rglob("*")再按后缀过滤。MD5 计算会读取整个文件,图片多时耗时几分钟,属于正常现象。
3.2 分层划分与类别平衡:划分前不做去重就没有后悔药
清洗完就可以划分数据集。很多人直接对整个数据集 shuffle 之后按比例切,这会带来一个隐蔽问题:样本少的类别可能在验证集里只有三五张,评估结果波动巨大;更糟的是如果前面 MD5 漏掉相似图,同一内容可能同时落在 train 和 val 里。正确做法是按类别分别划分,固定随机种子,先 train/val/test 再检查交集。
import random from pathlib import Path import shutil random.seed(42) # 固定种子,保证每次划分结果一致 root = Path("dataset/images") out_root = Path("dataset/images_split") for split_name in ["train", "val", "test"]: (out_root / split_name).mkdir(parents=True, exist_ok=True) for class_dir in sorted([p for p in (root / "train").iterdir() if p.is_dir()]): images = sorted(class_dir.glob("*.*")) if len(images) < 10: print(f"[警告] 类别 {class_dir.name} 只有 {len(images)} 张,建议补数据") random.shuffle(images) n_train = int(len(images) * 0.8) n_val = int(len(images) * 0.9) for img in images[:n_train]: shutil.copy2(img, out_root / "train" / class_dir.name / img.name) for img in images[n_train:n_val]: shutil.copy2(img, out_root / "val" / class_dir.name / img.name) for img in images[n_val:]: shutil.copy2(img, out_root / "test" / class_dir.name / img.name)关键在random.seed(42)。不固定种子,每次跑划分出来的集合不同,实验结果不可复现,后面调参时很难判断提升到底来自数据变化还是模型变化。分层划分保证每个类别在三个集合里的比例一致,避免某类在验证集里消失。代码里用copy2而不是rename,因为后面还要回到原始目录核对,复制一份去训练不影响原始数据。
参数说明:8:1:1 是图片分类常用的比例;如果数据量小,可以用 7:2:1 或者把 val 凑到 20%。少于 10 张的类直接报警,这时硬训练会让模型少数类根本学不到。我一般还会在划分后打印 train 和 val 的图片文件名集合交集,如果交集非空就回头查重复。
3.3 数据增强的边界:翻转可以做,色彩抖动要克制
垃圾分类作为图像分类任务,数据增强该做,但不能照搬 ImageNet 那套。垃圾分类的判别信息往往在材质、瓶身印刷、封装形态上——塑料瓶和玻璃瓶的颜色可以一样,区分点在高光和纹路;纸盒和纸箱的区别有时就是盒盖边缘那一点形状。这就意味着色彩抖动这类增强如果调太狠,等于人为抹掉关键特征。
from torchvision import transforms train_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.7, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])RandomResizedCrop 相当于模拟不同拍摄距离,scale下限设 0.7 是防止裁到只有局部纹理、丢了整体形态。ColorJitter 里 saturation 设 0.1、hue 设 0,因为垃圾分类中色相信息很关键——有色垃圾袋、塑料瓶的颜色都是判别依据。验证集和测试集不要用任何随机增强,只做 Resize + CenterCrop + Normalize,保证评估结果没有随机性。
4. 垃圾分类训练与部署中的 5 个高频踩坑点:现象、原因与解决
即使数据集干净、脚本没写错,训练和部署阶段仍然有几个高频问题,几乎每个做垃圾分类的人都会遇到一两个。这里把现象、原因和解决方式拆开讲。
4.1 训练 loss 下降但验证准确率不涨:先看混淆矩阵再动学习率
现象:loss 曲线一路往下,训练准确率接近 99%,验证准确率卡在 70% 上下不动。
原因:最常见的是类别不平衡,交叉熵 loss 被样本多的类别主导,模型把所有图都往多数类猜。多数类对了 loss 就小,少数类全错也不影响整体数值。只看准确率平均分看不出问题。
解决:打印 per-class 的 precision/recall,同时训练时给 CrossEntropyLoss 传 weight 参数,按样本数量的反比设置:
from sklearn.metrics import classification_report # y_true 为验证集真实类别索引,y_pred 为模型预测类别索引 print(classification_report(y_true, y_pred, target_names=class_names)) # 训练时给少数类加权,n_samples 为每个类别的样本数 import torch weights = 1.0 / torch.tensor(n_samples, dtype=torch.float) weights = weights / weights.mean() # 归一化到均值 1,保持 loss 量级 criterion = torch.nn.CrossEntropyLoss(weight=weights.to(device))classification_report 直接给出每个类别的 precision、recall、f1,如果某个类 recall 是 0,问题一目了然。权重归一化到均值 1 是为了不让 loss 整体变大太多,学习率不用重调。weight 要和类别顺序一一对应,也就是 classes.txt 的顺序,对不上等于白写。
4.2 验证集里混进了训练集的图:划分前必须按内容去重
现象:验证准确率高达 96%,模型看起来完美,但上线后表现很普通。
原因:数据包里同一张图存在不同尺寸的版本,或同一物体被连续拍了好几张几乎一样的照片。按文件名划分时这些相似图被拆到了不同集合,模型在验证集上等于“背答案”。MD5 去重只能处理完全相同的内容,缩放或转码后的图需要感知哈希。
解决:在 MD5 之外再补一道感知哈希去重,把像素相似度高的图也归到一边。轻量做法是把图片缩到 8x8 灰度格,用均值比较生成哈希,hamming 距离小于阈值视为重复:
import imagehash from PIL import Image def perceptual_hash(img_path, hash_size=8): img = Image.open(img_path).convert("L").resize((hash_size, hash_size)) return imagehash.phash(img) # 两两比较时,hamming 距离 <= 5 视为相似图,移入隔离目录做数据划分时养成习惯:划分前先跑一遍相似度扫描,把重复和近似图挪到 suspected 目录,划分后再打印两个集合的文件名交集做最终确认。这一步没有多复杂,但能省掉后面所有“模型泛化差”的排查时间。
4.3 保存的模型在另一台机器上加载报错:只存 state_dict 不要存整个模型
现象:训练机上跑得好好的,换台电脑加载时报Can't get attribute 'MyDataset'或ModuleNotFoundError。
原因:保存时直接torch.save(model),整个模型对象连同里面引用的自定义 Dataset 类、脚本路径被打包在一起;换环境后类定义不存在,加载自然失败。GPU 训练的模型在 CPU 机器上加载时还会报 CUDA 相关的 device mismatch。
解决:只保存 state_dict,加载时先构造同样的模型结构,再 load_state_dict:
# 保存 torch.save({ "model": model.state_dict(), "class_names": class_names, "input_size": 224, }, "checkpoints/best.pt") # 加载(跨机器、跨 CPU/GPU 都安全) ckpt = torch.load("checkpoints/best.pt", map_location="cpu") model = build_model(num_classes=len(ckpt["class_names"])) model.load_state_dict(ckpt["model"]) model.eval()把 class_names 也存进 checkpoint,这样推理时就知道模型输出索引对应哪个类别,避免“模型还能用但不知道输出序号对不对”的尴尬。map_location="cpu"让模型先落到 CPU 再搬运,兼容性最好。
4.4 “纸盒”和“纸箱”、“矿泉水瓶”和“易拉罐”错得离谱:标签口径与细分类合并策略
现象:验证集整体准确率还行,但混淆矩阵里纸盒和纸箱、矿泉水瓶和易拉罐互相串,错误高度集中在某几对类别。
原因:标签口径不一致,标注时有人按材质分、有人按用途分;或者类别本身太细,人类标注员都未必能分清。另一个因素是类别粒度越细,类间相似度越高,对图像分辨率的要求也越高。
解决:第一步先给每个类别写一句判定规则,例如“纸盒 = 有硬质纸壳的包装,纸箱 = 棕色瓦楞纸容器”,按规则重新翻一遍数据。第二步,如果规则仍然兜不住,把易混淆的类合并成一个父类,例如统一叫“纸质包装”。合并比硬扛更划算,模型输出的类别越清晰,后续业务方越好用:
# 类别合并映射:旧类名 -> 合并后的新类名 merge_map = { "cardboard_box": "paper_packaging", "carton": "paper_packaging", "paper_cup": "paper_packaging", "plastic_bottle": "plastic_packaging", "aluminum_can": "metal_packaging", } # 按映射重建 train/val 目录,重跑一遍清洗和划分脚本合并会损失一点类别粒度,但换来的是标注一致性,模型 recall 通常显著上升。这一步在训练大改前先做,因为它直接改变数据集结构,代价最小。
4.5 CPU 推理慢得让人怀疑人生:先量化耗时再决定换网络
现象:训练时用 GPU 没感觉,模型部署到 CPU 上预测一张图要 1.5 秒,根本没法用。
原因:网络输入分辨率偏大、模型骨架过重(ResNet50 起步)、推理时每张图单独过 forward 没有 batch、CPU 线程数没设置。最常见的是压根没做速度测试,直接拿训练时的输入尺寸跑推理。
解决:先写一段 benchmark 脚本,用固定尺寸的 dummy 输入跑几十次,记录单张耗时,再决定瓶颈在哪:
import torch, time model.eval() dummy = torch.randn(1, 3, 224, 224) with torch.no_grad(): for _ in range(10): # warmup,排除显存/内存分配抖动 model(dummy) t0 = time.time() for _ in range(50): model(dummy) print(f"单张耗时: {(time.time() - t0) / 50 * 1000:.1f} ms")warmup 必做,否则第一次推理包含初始化开销,测出来的耗时会明显偏大。如果单张耗时在 100ms 以内,无需优化;在几百 ms 级别,可以先换 ResNet18 或 MobileNet 再测。模型压测和训练是两套流程,训练追求准确率,部署追求延迟,这两个指标要通盘考虑。
5. 训练脚本怎么写顺手:从预训练权重到评估指标
数据整理完毕,踩坑预案也清楚了,接下来把训练脚本写利索。这里以 PyTorch 和 torchvision 为例,垃圾分类是单标签图像分类,用 ImageFolder 直接加载目录,整体代码量不大,但每个参数都值得解释。
5.1 最小训练脚本:resnet18 改造分类头,从示例代码改成自己的数据
不要一上来就用 ResNet50 或更重的网络。垃圾分类这类小数据集,每类几百到几千张,ResNet18 已经足够判断数据质量。网络再大,只会更早过拟合,训练时间还翻倍。模型部分先加载 ImageNet 预训练权重,然后把最后的全连接层改成自己的类别数:
import torch import torch.nn as nn from torchvision import models, transforms from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) num_classes = 4 # 按自己的类别数改 model.fc = nn.Linear(model.fc.in_features, num_classes) train_dataset = ImageFolder("dataset/images/train", transform=train_transform) val_dataset = ImageFolder("dataset/images/val", transform=val_transform) train_loader = DataLoader( train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True, drop_last=False, ) val_loader = DataLoader( val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True, ) optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-4) criterion = torch.nn.CrossEntropyLoss() scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode="min", factor=0.5, patience=3 )train_dataset 里的 ImageFolder 会直接按目录名生成 label,类别顺序和目录名的字母序一致,所以 classes.txt 要和目录名保持一致。AdamW 是常见选择,微调场景学习率 3e-4 起步,比从头训练低一个数量级。ReduceLROnPlateau 在验证 loss 连续 3 个 epoch 不降时把学习率减半,这是处理学习率玄学最简单省事的办法。
提示:Windows 上 DataLoader 的 num_workers 设置过高会导致脚本卡死或反复重启,建议从 2 开始调。
训练主循环里,每个 epoch 结束后在验证集上跑一遍,记录 loss 和准确率:
best_acc = 0.0 for epoch in range(30): model.train() running_loss = 0.0 for images, labels in train_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) model.eval() correct = 0 total = 0 val_loss = 0.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) _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) val_acc = correct / total print(f"epoch {epoch+1} | train_loss {running_loss/len(train_dataset):.4f} | " f"val_loss {val_loss/len(val_dataset):.4f} | val_acc {val_acc:.4f}") scheduler.step(val_loss) if val_acc > best_acc: best_acc = val_acc torch.save({"model": model.state_dict(), "class_names": class_names}, "checkpoints/best.pt")验证集里的torch.max(outputs, 1)得到预测类别索引,这个索引就是 classes.txt 里的行号。保存 checkpoint 时只存 state_dict,并按前面的约定带上 class_names,这样下次加载不需要重新数类别。scheduler.step(val_loss)传入的是验证 loss 而不是训练 loss,因为 ReduceLROnPlateau 关注的是泛化表现。
5.2 训练过程怎么看:验证准确率之外还要盯住 per-class 指标
训练日志里只打一个 val_acc 很容易骗人:某个类别占验证集 40%,只要这个类别对了,整体 acc 就有 40% 的保底。所以每轮验证时把预测结果收集起来,训练结束后用混淆矩阵看每个类别的表现,尤其是 recall 偏低的那几个类:
from sklearn.metrics import confusion_matrix, classification_report all_preds, all_labels = [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images = images.to(device) outputs = model(images) preds = torch.argmax(outputs, dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_names=class_names)) cm = confusion_matrix(all_labels, all_preds) print(cm)classification_report 输出的每一行对应一个类别,support 列能看到这个类在验证集里到底有多少张图。如果某类 support 只有 5,那它的 recall 再高也说明不了问题。混淆矩阵重点看对角线以外数值集中的格子,那个位置就是需要的类别对,对应前面说的合并策略。
5.3 checkpoint 里该存什么:断点续训不丢类别顺序的存档方法
训练到一半断电、服务器重启是常事,所以训练脚本从第一天就要带断点续训。每轮保存时把 model.state_dict、optimizer.state_dict、epoch、best_acc、class_names 一起存进去。只存模型权重的话,续训时 optimizer 的学习率状态会丢失,后面的收敛节奏基本要重来:
ckpt = { "epoch": epoch + 1, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "best_acc": best_acc, "class_names": class_names, } torch.save(ckpt, f"checkpoints/epoch{epoch+1:02d}_acc{val_acc:.3f}.pt") # 续训 ckpt = torch.load(resume_path, map_location=device) model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) scheduler.load_state_dict(ckpt["scheduler"]) start_epoch = ckpt["epoch"]文件命名带上 epoch 和 acc,找最优模型时不用逐个加载看。class_names 必须每次保存都带上,因为目录里类别顺序一旦变动,模型输出和真实类别的对应关系就全乱了,这是很多人续训时不知不觉犯的错。
6. 把模型用起来:推理脚本、checkpoint 复用与检测扩展
模型训练完,落地第一步是写一个不受训练代码污染的推理脚本。最容易出错的地方是预处理和训练时不一致,训练用了 RandomResizedCrop,推理时忘了用 CenterCrop,预测结果当然差。下面这个函数可以直接贴进服务端调用。
6.1 推理脚本骨架:预处理和训练时保持一致,输出 top-k 而不是单一标签
from PIL import Image def predict_one(model, img_path, class_names, device="cpu"): img = Image.open(img_path).convert("RGB") img = val_transform(img).unsqueeze(0).to(device) # 用 val_transform,不做随机增强 model.eval() with torch.no_grad(): prob = torch.softmax(model(img), dim=1)[0] topk_idx = torch.argsort(prob, descending=True)[:3] return [(class_names[i], prob[i].item()) for i in topk_idx] # 调用示例 for name, score in predict_one(resnet18, "test_img.jpg", class_names): print(f"{name}: {score:.3f}")softmax 之后的概率可以当作置信度输出,排序后取 top-3 比只给一个标签更适合对接业务。如果最高置信度只有 0.4,前端可以直接提示“无法确定”。推理脚本里不写任何随机增强,这是硬性要求。
6.2 从分类到检测:YOLOv8 训练自己的数据集的扩展路线
拍到一张照片里有多个垃圾,分类模型就无能为力了,这时候要往检测方向走。YOLOv8 训练自己的数据集是现在最常见的扩展路线,数据从文件夹结构换成“图片 + txt 标签”的形式。每张图对应一个 txt,每行格式是:类别编号、归一化中心 x、归一化中心 y、归一化宽、归一化高。
images/ ├── img_001.jpg ├── img_001.txt └── img_002.jpg ... labels/ ├── img_001.txt └── img_002.txt txt 格式示例: 0 0.5 0.5 0.2 0.3 1 0.1 0.8 0.4 0.2类别编号在 yaml 配置里定义,训练前把数据集路径和类别数量填进去,训练脚本就会按这个约定读取。分类项目的数据没法直接喂给检测模型,需要把之前标注过的图片用框重新标一遍,这是一笔不小的成本,所以最开始选型时就要想清楚到底做分类还是检测。
这里也说说我的习惯:无论分类还是检测,我都坚持先跑通最小示例代码,再往里面加自己的数据。最小示例代码能跑通,说明环境、依赖、数据读取链路是通的,这时候替换成自己的数据,问题定位范围会小很多。
我最早做垃圾分类时就栽在数据划分上,从网上搜罗了一堆数据集下载下来,没清洗没去重直接开训,第二天起来验证集准确率卡在 76%。后来把重复图清掉、按类别分层重划分,同样一套网络和参数,验证集直接跳到 89%。那以后我做任何图像分类项目都先做去重再做划分,任何 checkpoint 里都存一份 class_names。这两步看着不起眼,最浪费时间。希望这篇能帮你少走一点弯路,帮到你。
本文还有配套的精品资源,点击获取