☰
甘蔗病害图像分类实战:从19,000张标注数据到模型部署
2026/9/26 2:12:29 网站建设 项目流程

简介:这是一份面向计算机视觉与深度学习学习者的甘蔗植物病害图像分类数据集,覆盖红腐病、锈病、健康、枯萎病等6个类别,共约19,000张已标注图片。数据已按训练集与测试集分别存放,json文件记录了类别映射与划分信息,方便直接用于模型训练和效果评估。资源共2000个文件,包括1998张jpg图像、1个可视化脚本py和1个json配置文件,压缩包总大小约424.78MB;运行配套show脚本即可快速查看各类别样例,降低数据预处理门槛。目前已有196人学习下载,适合希望在标注数据上开展图像分类实验、对比不同网络结构或复现论文结果的初学者与研究人员。利用这份数据集,可省去采集与清洗图片的时间,直接聚焦模型设计与调优,同时还可结合作者发布的图像分类网络改进专题和计算机视觉项目内容进行深入实践。

1. 甘蔗病害图像分类:19,000张标注数据能解决什么问题

做农业视觉项目的人,最怕的不是模型选型,而是数据集根本没法用。甘蔗植物病害图像分类数据集,带标注、约19,000张,正是冲着这个痛点来的。甘蔗是我国糖料作物的主力,种植面积大,但梢腐病、赤腐病、锈病、白条病这些病害在田间表现高度相似,靠人眼巡检效率低,误判率高。一套按图像分类任务整理好的标注数据集,直接决定了你能不能快速训练出一个能上无人机或手机端的病害识别模型。

这个数据集适合谁?一类是做农业AI落地项目的工程师,需要用真实田间数据训练分类模型;一类是高校和研究所的团队,做甘蔗病害识别算法研究,但没条件自己下田采集、标注;还有一类是刚转行做图像分类的开发者,拿一套完整标注数据走通从训练到部署的流程,比自己在网上拼凑零散图片省力得多。19,000张的量级,对二分类或多分类的病害识别来说,既不会因为数据太少导致过拟合,又不会因为数据太大而让单卡训练耗时过长。

这套数据能支撑的核心任务很明确:图像分类。也就是给定一张叶片或茎秆照片,模型输出它属于健康还是某一类病害。沿这个方向,也能延伸出目标检测、语义分割的预训练需求。下面我会从数据格式、标注质量、模型选型、训练参数到部署推理,把整个落地路径拆开讲,并标出哪些地方最容易翻车。

2. 先摸清数据集底细:目录结构、标注格式与验证方法

2.1 拿到数据集后第一件事:核对目录和标签

甘蔗病害数据集的常见组织方式是按类别分文件夹,也可能带一份CSV或JSON格式的标注文件。无论哪种,第一步都是在本地把目录结构打印出来,确认图片数量与标注文件是否对得上。先跑一段脚本扫一遍:

import os from collections import Counter root = "sugarcane_disease_dataset" label_counts = Counter() image_exts = {".jpg", ".jpeg", ".png", ".bmp", ".webp"} total_images = 0 for cls_name in sorted(os.listdir(root)): cls_path = os.path.join(root, cls_name) if not os.path.isdir(cls_path): continue imgs = [f for f in os.listdir(cls_path) if os.path.splitext(f)[1].lower() in image_exts] label_counts[cls_name] = len(imgs) total_images += len(imgs) print("类别分布:", dict(label_counts)) print("图片总数:", total_images)

这段代码的核心作用是快速验证「目录结构是否完整」和「类别是否平衡」。如果发现某几个病害类别图片数只有几十张,而另一些有一两千张,那后续训练就要考虑类别权重或数据增强。19,000张如果平均分到5类,每类约3,800张,属于比较舒服的量级;但如果分布极度不均衡,A类8,000张、B类500张,那就不能直接拿来训。

还要核对每张图片能否被正常解码。实际项目中经常会遇到某些图片文件后缀是.jpg但实际是损坏文件,训练时会在数据加载阶段直接报错。可以用PIL批量验证:

from PIL import Image import os bad_images = [] for cls_name in os.listdir(root): cls_path = os.path.join(root, cls_name) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): img_path = os.path.join(cls_path, img_name) try: with Image.open(img_path) as im: im.load() except Exception as e: bad_images.append((img_path, str(e))) print("损坏图片数量:", len(bad_images)) for path, err in bad_images[:10]: print(path, err)

