news 2026/10/2 3:51:00

PyTorch图像分类实战:用迁移学习训练12类垃圾识别模型

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch图像分类实战:用迁移学习训练12类垃圾识别模型

简介:这套12种生活垃圾分类图像数据集面向图像分类与目标检测场景,覆盖衣服、塑料、废纸等12个常见生活垃圾类别,适合需要标准分类数据进行算法练习、课程设计或项目实验的学习者与开发者。资源压缩包文件条目共2000个,绝大多数为jpg图片,另附1个json分类字典文件和1个可视化py脚本,压缩包总大小约231MB;数据在data目录下按类别存放,分为train与test两个子集,训练样本合计12415张、验证样本3100张,可直接用ImageFolder读取或作为YOLOv5分类数据使用,无需额外整理。目前已有1370人浏览学习。下载后借助json字典可明确对应12个类别标签,运行可视化脚本即可随机抽取4张图片快速预览并保存结果,便于检查图像质量与类别均衡性;整体目录结构清晰,适合快速搭建垃圾分类识别模型并进行训练、验证与调参。

1. 12类生活垃圾分类数据集开跑图像分类:先解决数据,再谈模型

有一次接了个垃圾分类的小 demo,在网上翻了半天才发现:图像分类的公开数据要么是 MNIST、CIFAR-10 这种玩具尺度,要么是带标注框的目标检测数据集,起步就是几十万张图,压根不是给“先跑通流程”准备的。最后用的是一个 12 类生活垃圾分类数据集,类别正好覆盖可回收、厨余、有害、其他四大体系,而且训练集、验证集划分得清清楚楚,解压就能喂给模型。

这篇文章就把这条路径完整走一遍:先搞清楚这份数据集的目录和标签是怎么组织的,然后用 PyTorch 跑通一个最精简的图像分类训练流程,再把标签错位、数据泄露、类别不均衡这些容易翻车的点单独拿出来排查。第一次搭分类流水线的新手可以直接照着跑,老手也能用这里的查重和泛化验证方法给手里的数据补一道保险。

2. 看懂训练集与验证集的目录结构:12类标签怎么对齐

2.1 一份能直接开训的数据集通常长什么样

我拿到的这个版本,整理方式非常标准:根目录下只有train/和val/两个文件夹,每个文件夹里按类别再建 12 个子目录,图片直接平铺在类别目录下。这种结构最大的好处是 PyTorch 的ImageFolder能直接吃,不用自己写 Dataset 去解析 CSV 或 JSON 标注。

split_data/ ├── train/ │ ├── 01_cardboard/ # 纸板 │ ├── 02_glass/ # 玻璃 │ ├── 03_metal/ # 金属 │ ├── 04_plastic/ # 塑料 │ ├── 05_trash_food/ # 厨余剩饭 │ ├── 06_fruit_peel/ # 果皮 │ ├── 07_battery/ # 电池 │ ├── 08_lamp/ # 灯管 │ ├── 09_clothes/ # 衣物 │ ├── 10_shoes/ # 鞋子 │ ├── 11_tissue/ # 卫生纸 │ └── 12_cigarette_butt/ # 烟头 └── val/ └── (与 train 完全相同的 12 个子目录)

注意一点:不管数据集是从哪下载的,解压之后第一件事永远是打开train和val两个目录数一下数量,别急着训练。常见的问题包括验证集某个类别文件夹是空的、或者某个文件夹里混进了.txt说明文件——ImageFolder会尝试把.txt当作图片解码,到时候报错报得莫名其妙。

2.2 12类标签的体系与命名对齐

这 12 个类别并不是随便凑的,它们基本遵循国内生活垃圾“可回收物、厨余垃圾、有害垃圾、其他垃圾”的四分法:

