Python深度学习实战:基于MobileNet的中草药植物识别系统源码解析

前言

在深度学习领域,图像分类一直是热门研究方向。本文将深入解析一个基于 MobileNet 架构的中草药植物识别系统,该系统能够识别 67种 常见中草药植物,并提供了完整的 PyQt5 图形界面。本文将从项目架构、核心算法实现、模型设计等多个维度进行源码级别的深度剖析,适合有一定深度学习基础的开发者学习参考。

关键词:Python、毕业设计、深度学习、MobileNet、PyTorch、图像分类、中草药识别


一、项目整体架构分析

1.1 目录结构设计

项目采用了清晰的分层架构设计,主要目录结构如下:

mobile_net_plant_02/
├── models/              # 模型定义模块
│   └── mobilenet.py     # MobileNet网络结构实现
├── all_data/            # 数据集目录(67个植物类别)
├── weights/             # 训练好的模型权重
│   └── plant-best-epoch.pth
├── ui/                  # UI资源文件
├── font/                # 中文字体文件
├── train.py             # 模型训练脚本
├── predict.py           # 单张图片预测脚本
├── 主界面.py            # PyQt5主界面程序
├── my_dataset.py        # 自定义数据集类
├── utils.py             # 工具函数集合
└── class_indices.json   # 类别索引映射文件

这种模块化的设计使得代码职责清晰,便于维护和扩展。

在这里插入图片描述

1.2 技术栈选型

  • 深度学习框架:PyTorch(灵活的动态图机制,便于调试)
  • GUI框架:PyQt5(跨平台桌面应用开发)
  • 图像处理:PIL/Pillow、torchvision
  • 数据处理:NumPy、Pandas
  • 可视化:Matplotlib

二、MobileNet网络架构深度解析

2.1 深度可分离卷积(Depthwise Separable Convolution)

MobileNet的核心创新在于使用深度可分离卷积替代标准卷积,大幅减少参数量和计算量。让我们看看源码实现:

class DepthSeperabelConv2d(nn.Module):
    """深度可分离卷积实现"""
    def __init__(self, input_channels, output_channels, kernel_size, **kwargs):
        super().__init__()
        # 第一步:深度卷积(Depthwise Convolution)
        # 每个输入通道独立进行卷积,groups=input_channels实现通道分离
        self.depthwise = nn.Sequential(
            nn.Conv2d(
                input_channels,
                input_channels,  # 输出通道数等于输入通道数
                kernel_size,
                groups=input_channels,  # 关键:分组数等于输入通道数
                **kwargs),
            nn.BatchNorm2d(input_channels),
            nn.ReLU(inplace=True)
        )
        
        # 第二步:逐点卷积(Pointwise Convolution)
        # 1x1卷积进行通道融合
        self.pointwise = nn.Sequential(
            nn.Conv2d(input_channels, output_channels, 1),  # 1x1卷积
            nn.BatchNorm2d(output_channels),
            nn.ReLU(inplace=True)
        )
    
    def forward(self, x):
        x = self.depthwise(x)   # 先进行深度卷积
        x = self.pointwise(x)    # 再进行逐点卷积
        return x

算法原理

  • 标准卷积:参数量 = kernel_size² × input_channels × output_channels
  • 深度可分离卷积:参数量 = kernel_size² × input_channels + input_channels × output_channels
  • 参数量减少比例:约为 1/output_channels + 1/kernel_size²,通常可减少 8-9倍 参数量

[插入图片:深度可分离卷积示意图]

2.2 MobileNet完整网络结构

class MobileNet(nn.Module):
    def __init__(self, width_multiplier=1, class_num=100):
        super().__init__()
        alpha = width_multiplier  # 宽度乘子,用于控制网络宽度
        
        # Stem层:初始特征提取
        self.stem = nn.Sequential(
            BasicConv2d(3, int(32 * alpha), 3, padding=1, bias=False),
            DepthSeperabelConv2d(int(32 * alpha), int(64 * alpha), 3, 
                                 padding=1, bias=False)
        )
        
        # 四个下采样阶段,逐步降低特征图尺寸,增加通道数
        self.conv1 = nn.Sequential(...)  # 128通道
        self.conv2 = nn.Sequential(...)  # 256通道
        self.conv3 = nn.Sequential(...)  # 512通道(6个深度可分离卷积块)
        self.conv4 = nn.Sequential(...)  # 1024通道
        
        # 全局平均池化 + 全连接层
        self.avg = nn.AdaptiveAvgPool2d(1)  # 自适应平均池化到1x1
        self.fc = nn.Linear(int(1024 * alpha), class_num)  # 分类头
    
    def forward(self, x):
        x = self.stem(x)
        x = self.conv1(x)  # 特征图尺寸减半
        x = self.conv2(x)  # 特征图尺寸再减半
        x = self.conv3(x)  # 特征图尺寸再减半
        x = self.conv4(x)  # 特征图尺寸再减半
        x = self.avg(x)    # [B, 1024, 1, 1]
        x = x.view(x.size(0), -1)  # 展平:[B, 1024]
        x = self.fc(x)     # [B, class_num]
        return x