参数说明:im.load()会真正把图片像素数据读入内存,只调用Image.open()而不load(),遇到部分损坏文件可能不会报错。这一步建议在训练前必做,不要等到DataLoader抛异常了才回头排查。

2.2 标注格式怎么选:单标签多分类最常见

甘蔗病害图像分类数据集的标注,常见有三种形态:

第一种是「目录即标签」。这是最简单的形态,train/healthy/xxx.jpg、train/rust/xxx.jpg,PyTorch的ImageFolder和TensorFlow的image_dataset_from_directory都能直接读。第二种是「CSV/Excel索引表」,每行是图片文件名,类别ID,适合需要自定义标签映射的管理方式。第三种是「JSON/XML标注」,一般包含病害位置框或区域分割信息,这类数据虽然也能做分类,但往往是为检测或分割任务准备的。

「已标注」三个字具体指哪种,直接决定你要不要再做预处理。我建议拿到数据后先自己写一个标签映射文件,把中文类别名转成稳定的英文ID,避免后续模型训练时因为类别名编码问题翻车。常见做法是维护一个class_to_id.json:

{ "healthy": 0, "leaf_scald": 1, "rust": 2, "red_rot": 3, "pokkah_boeng": 4 }

这样做的原因是Deep Learning框架对字符串标签的兼容性参差不齐,训练脚本里统一用整数ID,推理时再用映射表转回可读名称,部署阶段会省掉很多麻烦。

2.3 划分训练集/验证集/测试集:不要随机打乱后直接split

19,000张的数据集,常规划分比例是训练集70%、验证集15%、测试集15%。但直接random.split有个隐患:如果同一株甘蔗或同一批次的照片出现在训练集和验证集,验证分数会虚高。做农业图像数据,更稳妥的做法是事先看图片文件名里有没有采集批次或地块编码,有的话按批次划分。

import random import shutil from pathlib import Path random.seed(42) src_root = Path("sugarcane_disease_dataset") out_root = Path("sugarcane_splits") for split in ["train", "val", "test"]: (out_root / split).mkdir(parents=True, exist_ok=True) for cls_dir in src_root.iterdir(): if not cls_dir.is_dir(): continue images = list(cls_dir.glob("*.jpg")) + list(cls_dir.glob("*.png")) random.shuffle(images) n = len(images) n_train, n_val = int(n * 0.7), int(n * 0.15) for img in images[:n_train]: dest = out_root / "train" / cls_dir.name / img.name dest.parent.mkdir(parents=True, exist_ok=True) shutil.copy(img, dest) for img in images[n_train:n_train + n_val]: dest = out_root / "val" / cls_dir.name / img.name dest.parent.mkdir(parents=True, exist_ok=True) shutil.copy(img, dest) for img in images[n_train + n_val:]: dest = out_root / "test" / cls_dir.name / img.name dest.parent.mkdir(parents=True, exist_ok=True) shutil.copy(img, dest) print("数据集划分完成")

这段脚本里random.seed(42)保证每次执行划分结果一致,复现实验时这是关键参数。划分完成后,建议统计每个split的类别分布,确认比例没跑偏。如果某些类别图片太少,可以考虑按类别数量加权划分。当发现某个类别只有几百张时,验证集中该类可能只有几十张,评估指标会剧烈波动,这时候需要用StratifiedSplit保证每个类别在验证集和测试集中都有足够样本。

3. 训练图像分类模型:从ResNet到ConvNeXt的选型与实践

3.1 模型选型:不是越新越好,要看你的部署环境

甘蔗病害图像分类,核心是细粒度图像分类问题。不同病害的病斑纹理、颜色、形状差异有时很明显,有时只差几个像素。选模型时,我一般按部署目标分三条线:

  • 如果跑在云服务器或实验室GPU上,优先考虑ResNet50、ResNeXt或EfficientNet。ResNet50作为baseline最稳,预训练权重好找,训练技巧成熟。
  • 如果跑在手机端或边缘设备(如无人机、田间摄像头),考虑MobileNetV3、ShuffleNetV2或轻量化的ConvNeXt。这类模型在保持精度的同时,推理延迟能控制在几十毫秒内。
  • 如果设备性能中等、但对精度要求高,EfficientNetV2和ConvNeXt-Tiny是性价比不错的选择。

