ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

MobileNetV2优化实现中草药识别系统

MobileNetV2优化实现中草药识别系统 1. 项目概述当MobileNet遇上中草药识别去年帮学弟调试这个中草药识别系统时我们测试了7种不同的轻量级网络最终MobileNetV2以87.3%的准确率和23ms的单图预测速度胜出。这个基于PyTorchPyQt5的解决方案完美平衡了算法精度和工程落地需求特别适合需要本地化部署的毕业设计场景。系统核心是67类常见中草药的识别包含当归、黄芪等典型药材。与常规植物识别不同中草药存在以下特征叶片纹理相似度高如薄荷与留兰香干燥药材颜色失真拍摄角度差异大针对这些痛点我们在标准MobileNet基础上做了三点改进在最后一个卷积层后增加SE注意力模块采用Focal Loss解决类别不平衡添加旋转/色彩抖动的数据增强2. 环境搭建避坑指南2.1 PyTorch环境配置推荐使用conda创建虚拟环境conda create -n herb python3.8 conda install pytorch1.12.1 torchvision0.13.1 cudatoolkit11.3 -c pytorch重要提示如果使用30系显卡必须搭配CUDA11.x以上版本否则会报CUDA error: no kernel image is available错误2.2 PyQt5安装常见问题遇到Could not find a version that satisfies the requirement PyQt5错误时改用以下命令pip install --pre pyqt5 -f https://www.riverbankcomputing.com/pyqt/download实测在Windows 11Python3.8环境下需要额外安装pip install pyqt5-tools3. 数据准备与增强策略3.1 数据集构建我们使用的67类中草药数据集包含每类300-500张原始图像包含叶片特写、整体植株、干燥药材三种形态分辨率统一调整为224x224数据目录结构示例herb_dataset/ ├── angelica/ │ ├── 001.jpg │ ├── 002.jpg ├── astragalus/ │ ├── 001.jpg ...3.2 数据增强方案在torchvision.transforms中配置train_transform transforms.Compose([ transforms.RandomRotation(30), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])经验之谈中草药识别最关键的是颜色保真度建议将ColorJitter的hue参数设为0避免色相变化影响识别4. MobileNet模型优化实战4.1 网络结构调整原始MobileNetV2的修改点class HerbMobileNet(nn.Module): def __init__(self, num_classes67): super().__init__() self.base mobilenet_v2(pretrainedTrue) # 修改最后一层 self.base.classifier[1] nn.Linear(1280, num_classes) # 添加SE模块 self.se nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(1280, 1280//16, 1), nn.ReLU(), nn.Conv2d(1280//16, 1280, 1), nn.Sigmoid() ) def forward(self, x): x self.base.features(x) se self.se(x) x x * se # 注意力加权 x nn.functional.adaptive_avg_pool2d(x, (1, 1)) return self.base.classifier(x.flatten(1))4.2 训练参数配置关键训练参数optimizer torch.optim.Adam(model.parameters(), lr0.001, weight_decay1e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size5, gamma0.1) criterion FocalLoss(gamma2, alpha0.25) # 替代交叉熵Focal Loss实现class FocalLoss(nn.Module): def __init__(self, gamma2, alpha0.25): super().__init__() self.gamma gamma self.alpha alpha def forward(self, inputs, targets): BCE_loss F.cross_entropy(inputs, targets, reductionnone) pt torch.exp(-BCE_loss) loss self.alpha * (1-pt)**self.gamma * BCE_loss return loss.mean()5. PyQt5界面开发技巧5.1 界面布局设计使用Qt Designer创建的主界面包含图像显示区域QLabel结果展示表格QTableWidget模型加载进度条QProgressBar批量识别按钮QPushButton关键代码片段class MainWindow(QMainWindow): def __init__(self): super().__init__() self.model None self.initUI() def initUI(self): # 中央组件 self.image_label QLabel() self.image_label.setAlignment(Qt.AlignCenter) self.image_label.setStyleSheet(border: 2px dashed #aaa;) # 结果表格 self.result_table QTableWidget() self.result_table.setColumnCount(2) self.result_table.setHorizontalHeaderLabels([药材名称, 置信度]) # 布局设置 central_widget QWidget() layout QHBoxLayout() layout.addWidget(self.image_label, 60) layout.addWidget(self.result_table, 40) central_widget.setLayout(layout) self.setCentralWidget(central_widget)5.2 模型加载优化采用多线程防止界面卡死class LoadModelThread(QThread): finished pyqtSignal(bool) def __init__(self, model_path): super().__init__() self.model_path model_path def run(self): try: model torch.load(self.model_path) model.eval() self.finished.emit(True) except Exception as e: print(f加载失败: {str(e)}) self.finished.emit(False)6. 模型部署与性能优化6.1 ONNX转换将PyTorch模型转为ONNX格式提升推理速度dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export(model, dummy_input, herb_mobilenet.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}})6.2 TensorRT加速在支持NVIDIA GPU的设备上trtexec --onnxherb_mobilenet.onnx --saveEngineherb_mobilenet.trt --fp16实测性能对比设备原始PyTorchONNX RuntimeTensorRTRTX 306018ms12ms8msJetson Nano230ms180ms95ms7. 常见问题解决方案7.1 内存泄漏排查在PyQt5中频繁加载图像可能导致内存泄漏正确做法def load_image(self, path): # 先清除原有QPixmap self.image_label.clear() pixmap QPixmap(path) self.image_label.setPixmap(pixmap.scaled( self.image_label.size(), Qt.KeepAspectRatio, Qt.SmoothTransformation ))7.2 模型推理异常当出现RuntimeError: expected scalar type Float but found Byte错误时# 在图像预处理时确保转为float32 image Image.open(path).convert(RGB) image_tensor train_transform(image).unsqueeze(0).float() # 关键.float()8. 项目扩展方向多模态识别加入药材气味传感器数据三维特征提取对药材进行3D扫描建模移动端部署使用TorchScript转换模型开发Android应用知识图谱整合关联药材功效、用法禁忌等信息这个项目最让我惊喜的是MobileNet在轻量级场景的潜力——经过适当调整它在保持高效率的同时对细粒度特征的捕捉能力不输ResNet等重型网络。特别是在Jetson Nano这类边缘设备上经过TensorRT加速后完全可以实现实时识别这对中医药数字化是很有意义的实践。
返回列表