网络设计亮点

  1. 渐进式下采样:通过stride=2的卷积逐步降低特征图尺寸(224→112→56→28→14→7)
  2. 通道数递增:64→128→256→512→1024,符合特征提取的层次化设计
  3. 自适应池化AdaptiveAvgPool2d(1)确保无论输入尺寸如何,都能输出固定维度特征

三、数据集处理与数据增强策略

3.1 自定义数据集类实现

class MyDataSet(Dataset):
    """自定义数据集类,继承PyTorch的Dataset基类"""
    def __init__(self, images_path: list, images_class: list, transform=None):
        self.images_path = images_path      # 图片路径列表
        self.images_class = images_class    # 对应的类别标签列表
        self.transform = transform          # 数据增强变换
    
    def __len__(self):
        return len(self.images_path)  # 返回数据集大小
    
    def __getitem__(self, item):
        img = Image.open(self.images_path[item])
        # 确保图片为RGB格式,避免灰度图或RGBA格式导致的问题
        if img.mode != 'RGB':
            raise ValueError("image: {} isn't RGB mode.".format(
                self.images_path[item]))
        label = self.images_class[item]
        
        # 应用数据增强(训练时)或预处理(验证时)
        if self.transform is not None:
            img = self.transform(img)
        
        return img, label
    
    @staticmethod
    def collate_fn(batch):
        """自定义批处理函数,将多个样本打包成batch"""
        images, labels = tuple(zip(*batch))
        images = torch.stack(images, dim=0)  # [B, C, H, W]
        labels = torch.as_tensor(labels)      # [B]
        return images, labels

设计要点

  • __getitem__方法实现懒加载,只在需要时读取图片,节省内存
  • collate_fn静态方法统一处理batch数据,确保维度正确

3.2 数据增强策略

# 训练集数据增强:增强模型泛化能力
data_transform = {
    "train": transforms.Compose([
        transforms.RandomResizedCrop(224),      # 随机裁剪并缩放到224x224
        transforms.RandomHorizontalFlip(),     # 随机水平翻转(50%概率)
        transforms.ToTensor(),                  # 转换为Tensor并归一化到[0,1]
        transforms.Normalize([0.485, 0.456, 0.406],  # ImageNet均值
                           [0.229, 0.224, 0.225])   # ImageNet标准差
    ]),
    "val": transforms.Compose([
        transforms.Resize(int(224 * 1.143)),   # 先放大到256
        transforms.CenterCrop(224),            # 中心裁剪到224
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406],
                           [0.229, 0.224, 0.225])
    ])
}

数据增强的作用

  • RandomResizedCrop:模拟不同视角和尺度,提高模型鲁棒性
  • RandomHorizontalFlip:增加数据多样性,对植物识别特别有效(左右对称)
  • Normalize:使用ImageNet预训练统计值,有助于模型收敛

四、训练流程核心代码解析

4.1 训练循环实现