这里说一个常见误区:不要一上来就选Swin Transformer或ViT。甘蔗叶片图像通常在复杂背景中拍摄,光照、泥土、叶片遮挡干扰很多。ViT类模型在小规模数据集上如果没有充分的预训练或数据增强,很容易训不过ResNet。19,000张说多不算多,说少不少,除非你打算用大规模预训练权重做微调,否则CNN是更稳妥的起点。

3.2 用ResNet50做基线训练:完整的PyTorch训练脚本

下面是一套我自己常用的训练骨架,输入是224x224的RGB图像,输出是5类病害概率。这套代码训练一个baseline,跑通后再换更强模型。

import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models from torch.optim import AdamW device = torch.device("cuda" if torch.cuda.is_available() else "cpu") transform_train = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) transform_val = 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]) ]) train_ds = datasets.ImageFolder("sugarcane_splits/train", transform=transform_train) val_ds = datasets.ImageFolder("sugarcane_splits/val", transform=transform_val) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1) model.fc = nn.Linear(model.fc.in_features, len(train_ds.classes)) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) def train_one_epoch(model, loader, optimizer, criterion): model.train() total_loss, correct, total = 0, 0, 0 for images, labels in loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) return total_loss / total, correct / total def evaluate(model, loader, criterion): model.eval() total_loss, correct, total = 0, 0, 0 with torch.no_grad(): for images, labels in loader: images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) total_loss += loss.item() * images.size(0) _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) return total_loss / total, correct / total

关键参数说明:

  • batch_size=32,如果GPU显存只有8GB可能吃不消,调成16即可。甘蔗叶图片分辨率高,ResNet50输入224x224,一个batch的中间激活值比较大。
  • lr=1e-4是迁移学习常用的微调学习率,如果从零训练可以提到1e-3。
  • RandomResizedCrop是这对数据里最有效的增强手段,它能模拟叶片不同部位被拍摄时裁剪框变化。别小看这个,农业图像里目标在画面中的位置和尺度本来就变化很大。
  • num_workers=4在Windows上如果报DataLoader worker错误,改成num_workers=0。

3.3 从ResNet50换到ConvNeXt时,哪些参数要重新调

很多人从ResNet切到ConvNeXt,只改模型名不改训练配置,结果效果反而变差,然后骂模型不行。实际上ConvNeXt的训练配方和ResNet差异很大。

ConvNeXt在19,000张这种中等规模数据集上,要特别关注以下几点:

第一,优化器建议用AdamW而不是SGD。ConvNeXt的架构设计本来就和Transformer对齐,AdamW的稳定性更好。第二,weight_decay可以比ResNet稍大,1e-4到5e-4都行。第三,学习率要从1e-4起步,但Loss下降到平台期时直接乘0.1衰减,而不是等它自己震荡收敛。第四,数据增强要多加一个RandomErasing或CutMix,因为ConvNeXt更吃数据多样性。

try: import timm model = timm.create_model("convnext_tiny", pretrained=True, num_classes=5) except ImportError as e: print("需要安装timm库: pip install timm")

timm是图像分类领域事实上的模型库,star量很高,模型实现质量比torchvision里杂七杂八的第三方实现稳定得多。用timm.create_model时,pretrained=True会加载ImageNet-1K预训练权重,num_classes=5会替换分类头。这里要提醒一个坑:有些版本的timm模型对输入尺寸有最小限制,ConvNeXt默认是224x224,如果你用了320x320输入,需要确认模型内部的patch大小是否匹配。

3.4 训练过程中的监控指标:别只看准确率

甘蔗病害分类的评估不能只盯着Top-1 Accuracy。农业场景里,漏检一个病株比误判一个健康株代价更高,所以必须额外关注Precision、Recall、F1-score,尤其是每个类别的Recall。

from sklearn.metrics import classification_report, confusion_matrix def full_evaluate(model, loader, class_names): model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for images, labels in loader: images = images.to(device) outputs = model(images) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_names=class_names)) print("Confusion Matrix:") print(confusion_matrix(all_labels, all_preds))

classification_report会输出每个类别的precision、recall、f1-score,还有macro avg和weighted avg。如果某个病害类的Recall特别低,比如赤腐病只有0.6,说明模型把大量赤腐病图片误判成了别的类,这时候不能靠盲目调学习率解决,需要回去查两类图片在视觉上到底有多像。看混淆矩阵能很直观地找到「哪两类互相打架」。