类别编号类别名常见实物示例容易混淆的对象
01纸板快递纸箱、纸盒卫生纸(材质差异大,但颜色接近)
02玻璃玻璃瓶、碎玻璃塑料瓶(透明瓶身)
03金属易拉罐、金属盖塑料瓶(同为饮料容器)
04塑料塑料瓶、塑料袋、塑料餐具玻璃、金属容器
05厨余剩饭剩菜、面条、骨头果皮(同为湿垃圾)
06果皮香蕉皮、苹果核、橙子皮厨余剩饭
07电池干电池、充电电池金属(外壳是金属)
08灯管荧光灯管、节能灯玻璃
09衣物旧T恤、裤子、毛巾布料材质的其他垃圾
10鞋子运动鞋、拖鞋衣物
11卫生纸用过的纸巾、厨房纸纸板
12烟头烟蒂、烟灰无明显混淆对象

拿到数据集之后先做一次“人对表”:随便打开01_cardboard和11_tissue各看二十张图,确认目录名和图片内容对得上。图像分类数据集的标签错位问题比大部分人想象的严重,网上拉下来的数据经常出现“文件夹叫塑料,里面一半是玻璃瓶”的情况,这种脏数据越早发现损失越小。

2.3 训练集与验证集比例不对怎么办:自己重新划分

原版数据集的划分比例通常是 8:2 或 9:1,训练集占大头。但有时候你从别人的网盘或者 GitHub 上拿到的版本,验证集可能被抽得太少——比如某个类别验证集只有十几张图,跑出来的准确率波动极大,换一个随机种子结果就完全不同。

碰到这种情况,我一般会重新做一次分层划分。核心是保证每个类别里训练集和验证集的比例一致,不能直接全局乱序切。

import os import random import shutil from collections import defaultdict src = "raw_data" # 原始数据根目录,下面直接是12个类别文件夹 dst_root = "split_data" # 重新划分后的输出目录 train_ratio = 0.8 random.seed(42) # 固定随机种子,让划分结果可复现 # 1. 按类别收集所有图片路径 class_files = defaultdict(list) for cls in os.listdir(src): cls_dir = os.path.join(src, cls) if not os.path.isdir(cls_dir): continue for f in os.listdir(cls_dir): if f.lower().endswith((".jpg", ".jpeg", ".png")): class_files[cls].append(os.path.join(cls_dir, f)) # 2. 逐类别打乱并切分 for cls, files in class_files.items(): random.shuffle(files) split_idx = int(len(files) * train_ratio) for phase, part in [ ("train", files[:split_idx]), ("val", files[split_idx:]), ]: out_dir = os.path.join(dst_root, phase, cls) os.makedirs(out_dir, exist_ok=True) for img_path in part: dst_path = os.path.join(out_dir, os.path.basename(img_path)) shutil.copy(img_path, dst_path) # 用copy而不是move,出事还有后悔药 print(f"{cls}: train={split_idx}, val={len(files) - split_idx}")

这段脚本的逻辑很简单:先按类别分组,再逐类别打乱、按比例切分。这里用copy而不用move是有意的——如果切分脚本有 bug,或者发现验证集里混入了不该出现的图片,原始数据还在,可以重新来;如果你直接move,数据被物理移动之后就很难追溯了。

参数上的两个关键点:train_ratio = 0.8适合万级数据量,如果整个数据集只有两三千张图,建议提到 0.9,否则验证集每类不到二十张,评估结果统计意义太弱;random.seed(42)必须固定,否则每次跑出来的划分都不一样,实验结果无法复现。

3. 用PyTorch加载训练集跑通图像分类最小流程:ImageFolder、Transforms与参数设置

3.1 直接调用ImageFolder,而不是手写Dataset

很多人一上来就想写一个自定义 Dataset 去读 CSV 标签,其实在文件夹结构已经按类别组织好的情况下完全没必要。torchvision 自带的ImageFolder就是为这种目录结构设计的,它会自动递归扫描子目录,把每个文件夹当作一个类别,按字母序生成标签。

from torchvision import datasets data_root = "split_data" train_ds = datasets.ImageFolder(f"{data_root}/train") val_ds = datasets.ImageFolder(f"{data_root}/val") # 这一步必须打印出来核对,标签对齐全靠它 print("train classes:", train_ds.class_to_idx) print("val classes:", val_ds.class_to_idx) print("train size:", len(train_ds)) print("val size:", len(val_ds))