def train_one_epoch(model, optimizer, data_loader, device, epoch):
    """单轮训练函数"""
    model.train()  # 设置为训练模式(启用Dropout、BatchNorm等)
    loss_function = torch.nn.CrossEntropyLoss()  # 交叉熵损失函数
    accu_loss = torch.zeros(1).to(device)  # 累计损失
    accu_num = torch.zeros(1).to(device)   # 累计正确预测数
    
    optimizer.zero_grad()  # 清零梯度
    sample_num = 0
    
    # 使用tqdm显示训练进度条
    data_loader = tqdm(data_loader, file=sys.stdout)
    for step, data in enumerate(data_loader):
        images, labels = data
        sample_num += images.shape[0]  # 累计样本数
        
        # 前向传播
        pred = model(images.to(device))  # [B, num_classes]
        pred_classes = torch.max(pred, dim=1)[1]  # 获取预测类别
        accu_num += torch.eq(pred_classes, labels.to(device)).sum()  # 统计正确数
        
        # 计算损失
        loss = loss_function(pred, labels.to(device))
        loss.backward()  # 反向传播,计算梯度
        accu_loss += loss.detach()  # 累计损失(detach避免梯度追踪)
        
        # 更新参数
        optimizer.step()  # 根据梯度更新模型参数
        optimizer.zero_grad()  # 清零梯度,准备下一轮
        
        # 更新进度条显示
        data_loader.desc = "[train epoch {}] loss: {:.3f}, acc: {:.3f}".format(
            epoch, accu_loss.item() / (step + 1), accu_num.item() / sample_num)
    
    return accu_loss.item() / (step + 1), accu_num.item() / sample_num

关键点解析

  1. model.train():启用训练模式,BatchNorm使用当前batch统计,Dropout随机丢弃神经元
  2. 梯度清零:每次迭代前必须清零,否则梯度会累积
  3. loss.detach():分离计算图,避免内存泄漏

4.2 验证流程实现

@torch.no_grad()  # 装饰器:禁用梯度计算,节省内存和计算
def evaluate(model, data_loader, device, epoch):
    """验证函数"""
    loss_function = torch.nn.CrossEntropyLoss()
    model.eval()  # 设置为评估模式(禁用Dropout,BatchNorm使用全局统计)
    
    accu_num = torch.zeros(1).to(device)
    accu_loss = torch.zeros(1).to(device)
    sample_num = 0
    
    for step, data in enumerate(data_loader):
        images, labels = data
        sample_num += images.shape[0]
        
        pred = model(images.to(device))
        pred_classes = torch.max(pred, dim=1)[1]
        accu_num += torch.eq(pred_classes, labels.to(device)).sum()
        
        loss = loss_function(pred, labels.to(device))
        accu_loss += loss
    
    return accu_loss.item() / (step + 1), accu_num.item() / sample_num

@torch.no_grad()的作用

  • 禁用自动求导,大幅减少内存占用(验证时不需要计算梯度)
  • 提升推理速度约 20-30%

4.3 最佳模型保存策略

# 在训练主循环中
if val_acc == max(val_acc_list):
    print('save-best-epoch:{}'.format(epoch))
    # 保存最佳验证准确率对应的模型
    torch.save(model.state_dict(), "./weights/plant-best-epoch.pth")

保存策略:只保存验证集上表现最好的模型,避免过拟合。


五、PyQt5界面设计与预测流程

5.1 模型加载与初始化

class MainWindow(QtWidgets.QMainWindow, Ui_MainWindow):
    def __init__(self, parent=None):
        super(MainWindow, self).__init__(parent)
        self.setupUi(self)
        
        # 设置计算设备(优先使用GPU)
        self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
        
        # 加载类别索引映射文件
        json_path = './class_indices.json'
        json_file = open(json_path, "r")
        self.class_indict = json.load(json_file)
        
        # 创建模型并加载权重
        self.model = create_model(class_num=67).to(self.device)
        model_weight_path = "weights/plant-best-epoch.pth"
        self.model.load_state_dict(torch.load(model_weight_path, 
                                             map_location=self.device))
        self.model.eval()  # 设置为评估模式

5.2 单张图片预测实现

def img_detect(self):
    """单张图片检测函数"""
    img_size = 224
    # 定义与训练时一致的预处理流程
    data_transform = transforms.Compose([
        transforms.Resize(int(img_size * 1.143)),  # 256
        transforms.CenterCrop(img_size),          # 224
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], 
                           [0.229, 0.224, 0.225])
    ])
    
    img = Image.open(self.img_path)
    img = data_transform(img)  # 预处理
    img = torch.unsqueeze(img, dim=0)  # 增加batch维度:[1, C, H, W]
    
    # 模型推理
    with torch.no_grad():  # 禁用梯度计算
        output = torch.squeeze(self.model(img.to(self.device))).cpu()
        predict = torch.softmax(output, dim=0)  # 转换为概率分布
        predict_cla = torch.argmax(predict).numpy()  # 获取最大概率类别
    
    # 获取预测结果和置信度
    res = self.class_indict[str(list(predict.numpy()).index(max(predict.numpy())))]
    confidence = "%.2f" % (max(predict.numpy()) * 100) + "%"
    
    # 更新UI显示
    self.label_res.setText(name_list[int(res.split('_')[-1])])
    self.label_pro.setText(confidence)