4. 数据增强与样本不平衡:让19,000张发挥出190,000张的效果

4.1 农业图像增强组合:模拟田间光照、尺度、遮挡变化

甘蔗叶片在田间的照片,光照变化极大,上午逆光、中午强光直射、傍晚偏暗,露水反光也会影响纹理。常规的随机翻转和裁剪远远不够。我自己的增强策略分三组:

第一组是光照扰动,用ColorJitter和RandomAutoContrast,把亮度扰动范围放到plusmn0.3,对比度plusmn0.2,饱和度plusmn0.2。这会强迫模型学习色温变化下的病害特征,而不是靠颜色整体偏移判断。第二组是几何扰动,除了RandomResizedCrop,还可以用RandomAffine(degrees=10, translate=(0.1, 0.1))模拟拍摄角度和位置偏移。第三组是遮挡模拟,RandomErasing(p=0.3, scale=(0.02, 0.2))会随机擦除一小块区域,模拟叶片被其他叶子遮挡的情况。

transform_train_advanced = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.7, 1.0), ratio=(0.75, 1.33)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.15), transforms.RandomAffine(degrees=10, translate=(0.1, 0.1)), transforms.ColorJitter(brightness=0.3, contrast=0.2, saturation=0.2, hue=0.05), transforms.RandomErasing(p=0.3, scale=(0.02, 0.2), ratio=(0.3, 3.3)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])

参数说明:RandomVerticalFlip(p=0.15)的概率不要设太高,因为甘蔗叶片的上下方向有一定生物学意义,倒过来的叶片会引入伪样本。RandomAffine的translate参数控制平移比例,0.1表示最多平移图像宽高的10%。RandomErasing的scale是擦除区域相对整图的面积比例,0.02到0.2是经验值,太小起不到遮挡效果,太大会把关键病斑区域整个擦掉。

4.2 类别不均衡:加权损失比复制样本更靠谱

如果某个病害类别图片特别少,最简单的办法是复制少数类图片。但复制样本容易让模型对复制样本过拟合,验证时反而露馅。更推荐两种做法:

第一种是用类别权重调整损失函数。PyTorch的CrossEntropyLoss自带weight参数,权重按总样本数 / (类别数 * 该类样本数)计算。第二张是用WeightedRandomSampler做采样的类别均衡。

from torch.utils.data import WeightedRandomSampler labels = [train_ds.targets[i] for i in range(len(train_ds.targets))] class_counts = torch.bincount(torch.tensor(labels)) weight_per_class = 1.0 / class_counts.float() sample_weights = torch.tensor([weight_per_class[label] for label in labels]) sampler = WeightedRandomSampler(sample_weights, num_samples=len(sample_weights), replacement=True) train_loader = DataLoader(train_ds, batch_size=32, sampler=sampler, num_workers=4, pin_memory=True)

WeightedRandomSampler的replacement=True意味着每个epoch采样时可以重复选同一个样本,少数类被抽到的概率会显著上升。注意,如果用了sampler,DataLoader里的shuffle参数必须设为False,否则会冲突报错或者行为诡异。这种方式的代价是每个epoch里少数类会被反复看到,所以训练epoch数要适当减少,比如从30降到20。

4.3 MixUp与CutMix:在中等数据集上提升泛化能力

在19,000张量级的数据上,MixUp和CutMix是非常有效的正则化手段。MixUp将两张图片按比例混合,标签也对应混合;CutMix则是把一张图的区域粘贴到另一张图上,标签按面积比例混合。

def mixup_data(x, y, alpha=0.2): lam = torch.distributions.Beta(alpha, alpha).sample() index = torch.randperm(x.size(0), device=x.device) mixed_x = lam * x + (1 - lam) * x[index] y_a, y_b = y, y[index] return mixed_x, y_a, y_b, lam # 在训练循环中使用 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) mixed_images, y_a, y_b, lam = mixup_data(images, labels) optimizer.zero_grad() outputs = model(mixed_images) loss = lam * criterion(outputs, y_a) + (1 - lam) * criterion(outputs, y_b) loss.backward() optimizer.step()