class_to_idx是ImageFolder生成的类别映射字典,它的排序规则是字符串字母序,而不是你在业务上的类别编号。比如01_cardboard会排到07_battery前面,因为字符串"07"小于"01"在字典序上不成立——等等,这里要注意,"01_cardboard"和"07_battery"按 ASCII 码比较,"0"相同,"1"小于"7",所以顺序是对的,但如果你有的目录带中文名、有的带英文名,排序结果就完全不可预测了。

所以在训练之前,务必要对一下class_to_idx和你心里的类别顺序。后续所有评估、混淆矩阵、导出模型时的类别名,都以这个字典为准。这个习惯能帮你避免一个非常隐蔽的坑:模型预测“玻璃”,实际上 idx=2 指向的是金属。

3.2 训练集和验证集的Transforms为什么必须分开写

垃圾图片的拍摄条件极其不统一,有的在室内灯光下拍的,有的是在垃圾桶里仰拍的,还有的是隔着垃圾袋拍的。训练集要做足够的增强来模拟这种多样性,但验证集必须保持稳定、可复现,所以两套 transform 不能混用。

from torchvision import transforms # ImageNet预训练模型要求的归一化参数,迁移学习必须用这套 mean = [0.485, 0.456, 0.406] std = [0.229, 0.224, 0.225] train_tf = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3), transforms.ToTensor(), transforms.Normalize(mean, std), ]) val_tf = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean, std), ])

几个参数值得单独说。RandomResizedCrop的scale=(0.6, 1.0)我故意没有按默认的(0.08, 1.0)来写,因为垃圾图片的主体通常占画面比例很大,裁剪太狠会把瓶子拦腰截断,反而把最关键的表征切没了。ColorJitter的三个取值 0.3 是经验值,生活垃圾在黄光、白光、户外自然光下色温差异很大,不加颜色抖动的话,模型容易把“灯光颜色”当成类别特征。验证集的Resize(256) + CenterCrop(224)是 ImageNet 时代的标准评估范式,保持输入分布稳定,不引入随机性。

这里提一句:如果以后你想把同样的数据迁移到 YOLO 训练目标检测模型,transforms 的逻辑会完全不一样,但数据清洗的思路是相通的——先保证训练集和验证集同源且无泄露,再谈增强策略。

3.3 训练主循环:一个Epoch里盯住三个指标

训练代码不需要花哨,一个 ResNet-18 加交叉熵损失就足以在这个数据集上拿到一个能看的基线。代码的核心是把训练和评估拆开:训练阶段开 dropout 和增强,评估阶段必须关掉它们。

import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import models loader = { "train": DataLoader(train_ds, batch_size=64, shuffle=True, num_workers=4), "val": DataLoader(val_ds, batch_size=64, shuffle=False, num_workers=4), } model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) model.fc = nn.Linear(model.fc.in_features, 12) # 把1000类换成12类 criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) def eval_acc(model, val_loader): model.eval() correct = 0 total = 0 with torch.no_grad(): for x, y in val_loader: pred = x.argmax(dim=1) # 注意这里 correct += (pred == y).sum().item() total += y.size(0) return correct / total for epoch in range(20): model.train() running_loss = 0.0 for x, y in loader["train"]: optimizer.zero_grad() out = model(x) loss = criterion(out, y) loss.backward() optimizer.step() running_loss += loss.item() val_acc = eval_acc(model, loader["val"]) print(f"epoch={epoch}, loss={running_loss / len(loader['train']):.3f}, val_acc={val_acc:.3f}")

model.fc = nn.Linear(...)是迁移学习的关键操作,将 ResNet-18 最后接的 1000 类全连接层替换成 12 类。CrossEntropyLoss内部已经带了 softmax,所以模型输出层不要额外加。lr=1e-4是迁移学习的安全起点,比从头训练的 1e-3 低一个量级,因为预训练权重已经在一个大分布上收敛过,学习率太大会把学好的特征直接冲垮。