关键步骤

  1. 预处理一致性:必须与训练时使用相同的预处理流程
  2. torch.unsqueeze:添加batch维度,因为模型期望输入shape为[B, C, H, W]
  3. torch.softmax:将logits转换为概率,便于理解置信度

在这里插入图片描述

5.3 批量预测优化策略

def detect_folder_timer(self):
    """使用定时器实现批量预测,避免界面卡顿"""
    if self.folder_num == len(self.folder_list):
        self.timer_folder.stop()  # 所有图片处理完成
    else:
        img_path = self.folder_list[self.folder_num]
        # 处理单张图片
        res, pro = self.img_detect2(img_path)
        # 更新表格显示
        for column, data in enumerate([str(self.folder_num + 1), 
                                       img_path, res, pro]):
            self.tableWidget.setItem(self.folder_num, column, 
                                   QtWidgets.QTableWidgetItem(str(data)))
        self.folder_num += 1

优化技巧

  • 使用QTimer分时处理,每30ms处理一张图片
  • 避免一次性处理所有图片导致界面冻结
  • 实时更新UI,提供良好的用户体验

六、关键技术难点与解决方案

6.1 数据集类别不平衡问题

utils.py的数据读取函数可以看到,项目通过随机采样划分训练集和验证集:

val_path = random.sample(images, k=int(len(images) * val_rate))

潜在问题:如果某些类别样本数过少,可能导致验证集代表性不足。

改进建议

  • 使用分层采样(Stratified Sampling)确保每个类别在验证集中都有代表
  • 对少数类别进行数据增强(旋转、颜色抖动等)

6.2 模型推理性能优化

当前实现每次预测都要加载图片、预处理、模型前向传播,对于批量预测可以进一步优化:

# 优化建议:批量推理
def batch_predict(self, img_paths, batch_size=8):
    """批量预测,提升效率"""
    imgs = []
    for path in img_paths:
        img = Image.open(path)
        img = self.data_transform(img)
        imgs.append(img)
    
    # 堆叠成batch
    batch = torch.stack(imgs, dim=0)  # [B, C, H, W]
    
    with torch.no_grad():
        outputs = self.model(batch.to(self.device))  # 一次前向传播
        predicts = torch.softmax(outputs, dim=1)
    
    return predicts

性能提升:批量推理比单张推理快 3-5倍(充分利用GPU并行计算能力)

6.3 中文显示问题解决

项目在utils.py中使用了自定义字体解决matplotlib中文显示问题:

font_path = 'font/simsun.ttc'
font = matplotlib.font_manager.FontProperties(fname=font_path)
plt.xticks(range(len(flower_class)), 
          [name_list[int(i.split('_')[-1])] for i in flower_class],
          rotation=45, fontproperties=font)

这是处理matplotlib中文乱码的标准做法。


七、总结与展望

本文深入解析了基于MobileNet的中草药植物识别系统的核心源码实现。通过分析深度可分离卷积、数据增强、训练流程等关键环节,我们可以学习到:

  1. MobileNet架构的精妙设计:通过深度可分离卷积大幅减少参数量
  2. PyTorch训练的标准流程:数据加载、前向传播、反向传播、参数更新
  3. GUI应用的开发技巧:使用定时器避免界面卡顿,提升用户体验

技术拓展方向

  • 可以尝试MobileNetV2/V3等改进版本
  • 引入注意力机制(如SE-Net)提升识别准确率
  • 使用知识蒸馏技术进一步压缩模型

全套项目获取连接

本文从源码层面深度剖析了植物识别系统的实现细节,希望能帮助正在学习深度学习的同学更好地理解项目架构和算法原理。

完整项目源码、训练权重文件、详细报告及使用文档,可下方链接获取。项目包含完整的训练代码、预训练模型、数据集处理脚本等,适合作为毕业设计参考或深度学习实战项目学习。

项目链接:https://my.feishu.cn/wiki/Psatwgu0Qic4K3kDc8hc276ynXd?from=from_copylink

更多推荐