alpha=0.2是MixUp论文里对小数据集常用的参数,Beta(0.2, 0.2)的分布偏向0和1两端,也就是混合时大多以一张图为主。如果你的loss在MixUp下过拟合反而严重,可以试着把alpha调大,混合程度更强,正则化效果更明显。要留意的是,MixUp在验证时不能使用,推理时输入是单张图片,不存在混合问题。

4.4 数据增强的验证:对比增强前后验证集分布差异

增强策略不是越猛越好。加了过多遮挡或颜色扰动,可能导致模型学到的是「扭曲后的假甘蔗」,而不是「真实的甘蔗病害」。强烈建议做一次对照实验:一组只用基础增强,一组用全套增强,都训练相同epoch数,对比验证集上的F1-score。如果全套增强的F1反而掉了,那就说明增强力度过大,裁剪掉某个tranform后再试。

5. 训练避坑指南:甘蔗病害分类最常见的5个坑

5.1 图片加载到一半报错,训练进程直接崩溃

现象:训练跑到第2000步时,突然报OSError: image file is truncated,之前一切正常,进程白跑。

原因:数据集里包含损坏或不完整的JPEG文件,PIL在解码时才暴露问题。

解决:训练前跑一遍完整性校验脚本(前面2.1已经写过)。如果不想重跑,可以在训练时给PIL设置容错模式:

from PIL import ImageFile ImageFile.LOAD_TRUNCATED_IMAGES = True

但这只是补救,不建议作为常规方案,因为截断图片加载后可能显示为半截图,喂进模型会引入噪声样本。正确做法是把损坏图片直接剔除出数据集。

5.2 验证集准确率高,田间实拍却一塌糊涂

现象:测试集F1有0.95,但拿到新拍摄的田间照片,模型准确率骤降到0.6。

原因:训练集和测试集可能来自同一采集批次,背景、光照、拍摄设备高度相似,模型学到的是背景特征而不是病害特征。这是农业图像数据集最经典的翻车场景。

解决:第一,划分数据集时必须按采集批次分,不能按图片随机分。第二,训练时加入RandomResizedCrop降低背景占比。第三,如果条件允许,收集一批不同地区、不同光照条件的图片做外部测试集,单独评估模型的泛化能力。

5.3 训练loss下降但验证loss上升,明显过拟合

现象:训练top-1准确率很快到0.98,验证集准确率卡在0.88不再上涨,loss在epoch 15后开始反弹。

原因:模型容量相对于数据量过大。如果用的是ResNet50,19,000张其实够用但很紧,不加正则化就容易过拟合。

解决:优先加weight_decay,从1e-4提到5e-4;然后加RandomErasing增强;再不行就在全连接层前加Dropout=0.3。ResNet50的瓶颈在最后一层全连接,更换分类头时增加一层torch.nn.Dropout(0.3)会有效抑制过拟合。

5.4 标签错标:验证集里混入了错误标签

现象:训练结束后看混淆矩阵,发现「健康」和「枯梢病」互误率特别高,抽几张图片人工检查,发现有一批图片标签确实标反了。

原因:数据标注过程中,相似病斑被误标。尤其是病害早期阶段,肉眼很难区分。

解决:训练一版模型后,用模型预测验证集,把预测置信度高于0.9但预测标签与标注不一致的图片列出来,人工抽检。这些往往是标注错误的高发区。做法是输出一个error_analysis.csv,包含图片路径、真实标签、预测标签、置信度,逐条核对后修正再重新划分。

5.5 GPU显存溢出,batch_size改小后准确率浮动大

现象:RuntimeError: CUDA out of memory,把batch_size从32改成8后能跑,但准确率和之前实验差得离谱。

原因:batch_size缩小后,batch内样本多样性降低,BN层的统计量估计偏差变大,同时学习率没有等比调整。

解决:batch_size缩小为原来的四分之一,学习率也要对应缩小。常见做法是线性缩放规则:lr_new = lr_base * batch_size_new / batch_size_base。如果batch_size从32改成8,lr就从1e-4变成2.5e-5。另一个靠谱方案是用梯度累积模拟大batch:

accumulation_steps = 4 # 相当于batch_size=32 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): outputs = model(images) loss = criterion(outputs, labels) / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()

梯度累积能让小batch size下的训练效果接近大batch,但要注意loss要除以accumulation_steps,否则梯度会被放大,loss曲线会震荡。

6. 模型推理与落地部署:从训练完到田间用起来的最后一公里

6.1 ONNX导出与推理加速