eval_acc函数里有一行注释“注意这里”,是因为我在实际项目里犯过几次低级错误:比如忘了model.eval(),导致评估时 dropout 还在生效,acc 忽高忽低;或者把model(x)的结果直接argmax而没有经过softmax,虽然 argmax 不受 softmax 影响,但如果哪次改成取 top-k 概率,就容易写错。评估阶段用torch.no_grad()包住,节省显存和算力。

4. 图像分类模型怎么挑:预训练ResNet还是小Transformer

4.1 12类×几千张图片量级下的选型逻辑

很多新手第一个问题就是“Transformer 是不是一定比 CNN 强”。这个结论在大规模数据集上成立,但在这种一两万张图的生活垃圾分类任务里,不一定。12 个类别,每个类别几百到上千张图,这个量级最核心的矛盾是过拟合,而不是模型容量不够。

模型参数量是否需要预训练微调成本常见问题
ResNet-18约11M强烈建议低容量偏小,复杂场景稍吃力
ResNet-50约25M强烈建议中速度比18慢,但精度更稳
EfficientNet-B0约5M强烈建议低输入分辨率敏感,别乱改
ViT-Tiny约5M必须中收敛慢,数据增强不够会过拟合
Swin-Transformer-T约28M必须中高显存占用大,12类任务有点浪费

在这个量级上,我一般首选 ResNet-18 或 EfficientNet-B0。ViT 系模型在预训练+强增强拉满时确实能达到更高的上限,但它对训练策略更敏感,需要更长的 epoch、更精细的 learning rate schedule,如果你只是想快速跑通一个可用的分类系统,没有必要去承担这个调参成本。Transformer 图像分类的优劣势,在 CIFAR-10 这类小数据上已经被讨论很多了——数据量不足时它并不比同参数量 CNN 有优势。

4.2 用timm加载预训练权重做迁移学习的标准写法

torchvision 自带ResNet系列,但如果你要尝试 EfficientNet、ViT、Swin 这些模型,最省事的路径是timm。它封装了几百个预训练模型,而且统一了接口,换模型只改一行字符串。

import timm # 三种典型选择,按数据量来 model = timm.create_model("resnet18", pretrained=True, num_classes=12) # model = timm.create_model("efficientnet_b0", pretrained=True, num_classes=12) # model = timm.create_model("vit_tiny_patch16_224", pretrained=True, num_classes=12)

第一次运行pretrained=True时,timm 会从 HuggingFace 的模型仓库下载权重,网速不稳定时可能卡在下载阶段,需要确认网络能访问外网,或者提前用 download 脚本把权重拉下来放到缓存目录。这里有个细节:timm 和 torchvision 对预训练权重的命名不一样,"resnet18"在 torchvision 里要用weights=models.ResNet18_Weights.IMAGENET1K_V1显式指定,而 timm 直接用pretrained=True就行,两者不要混用。

换成 EfficientNet 后训练代码完全不用改,DataLoader 的输出尺寸还是 224。唯一要注意的是num_classes=12这个参数,timm 会帮你把模型的分类头替换掉,但如果你是在 torchvision 的模型上改,就要自己手动替换最后的全连接层。

4.3 三个必调的参数:学习率、batch_size、epoch数

面向小规模数据集的迁移学习,参数上有一套相对稳妥的组合:学习率用 1e-4 起步,验证集 loss 下降停滞时降到 1e-5 左右;batch size 按显存来,24G 显卡可以到 128,但 12 类别数据不追求大 batch;epoch 数 20 到 40 之间,配合早停。

from torch.optim.lr_scheduler import CosineAnnealingLR scheduler = CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6) # 在epoch循环末尾加上: # scheduler.step()

CosineAnnealingLR的T_max设置成和总 epoch 数一致,学习率就会从初始值按余弦曲线下降到接近eta_min。这比固定学习率更省心,最后几个 epoch 的小学习率能带来一到两个点的准确率提升,尤其是对分类头(就是那个新换上的 fc 层)非常有效。

