简介:本资源为常规茶叶叶片病害图像分类数据集,面向从事农业图像识别、深度学习入门与CNN分类实践的学生及研究人员,可用于训练和评估茶叶病害自动识别模型。数据集已标注,共划分5个类别,包括褐枯病、灰枯萎病、红点病等,具体类别信息可查看包内json文件;数据按训练集、验证集、测试集分别存放,每类图片归入对应目录,便于直接加载训练。压缩包共2000个文件,以1998张jpg图像为主,另含1个py脚本与1个json标注文件,整体约21.68MB,其中show脚本可用于数据集可视化预览。目前已有104人学习下载。读者可借助该数据集快速搭建分类实验流程,结合配套的CNN分类网络改进系列内容,完成数据加载、模型训练与效果对比,适合作为课程设计、论文实验或算法验证的基础数据。
1. 茶叶叶片病害分类数据集:4,000 张标注图能跑出什么名堂
手里有一份约 4,000 张的常规茶叶叶片病害图像分类数据集,已标注,标签覆盖常见叶部病害类别。这个量级放在图像分类任务里不算大,但足够把一条完整的训练链路跑通:从数据清洗、类别平衡、增强策略,到选一个 backbone、调通训练脚本、看混淆矩阵定位问题。它适合两类人:一类是想入门图像分类但不想从零标注的开发者,另一类是手里有类似农业图像数据、想验证方案可行性的工程师。4,000 张的规模决定了它不适合直接冲 SOTA,但非常适合做 baseline 验证、消融实验和部署前的可行性测试。下面按「数据怎么用 → 模型怎么选 → 训练怎么跑 → 坑在哪」的顺序拆开讲。
2. 先搞清楚这 4,000 张茶叶叶片图能怎么用
2.1 数据集的典型构成与类别分布
常规茶叶叶片病害图像分类数据集,通常按文件夹组织,每个子文件夹对应一个类别。常见类别包括:茶饼病、茶炭疽病、茶白星病、茶轮斑病、健康叶片等。约 4,000 张的体量,如果分 5 到 6 类,平均每类 600 到 800 张,但实际分布往往不均匀——健康叶片和常见病害样本多,稀有病害可能只有一两百张。
拿到数据后第一件事不是直接训练,而是统计类别分布。用下面这段脚本快速摸清家底:
import os from collections import Counter from pathlib import Path data_root = Path("tea_leaf_dataset") # 替换为实际路径 class_counts = Counter() for cls_dir in sorted(data_root.iterdir()): if cls_dir.is_dir(): imgs = [f for f in cls_dir.iterdir() if f.suffix.lower() in (".jpg", ".jpeg", ".png", ".bmp")] class_counts[cls_dir.name] = len(imgs) total = sum(class_counts.values()) print(f"总图片数: {total}") for cls, cnt in class_counts.most_common(): print(f"{cls}: {cnt} 张 ({cnt/total*100:.1f}%)") # 计算不平衡比 max_c, min_c = max(class_counts.values()), min(class_counts.values()) print(f"最大/最小类别比: {max_c/min_c:.1f}")这段代码遍历数据集根目录下的每个类别文件夹,统计图片数量并计算占比。关键参数是data_root,指向解压后的数据集根目录。输出里的「最大/最小类别比」很重要:如果超过 5:1,训练时就需要做类别加权或重采样,否则模型会偏向多数类。常见做法是用WeightedRandomSampler给少数类更高采样概率,或者在 loss 里传class_weights。
2.2 图像质量筛查:别让脏数据毁掉训练
农业图像数据集的一个通病是采集环境不可控:光照过曝、叶片重叠、背景杂乱、甚至混入非茶叶图片。4,000 张里如果有 5% 到 10% 的脏数据,验证集指标会明显虚高或抖动。我一般会做三步筛查:
第一步,检查图片尺寸分布。用 PIL 批量读取宽高,找出异常值:
from PIL import Image from pathlib import Path import numpy as np sizes = [] for img_path in Path("tea_leaf_dataset").rglob("*.jpg"): with Image.open(img_path) as im: sizes.append(im.size) ws, hs = zip(*sizes) print(f"宽度: min={min(ws)}, max={max(ws)}, mean={np.mean(ws):.0f}") print(f"高度: min={min(hs)}, max={max(hs)}, mean={np.mean(hs):.0f}") # 找出极端尺寸 for p, (w, h) in zip(Path("tea_leaf_dataset").rglob("*.jpg"), sizes): if w < 100 or h < 100: print(f"过小: {p} ({w}x{h})")宽度或高度低于 100 像素的图片基本没有训练价值,直接剔除。尺寸差异过大时,统一 resize 到 224×224 或 256×256 是标准操作,但要注意长宽比失真问题——茶叶叶片接近椭圆形,强行拉伸会改变形状特征。更稳妥的做法是短边 resize 后中心裁剪,或者 padding 到正方形再缩放。
第二步,肉眼抽查。从每个类别随机抽 20 张拼成网格图,快速过一遍。这一步没有脚本能替代,因为「叶片是否完整」「病害特征是否清晰」需要人判断。第三步,检查是否有重复图片。用感知哈希(pHash)去重:
import imagehash from PIL import Image from pathlib import Path from collections import defaultdict hashes = defaultdict(list) for img_path in Path("tea_leaf_dataset").rglob("*.jpg"): with Image.open(img_path) as im: h = str(imagehash.phash(im)) hashes[h].append(str(img_path)) duplicates = {k: v for k, v in hashes.items() if len(v) > 1} print(f"发现 {len(duplicates)} 组重复图片") for k, v in list(duplicates.items())[:5]: print(f"哈希 {k}: {v}")pHash 对轻微缩放和压缩不敏感,适合找近似重复。阈值方面,完全相同的哈希值就是重复图,直接保留一张即可。如果重复组很多,说明采集时可能有连拍或视频抽帧,需要警惕训练集和验证集之间的泄漏。
2.3 划分训练集、验证集、测试集的比例与策略
4,000 张的规模,推荐按 7:1.5:1.5 划分,即训练集约 2,800 张、验证集 600 张、测试集 600 张。如果类别不平衡,必须做分层抽样(stratified split),保证每个子集的类别比例一致。用 scikit-learn 的train_test_split配合stratify参数:
import shutil from sklearn.model_selection import train_test_split from pathlib import Path data_root = Path("tea_leaf_dataset") output_root = Path("tea_split") all_files, all_labels = [], [] for cls_dir in sorted(data_root.iterdir()): if cls_dir.is_dir(): for f in cls_dir.iterdir(): if f.suffix.lower() in (".jpg", ".jpeg", ".png"): all_files.append(f) all_labels.append(cls_dir.name) # 先分训练集和临时集 train_f, temp_f, train_l, temp_l = train_test_split( all_files, all_labels, test_size=0.3, stratify=all_labels, random_state=42) # 临时集再分验证和测试 val_f, test_f, val_l, test_l = train_test_split( temp_f, temp_l, test_size=0.5, stratify=temp_l, random_state=42) for split_name, files, labels in [("train", train_f, train_l), ("val", val_f, val_l), ("test", test_f, test_l)]: for f, l in zip(files, labels): dst = output_root / split_name / l dst.mkdir(parents=True, exist_ok=True) shutil.copy2(f, dst / f.name) print(f"训练集: {len(train_f)}, 验证集: {len(val_f)}, 测试集: {len(test_f)}")stratify参数确保每个子集的类别分布与原始数据一致,random_state固定后结果可复现。划分完成后,训练集用于梯度更新,验证集用于调超参和早停,测试集只在最后评估一次。常见错误是反复用测试集调参,导致指标虚高——测试集一旦用过就不再「干净」了。
3. 选哪个图像分类模型:从 ResNet 到 Transformer 的落地取舍
3.1 小数据集上 backbone 的选型逻辑
4,000 张图在深度学习里属于小样本范畴。模型越大,过拟合风险越高。我的经验是:优先选参数量在 5M 到 25M 之间的 backbone,配合强增强和正则化。具体推荐:
| 模型 | 参数量 | 输入尺寸 | 适合场景 | 注意事项 |
|---|---|---|---|---|
| ResNet-18 | 11.7M | 224 | 快速 baseline | 需要预训练权重 |
| ResNet-50 | 25.6M | 224 | 精度优先 | 过拟合风险中等 |
| EfficientNet-B0 | 5.3M | 224 | 部署友好 | 对增强敏感 |
| MobileNetV3-Small | 2.5M | 224 | 边缘设备 | 精度上限较低 |
| ViT-B/16 | 86M | 224 | 数据充足时 | 4,000 张不够,需强增强 |
| Swin-Tiny | 28M | 224 | 折中方案 | 训练时间较长 |
Transformer 类模型(ViT、Swin)在 ImageNet 上表现优异,但 4,000 张图直接从头训练基本会翻车。如果要用,必须加载大规模预训练权重,并且冻结前几层。我一般先用 ResNet-18 跑一个 baseline,确认数据 pipeline 没问题后,再换 EfficientNet-B0 或 Swin-Tiny 对比。
3.2 用 timm 加载预训练模型的最小示例
timm库统一了各种 backbone 的加载接口,切换模型只需改一个字符串:
import timm import torch import torch.nn as nn # 查看 timm 支持的模型列表(部分) # model_list = timm.list_models(pretrained=True) # print([m for m in model_list if "resnet" in m][:10]) def build_model(model_name="resnet18", num_classes=6, pretrained=True): model = timm.create_model( model_name, pretrained=pretrained, num_classes=num_classes, drop_rate=0.3, # 分类头 dropout drop_path_rate=0.1, # stochastic depth,对 Transformer 更有效 ) return model device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = build_model("resnet18", num_classes=6).to(device) # 统计参数量 total_params = sum(p.numel() for p in model.parameters()) trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"总参数: {total_params/1e6:.1f}M, 可训练: {trainable_params/1e6:.1f}M")timm.create_model的pretrained=True会从网络下载 ImageNet 预训练权重。drop_rate控制分类头前的 dropout 概率,小数据集建议设 0.2 到 0.5。drop_path_rate是 stochastic depth,对 ResNet 效果有限,但对 Swin Transformer 这类结构能明显抑制过拟合。num_classes必须与数据集类别数一致,否则最后的全连接层维度对不上。
如果显存有限,可以冻结 backbone 的前几个 stage,只训练后面的层和分类头:
# 冻结前两个 stage(以 ResNet 为例) for name, param in model.named_parameters(): if "layer1" in name or "layer2" in name or "conv1" in name: param.requires_grad = False # 验证可训练参数变化 trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"冻结后可训练参数: {trainable/1e6:.1f}M")冻结策略适合数据量极小(少于 1,000 张)或算力紧张的场景。4,000 张的规模,我建议先全量微调,如果验证集 loss 震荡再考虑冻结。
3.3 数据增强:小数据集的生命线
4,000 张图要撑起一个泛化能力可用的模型,增强策略必须到位。基础增强包括随机裁剪、水平翻转、颜色抖动;进阶增强可以用 Mixup、CutMix、RandAugment。用torchvision.transforms组合:
from torchvision import transforms train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.6, 1.0)), # 随机裁剪缩放 transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.2), # 叶片方向不固定 transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.05), transforms.RandomRotation(degrees=30), # 小角度旋转 transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), transforms.RandomErasing(p=0.25, scale=(0.02, 0.15)), # 随机遮挡 ]) val_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])RandomResizedCrop的scale=(0.6, 1.0)表示随机裁剪原图 60% 到 100% 的区域,模拟不同拍摄距离。RandomVerticalFlip概率设低一些(0.2),因为叶片正反面特征不同,过度翻转可能引入噪声。ColorJitter的hue控制在 0.05 以内,避免颜色失真导致病害特征被破坏。RandomErasing模拟叶片遮挡,提升模型对局部缺失的鲁棒性。验证集只用 resize 和中心裁剪,不做随机增强。
注意:归一化的 mean 和 std 必须与预训练模型一致。用 ImageNet 预训练权重时,就用车 ImageNet 的统计值;如果换其他预训练源,需要对应调整。
4. 训练脚本怎么写:从 DataLoader 到早停的完整链路
4.1 构建 Dataset 和 DataLoader
用torchvision.datasets.ImageFolder可以直接读取按文件夹组织的分类数据:
from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder train_dataset = ImageFolder("tea_split/train", transform=train_transform) val_dataset = ImageFolder("tea_split/val", transform=val_transform) test_dataset = ImageFolder("tea_split/test", transform=val_transform) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) print(f"类别映射: {train_dataset.class_to_idx}") print(f"训练批次数: {len(train_loader)}")ImageFolder要求每个类别一个子文件夹,文件夹名即类别名。batch_size=32是 4,000 张规模下的稳妥选择,显存不够就降到 16 或 8。num_workers设为 CPU 核心数的 1/4 到 1/2,太多反而会拖慢数据加载。pin_memory=True在 GPU 训练时能加速数据传输。
如果类别不平衡,用WeightedRandomSampler替代shuffle=True:
from torch.utils.data import WeightedRandomSampler import numpy as np targets = [label for _, label in train_dataset.samples] class_counts = np.bincount(targets) class_weights = 1.0 / class_counts sample_weights = class_weights[targets] sampler = WeightedRandomSampler( weights=sample_weights, num_samples=len(sample_weights), replacement=True ) train_loader = DataLoader(train_dataset, batch_size=32, sampler=sampler, num_workers=4, pin_memory=True)class_weights是类别频率的倒数,少数类权重更高。WeightedRandomSampler每个 epoch 按权重有放回地采样,使每个 batch 的类别分布更均衡。replacement=True表示允许重复采样同一样本。
4.2 训练循环与关键超参设置
训练循环包含前向传播、loss 计算、反向传播、参数更新四个步骤。加上验证和早停:
import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR import time def train_one_epoch(model, loader, criterion, optimizer, device): model.train() running_loss, correct, total = 0.0, 0, 0 for imgs, labels in loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(imgs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() * imgs.size(0) _, preds = outputs.max(1) correct += (preds == labels).sum().item() total += labels.size(0) return running_loss / total, correct / total @torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() running_loss, correct, total = 0.0, 0, 0 for imgs, labels in loader: imgs, labels = imgs.to(device), labels.to(device) outputs = model(imgs) loss = criterion(outputs, labels) running_loss += loss.item() * imgs.size(0) _, preds = outputs.max(1) correct += (preds == labels).sum().item() total += labels.size(0) return running_loss / total, correct / total # 超参设置 num_epochs = 50 lr = 1e-3 weight_decay = 1e-4 patience = 10 # 早停耐心值 criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay) scheduler = CosineAnnealingLR(optimizer, T_max=num_epochs, eta_min=1e-6) best_val_acc = 0.0 wait = 0 for epoch in range(num_epochs): t0 = time.time() train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc = evaluate(model, val_loader, criterion, device) scheduler.step() print(f"Epoch {epoch+1}/{num_epochs} | " f"train_loss={train_loss:.4f} train_acc={train_acc:.4f} | " f"val_loss={val_loss:.4f} val_acc={val_acc:.4f} | " f"lr={scheduler.get_last_lr()[0]:.2e} | {time.time()-t0:.1f}s") if val_acc > best_val_acc: best_val_acc = val_acc torch.save(model.state_dict(), "best_model.pth") wait = 0 else: wait += 1 if wait >= patience: print(f"早停于 epoch {epoch+1},最佳验证准确率: {best_val_acc:.4f}") breakAdamW的weight_decay=1e-4是 Transformer 和 CNN 的常用值,比 SGD 更容易调。CosineAnnealingLR让学习率从1e-3余弦衰减到1e-6,避免训练后期震荡。patience=10表示验证准确率连续 10 个 epoch 不提升就停止。best_model.pth保存验证集上表现最好的权重,而不是最后一个 epoch 的权重。
4.3 学习率与 batch size 的联动调整
学习率和 batch size 不是独立的。经验公式是:batch size 翻倍,学习率也翻倍(线性缩放规则)。4,000 张图、batch_size=32 时,lr=1e-3是合理起点。如果显存不够降到 batch_size=16,学习率应降到 5e-4 左右。如果换用 SGD,学习率要设到 0.01 到 0.1 量级,并配合 momentum=0.9。
另一个常被忽略的参数是 warmup。前几个 epoch 用极小的学习率线性增加到设定值,能避免预训练权重被大梯度破坏:
from torch.optim.lr_scheduler import LinearLR, SequentialLR warmup_epochs = 3 warmup_scheduler = LinearLR(optimizer, start_factor=0.01, total_iters=warmup_epochs) cosine_scheduler = CosineAnnealingLR(optimizer, T_max=num_epochs - warmup_epochs, eta_min=1e-6) scheduler = SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[warmup_epochs])start_factor=0.01表示从设定学习率的 1% 开始,total_iters=3表示 3 个 epoch 后达到设定值。之后切换到余弦退火。这套组合在微调预训练模型时几乎不会出错。
5. 避坑与排查:4,000 张茶叶叶片图训练中的血泪经验
5.1 验证集准确率远高于测试集
现象:训练时验证集准确率 95%,测试集只有 78%。
原因:最常见的是数据泄漏——训练集和验证集/测试集里有重复或近似重复的图片。茶叶叶片采集时可能连拍,同一片叶子出现在多个子集里。另一个原因是验证集被反复用于调参,模型间接「见过」验证集。
解决:用 pHash 去重后再划分数据集,确保同一片叶子的不同角度只出现在一个子集里。划分时用GroupShuffleSplit按采集批次分组。测试集只在最终评估时用一次,中间调参只看验证集。
5.2 模型只预测多数类
现象:训练 loss 下降但准确率卡在多数类占比附近,混淆矩阵显示少数类全被预测成多数类。
原因:类别不平衡 + 交叉熵 loss 没有加权。模型发现把所有样本预测为多数类就能获得较低的 loss,于是「躺平」。
解决:用WeightedRandomSampler或给CrossEntropyLoss传weight参数。weight 设为类别频率的倒数,归一化后传入:
class_weights = torch.tensor([1.0/c for c in class_counts], dtype=torch.float32) class_weights = class_weights / class_weights.sum() * len(class_counts) criterion = nn.CrossEntropyLoss(weight=class_weights.to(device))同时观察少数类的 recall,而不是只看整体 accuracy。
5.3 训练 loss 震荡不收敛
现象:loss 曲线剧烈抖动,验证准确率忽高忽低。
原因:学习率太大、batch size 太小、或者数据增强过猛。RandomResizedCrop的 scale 下限设到 0.3 时,裁剪出的区域可能只剩背景,模型学不到有效特征。
解决:先把学习率降一个数量级试试。增强参数逐步加,不要一次性全开。RandomResizedCrop的 scale 下限建议不低于 0.5。如果用了 Mixup 或 CutMix,先关掉确认 baseline 能收敛再加。
5.4 显存溢出(OOM)
现象:训练几个 batch 后报CUDA out of memory。
原因:batch size 太大、模型参数量太大、或者没有释放中间变量。
解决:降 batch size 是最直接的。如果不想降,可以用梯度累积模拟大 batch:
accum_steps = 4 # 等效 batch_size = 32 * 4 = 128 optimizer.zero_grad() for i, (imgs, labels) in enumerate(train_loader): imgs, labels = imgs.to(device), labels.to(device) outputs = model(imgs) loss = criterion(outputs, labels) / accum_steps loss.backward() if (i + 1) % accum_steps == 0: optimizer.step() optimizer.zero_grad()另外,验证阶段用@torch.no_grad()装饰器,避免计算图占用显存。
5.5 推理时预处理不一致
现象:训练时验证准确率正常,部署推理时结果完全不对。
原因:推理时的预处理与训练时的验证预处理不一致。常见错误包括:忘了归一化、用了不同的 mean/std、resize 方式不同、通道顺序搞反(BGR vs RGB)。
解决:把验证集的 transform 单独保存,推理时严格复用。用 OpenCV 读图默认是 BGR,需要转成 RGB 再送入模型。归一化的 mean/std 必须与训练时完全一致。写一个preprocess_image函数,训练和推理共用:
def preprocess_image(img_path, transform): from PIL import Image img = Image.open(img_path).convert("RGB") return transform(img).unsqueeze(0) # 增加 batch 维度推理时用val_transform,不要用train_transform。
6. 把 4,000 张图的价值榨干:进阶技巧与验证方法
6.1 用交叉验证替代单次划分
4,000 张图做单次 7:1.5:1.5 划分,验证集只有 600 张,指标波动可能达到 ±3%。更稳妥的做法是 5 折交叉验证:把数据分成 5 份,每次用 4 份训练、1 份验证,取 5 次结果的平均值。这样能更准确地评估模型性能,也能发现某些折上表现异常差的情况——那通常意味着数据分布有问题。
from sklearn.model_selection import StratifiedKFold skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42) for fold, (train_idx, val_idx) in enumerate(skf.split(all_files, all_labels)): print(f"Fold {fold+1}: train={len(train_idx)}, val={len(val_idx)}") # 按索引构建 Dataset 和 DataLoader,训练后记录 val_acc5 折交叉验证的训练时间是单次的 5 倍,但换来的是更可靠的结论。如果算力有限,至少做 3 折。
6.2 用混淆矩阵和 t-SNE 定位模型弱点
准确率只是一个数字,混淆矩阵能告诉你模型在哪些类别上犯错:
from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for imgs, labels in test_loader: imgs = imgs.to(device) outputs = model(imgs) _, preds = outputs.max(1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_names=test_dataset.classes)) sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", xticklabels=test_dataset.classes, yticklabels=test_dataset.classes) plt.xlabel("预测") plt.ylabel("真实") plt.tight_layout() plt.savefig("confusion_matrix.png", dpi=150)如果某两个类别互相混淆严重,说明它们的视觉特征太接近。解决办法包括:增加这两类的区分性样本、用更细粒度的标注(比如标注病斑位置)、或者换用更强的 backbone。t-SNE 可视化特征空间也能直观看到类别是否可分:
from sklearn.manifold import TSNE features = [] model.eval() with torch.no_grad(): for imgs, _ in test_loader: imgs = imgs.to(device) feat = model.forward_features(imgs) # timm 模型支持 features.append(feat.cpu().numpy()) features = np.concatenate(features, axis=0) tsne = TSNE(n_components=2, perplexity=30, random_state=42) embeddings = tsne.fit_transform(features.reshape(features.shape[0], -1)) plt.scatter(embeddings[:, 0], embeddings[:, 1], c=all_labels, cmap="tab10", s=5) plt.colorbar() plt.savefig("tsne.png", dpi=150)forward_features返回分类头之前的特征向量。t-SNE 降维后,如果同类样本聚成一团、不同类分开,说明模型学到了有区分度的特征;如果混在一起,说明 backbone 需要换或增强策略需要调整。
6.3 模型导出与推理速度测试
训练完成后,导出为 ONNX 格式可以在多种推理引擎上运行:
dummy_input = torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, "tea_leaf_model.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}}, opset_version=11 ) print("ONNX 模型已导出")dynamic_axes允许变长 batch 输入。opset_version=11兼容性较好。导出后用onnxruntime测推理速度:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession("tea_leaf_model.onnx") dummy = np.random.randn(1, 3, 224, 224).astype(np.float32) for _ in range(10): # 预热 sess.run(None, {"input": dummy}) import time t0 = time.time() for _ in range(100): sess.run(None, {"input": dummy}) print(f"平均推理耗时: {(time.time()-t0)/100*1000:.1f}ms")CPU 上 ResNet-18 的推理耗时通常在 20 到 50ms 之间,EfficientNet-B0 在 15 到 30ms。如果部署到边缘设备,MobileNetV3 能压到 10ms 以内。这些数据决定了方案能不能落地到实时检测场景。
我自己的习惯是:每次拿到新数据集,先跑一遍 ResNet-18 baseline,记录准确率和混淆矩阵,然后再决定要不要上更大的模型。4,000 张茶叶叶片图,ResNet-18 配合强增强通常能到 90% 以上的验证准确率;如果低于 85%,问题多半在数据质量或划分方式上,而不是模型不够大。希望帮到你。
本文还有配套的精品资源,点击获取