训练完的PyTorch模型不能直接部署到手机或边缘设备,标准做法是导出为ONNX格式,再用ONNX Runtime或TensorRT做推理加速。导出代码如下:

import torch import onnxruntime as ort from torchvision import models device = torch.device("cpu") model = models.resnet50(weights=None) model.fc = torch.nn.Linear(model.fc.in_features, 5) checkpoint = torch.load("best_model.pt", map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"]) model.eval().to(device) dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "sugarcane_model.onnx", input_names=["input"], output_names=["probs"], dynamic_axes={"input": {0: "batch_size"}, "probs": {0: "batch_size"}}, opset_version=17 ) ort_session = ort.InferenceSession("sugarcane_model.onnx")

参数说明:dynamic_axes指定batch维度为动态,这样同一份ONNX文件可以一次推理多张图。opset_version=17需要ONNX Runtime 1.10以上版本支持。导出时如果遇到Model exported with失败,多半是model.eval()没调用,训练模式下BatchNorm参数不会固化为常数。

推理时,ONNX Runtime在CPU上的速度比PyTorch原生推理快不少。如果跑在NVIDIA Jetson这类嵌入式GPU上,还可以进一步转成TensorRT引擎,把FP32精度降到FP16,延迟能缩短一半以上。

6.2 病害概率输出与置信度阈值:把分类结果变成可决策的信息

实际部署时,分类模型输出的5个概率值不能直接拿来当决策依据。田间使用场景,我们希望模型在把握不足时输出「不确定」,而不是硬选一个类。所以我在分类头后面加一个置信度过滤逻辑:

import numpy as np def predict_with_confidence(model, image_tensor, threshold=0.7): with torch.no_grad(): logits = model(image_tensor.unsqueeze(0)) probs = torch.softmax(logits, dim=1).cpu().numpy()[0] top_idx = np.argmax(probs) top_prob = probs[top_idx] if top_prob < threshold: return -1, top_prob # -1表示不确定 return top_idx, top_prob

当最高概率低于threshold=0.7时,系统提示「请专家复核」,这样比强行给出一个错误分类更负责任。无人机的巡检流程中,低于置信度的照片会自动标记为待人工确认,而不是直接进入统计报表。这个阈值怎么调?看验证集上F1-score和误检率的关系,画一条PR曲线,选误检率可以接受的那个点,我一般从0.6试到0.85。

6.3 连续帧去抖动:视频巡检场景的后处理技巧

如果要做视频级的巡检,单帧预测会有抖动。同一片叶子在第5帧被识别为锈病,第6帧被识别为健康,这在地图标记上会很难看。一个简单有效的后处理是滑动窗口多数投票:

from collections import deque class FrameVoteFilter: def __init__(self, window_size=5): self.window = deque(maxlen=window_size) def predict(self, cls_id, prob): if prob < 0.5: return None # 置信度过低,当前帧不参与投票 self.window.append(cls_id) if len(self.window) < self.window.maxlen: return None return max(set(self.window), key=self.window.count)

窗口大小为5时,一个误判帧最多影响20%的投票权重,不会被计入最终结果。这个技巧在无人机巡逻视频里非常实用,几乎零成本实现,却能把误报率压下去一个量级。注意,如果同一片叶子连续出现5帧且都属于不同病害,说明模型本身有问题,这时候不要靠投票兜底,而是回看训练数据里该病害类别的特征表达。

6.4 一套可参考的部署拓扑

我自己的落地经验是,把完整的推理链路拆成三个环节:边缘采集端、推理服务端、数据回流端。采集端是无人机或手机,将图片压缩到长边不超过1280像素后上传;推理服务端用ONNX Runtime接收图片、预处理、推理、输出结果;回流端把「低置信度图片 + 专家修正标签」定期收回来,作为下一轮微调的训练数据,形成闭环。

19,000张的甘蔗病害数据集,训练出一个可用模型完全足够。但真实田间环境的复杂度,决定了「训练完」不是终点。我的习惯是把验证时误判的图片全部存下来,每月挑一批加入训练集重新微调一轮,用两三个月时间,模型在特定农场的准确率能再提高几个点。做农业AI,数据集的质量和迭代闭环永远比单次训练的精度更值钱。

希望上面这些步骤和踩过的坑能帮到你,哪怕只省下一两轮返工的时间,这篇笔记就没白写。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询