再给一个进阶技巧:如果你发现直接全量微调效果不好,可以先把 backbone 冻结,只训练分类头跑几个 epoch,再解冻。冻结的方式是把requires_grad置为 False。解冻之后别忘了把学习率再降一个量级,否则 backbone 的预训练特征容易被破坏——这个坑在下一章的避坑部分会详细拆。

5. 生活垃圾分类数据集训练常见的5个翻车点:标签错位、类别不均与验证集泄露排查

5.1 训练loss正常下降,但验证集输出完全对不上业务标签

现象:loss 从 2.0 降到 0.3,打印预测结果却发现“纸板”被识别成“玻璃”,打开图片看模型明明分对了,但打印出来的标签名却是错的。

原因:ImageFolder按目录名的字符串排序生成索引,你数据集里的目录可能叫cardboard、glass,也可能是中文名,排序结果和业务方的标签顺序完全是两套。

解决:训练前强制打印class_to_idx,并建立自己的标签映射文件。我习惯写一个label_map.py,把class_to_idx的输出粘进去,后续所有推理、混淆矩阵、导出 ONNX 时的类别名,都以这份映射为准,绝不在代码里硬编码数字索引。

# label_map.py CLASS_TO_IDX = { "01_cardboard": 0, "02_glass": 1, "03_metal": 2, "04_plastic": 3, "05_trash_food": 4, "06_fruit_peel": 5, "07_battery": 6, "08_lamp": 7, "09_clothes": 8, "10_shoes": 9, "11_tissue": 10, "12_cigarette_butt": 11, }

用这份文件对齐之后,训练、评估、导出的标签就全程一致了,不会再出现“模型是对的,代码是错的”这种让人想砸电脑的情况。

5.2 验证集Acc比训练集高8到10个百分点,别高兴太早

现象:训练集 acc 84%,验证集 acc 92%,而且这个差距稳定复现。表面看模型泛化很好,实际上是数据泄露了。

原因:很多公开数据集在划分训练集和验证集时是直接按文件夹顺序或者随机切分的,同一个场景下连续拍摄的几十张同源图片会被同时分到两边。模型在训练时已经“记住”了这些图片,验证时只是做了一次记忆复现。

解决:先用 MD5 对全量图片做一次查重,把训练集和验证集中完全相同的图片找出来。

import hashlib import os def file_md5(path, chunk=4096): h = hashlib.md5() with open(path, "rb") as f: for block in iter(lambda: f.read(chunk), b""): h.update(block) return h.hexdigest() train_hashes = {} for cls in os.listdir("split_data/train"): cls_dir = os.path.join("split_data/train", cls) for f in os.listdir(cls_dir): p = os.path.join(cls_dir, f) train_hashes[file_md5(p)] = p dup_count = 0 for cls in os.listdir("split_data/val"): cls_dir = os.path.join("split_data/val", cls) for f in os.listdir(cls_dir): p = os.path.join(cls_dir, f) if file_md5(p) in train_hashes: dup_count += 1 print(f"重复图片: {p} == {train_hashes[file_md5(p)]}") print(f"验证集中重复图片总数: {dup_count}")

MD5 查重是第一步,做完之后还要警惕“近似重复”,也就是同一物体从不同角度、不同亮度拍出来的图片。这类图片 MD5 查不出来,需要用感知哈希或者直接抽 100 张人工比对。我的经验是,验证集 Acc 高过训练集 5 个点以上就要停下来怀疑一下,优先排查数据泄露。

5.3 某些类别永远分不好,先看混淆矩阵而不是急着加数据

现象:“果皮”和“厨余剩饭”天天互相混,“电池”和“金属”也纠缠不清,单独看每类的 acc,这两对永远垫底。

原因:垃圾类别本身就有视觉边界模糊的问题。香蕉皮和剩饭经常出现在同一个餐盒里,背景高度重叠;电池的外壳就是金属,模型看到的特征和“金属”类几乎一样。这种类别关系复杂的情况属于正常现象,不是模型坏了。

解决:先统计混淆矩阵,看错误是不是集中在某一两组类别之间。如果果皮和剩饭主要在彼此之间混,最简单的处理是把它们合并成一个“厨余垃圾”类别,重训一次。如果不想合并,就针对性地清洗标签——把果皮类别里那些背景是剩饭盒的图片单独挑出来重新标注或删除。盲目给其中一类加数据往往是白费功夫,因为问题出在类间区分度不够,不是样本量不够。

5.4 DataLoader worker报OSError,训练跑一半中断

现象:训练到第 3 个 epoch 突然报错,提示和DataLoader worker process相关,或者PIL无法识别图片文件。

原因:垃圾数据集的图片来源特别杂,很多是从网上爬取的,有些扩展名是.jpg但实际编码是 WebP,或者文件只有几百字节已经损坏。ImageFolder在扫描时只看扩展名,真正解码要等到训练过程中,所以错误会在 epoch 中途随机出现,很难复现。

解决:在送入 DataLoader 之前做一次图片有效性过滤。

from PIL import Image def is_valid_image(path): try: with Image.open(path) as im: im.load() # 触发实际解码 if im.mode not in ("RGB", "L"): return False # 过滤灰度图和RGBA图 return True except Exception: return False bad_files = [] for cls in os.listdir("split_data/train"): cls_dir = os.path.join("split_data/train", cls) for f in os.listdir(cls_dir): p = os.path.join(cls_dir, f) if not is_valid_image(p): bad_files.append(p) print(f"坏图片数量: {len(bad_files)}") for p in bad_files[:20]: print(p)

检查出来的坏图片可以直接删除,或者移动到单独的_corrupted/目录下留着备用。这一步看着麻烦,但能避免训练到一半白白浪费几个小时。注意im.mode的检查也很重要,RGBA 四通道图片如果直接喂给模型,会在 ToTensor 之后多出一个通道,导致模型输入维度不匹配。

5.5 迁移学习直接全量微调,效果反而不如冻结backbone

现象:用pretrained=True之后直接把整网扔进训练,跑了 20 个 epoch,验证集 acc 一直在 80% 上下浮动,还不如别人用冻结 backbone 只训练分类头的结果。

原因:学习率设置不合理。从头训练模型用 1e-3 没问题,但迁移学习中预训练特征已经很好了,大学习率会把这些特征在最初的几个 step 里彻底冲乱,后续再怎么训都很难恢复。生活垃圾与 ImageNet 的域差虽然存在,但底层纹理、边缘特征依然通用,不值得用大学习率去破坏。

解决:采用两阶段训练,先冻结 backbone,只让分类头学习。

for name, param in model.named_parameters(): if name != "fc.weight" and name != "fc.bias": param.requires_grad = False # 此时optimizer只包含fc层的参数,学习率可以大一些 optimizer = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=1e-3 ) # 训练5个epoch后,解冻backbone并重建optimizer,学习率降到1e-4 for param in model.parameters(): param.requires_grad = True optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

冻结阶段的学习率可以放到 1e-3,因为可训练参数只有全连接层,参数量小,不容易过拟合。解冻之后再降回 1e-4。如果还想再细一点,可以对 backbone 和分类头分别设置不同学习率,用二组参数传入同一个 AdamW,但普通任务中两阶段就够了。

6. 用你自己的照片验证模型泛化力:混淆矩阵和Bad Case才是数据集的照妖镜

6.1 收集20张手机实拍图做冒烟测试

训练跑完先别急着看测试集指标,我习惯用自己的手机拍 20 张照片做冒烟测试。重点拍训练集里“不会出现”的场景:压扁的塑料瓶、装了垃圾袋的垃圾桶、夜间灯光下的玻璃瓶。这些照片才是真正反映模型能不能用的标尺。

from PIL import Image def predict(image_path, model, val_tf, classes): img = Image.open(image_path).convert("RGB") x = val_tf(img).unsqueeze(0) # 加batch维度 with torch.no_grad(): logits = model(x) idx = logits.softmax(dim=1).argmax().item() return classes[idx], logits.softmax(dim=1).max().item() # classes直接取val_ds.classes # print(predict("test_photos/plastic_bottle.jpg", model, val_tf, val_ds.classes))

val_tf复用了之前验证集的预处理,不是训练集那套带增强的流程。.convert("RGB")必须写,因为手机照片可能是 RGBA 或者带 EXIF 方向信息,PIL 读取后需要统一到三通道。冒烟测试的目的不是追求 100% 准确率,而是看模型的错误是不是符合直觉——把瓶子认成金属可以接受,把瓶子认成食材就说明特征提取出了问题。

6.2 用混淆矩阵定位系统的真实弱点

测试集上算一个整体 acc 是不够的。垃圾类别的用户痛点往往集中在某几对类别上,所以我会把验证集预测结果收集起来,画一张归一化的混淆矩阵。

from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay # y_true和y_pred由上一章eval流程收集,注意索引要对应val_ds的classes cm = confusion_matrix(y_true, y_pred, normalize="true") disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=val_ds.classes) disp.plot()

normalize="true"让每一行和为 1,看的是“某个类别的真实样本被预测成了哪些类别”的分布。观察对角线之外的高亮区域,如果[果皮][厨余剩饭]这个格子特别亮,说明这两类确实分不开,可以明确去做标签合并;如果[塑料][玻璃]亮,说明透明容器在不同材质间是模型的主要困惑点,这时候加数据要优先加“透明塑料瓶 vs 透明玻璃瓶”这种成对的对比样本,而不是盲目扩充整个类。

我做这个流程最亏的一次教训,是拿到的数据表面分好了训练集和验证集,实际验证集里躺着 100 多张和训练集完全重复的图片,测试 acc 96%,现场演示直接翻车。从那以后,我的固定动作就是先跑一遍 MD5 查重再训练,这个习惯帮我在后面几个项目里躲过了类似的坑。希望帮到你。

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

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

Agent Skills与MCP的区别与协同:从原理到数据库巡检实战

先说一个我最近经常被问到的问题:“Agent Skills是不是要取代MCP?”问的人多了,我发现大家其实是被两个概念的命名带偏了——都带“Agent/模型”的影子,都像是给AI加能力,很容易被当成同一类东西。实际上在真实项目里&…

作者头像 李华
网站建设 2026/10/2 3:50:26

信息化主管的职责拆解与实操:从选型到数据治理的落地指南

做企业信息化主管这些年,我有一个很深的感触:这个岗位在公司里的位置相当微妙。老板觉得你是“搞IT的”,业务部门觉得你是“管系统的”,下属觉得你是“审批流程的”,只有你自己知道,你其实是在一座由老系统…

作者头像 李华
网站建设 2026/10/2 3:50:25

医疗器械集成供应链优化:从诊断到落地的咨询方案解析

医疗器械行业的供应链,和普通制造业完全是两回事。一台超声设备要出口到欧洲,对应的注册证、UDI编码、运输温湿度记录、当地代理商的服务能力备份,每一项都要对得上。迈瑞医疗这种体量的企业,产品线横跨生命信息与支持、体外诊断、…

作者头像 李华
网站建设 2026/10/2 3:49:45

SpringBoot+Docker+Jenkins:从零搭建CI/CD自动化部署流水线

做了几年 Java 后端,最烦的就是“本地编译没问题,一上线就各种崩”这种事。反复打 jar 包、传服务器、手动重启,一两次还能忍,项目一多,每周都能烧掉大半天。后来我把 SpringBoot、Docker、Jenkins 串成一条自动化的构…

作者头像 李华
网站建设 2026/10/2 3:49:19

视频运动检测工具DVR-Scan:用MOG2背景减除从监控录像中提取有效片段

简介:DVR-Scan视频运动检测工具.zip是一份面向计算机视觉学习者、毕业设计开发者及安防监控、交通管理等场景工程人员的资源包,用于对视频文件中的运动事件进行智能识别、标注与记录,核心依托OpenCV、机器学习与图像识别技术。压缩包共112个文…

作者头像 李华