分享一套锋哥原创的基于PyTorch的花卉图像识别系统(深度学习+PyQt6+ResNet18+ImageNet+迁移学习)

发布时间:2026/7/21 22:28:20
分享一套锋哥原创的基于PyTorch的花卉图像识别系统(深度学习+PyQt6+ResNet18+ImageNet+迁移学习) 大家好我是Java1234_小锋老师分享一套锋哥原创的基于PyTorch的花卉图像识别系统(深度学习PyQt6ResNet18ImageNet迁移学习)项目介绍随着深度学习技术的快速发展计算机视觉在农业信息化、植物科普、智能园艺等领域的应用日益广泛。传统花卉识别依赖人工经验效率低、主观性强难以满足大规模、实时化识别需求。本文设计并实现了一套基于PyTorch的花卉图像识别系统以Oxford Flowers102数据集为基础采用ResNet18卷积神经网络进行迁移学习构建面向102类花卉的图像分类模型并使用PyQt6开发桌面图形界面完成模型训练、图像识别、结果可视化等核心功能。系统在技术路线上重点结合了Python语言生态、ImageNet预训练知识与ResNet18残差网络结构。首先利用torchvision自动下载并管理Flowers102数据其次加载在ImageNet上预训练的ResNet18权重冻结卷积骨干网络仅替换并训练输出维度为102的全连接分类层以降低CPU环境下的训练成本最后通过Softmax概率输出与Top-5排序向用户展示中英文花卉名称及置信度。界面端支持训练超参数配置、后台线程训练、Loss/Accuracy曲线实时绘制以及单张图片识别预览。测试结果表明系统能够完整跑通“数据准备—模型训练—图像识别”流程界面交互清晰模块划分合理具备较好的可扩展性与教学演示价值。本文工作为本科毕业设计层面的深度学习应用提供了一套可落地的参考实现也可作为后续移动端部署、细粒度分类增强与多模型对比研究的基础平台。源码下载链接: https://pan.baidu.com/s/1MCPOPLcOAiiyRzxFBaT9bA?pwd1234提取码: 1234系统展示核心代码 模型定义模块 基于 ResNet18 的迁移学习花卉分类模型 from typing import Optional import torch import torch.nn as nn from torchvision.models import ResNet18_Weights, resnet18 from src.config import DEVICE, NUM_CLASSES def build_model(freeze_backbone: bool True) - nn.Module: 构建 ResNet18 迁移学习模型 加载 ImageNet 预训练权重替换全连接层为 102 类输出。 默认冻结卷积骨干仅训练全连接层适合 CPU 训练。 Args: freeze_backbone: 是否冻结骨干网络参数 Returns: 构建好的 ResNet18 模型 weights ResNet18_Weights.DEFAULT model resnet18(weightsweights) if freeze_backbone: for param in model.parameters(): param.requires_grad False in_features model.fc.in_features model.fc nn.Linear(in_features, NUM_CLASSES) if freeze_backbone: for param in model.fc.parameters(): param.requires_grad True return model def load_model( checkpoint_path: Optional[str] None, freeze_backbone: bool True, ) - nn.Module: 加载模型可选从检查点恢复权重 Args: checkpoint_path: 权重文件路径None 则仅加载预训练骨干 freeze_backbone: 是否冻结骨干 Returns: 加载权重后的模型 model build_model(freeze_backbonefreeze_backbone) model model.to(DEVICE) if checkpoint_path: checkpoint torch.load(checkpoint_path, map_locationDEVICE, weights_onlyFalse) if isinstance(checkpoint, dict) and model_state_dict in checkpoint: model.load_state_dict(checkpoint[model_state_dict]) else: model.load_state_dict(checkpoint) model.eval() return model 模型训练模块 提供 CPU 训练循环与进度回调 import json from datetime import datetime from pathlib import Path from typing import Callable, Dict, List, Optional import torch import torch.nn as nn from torch.utils.data import DataLoader from src.config import CLASS_INDEX_PATH, DEVICE, MODEL_DIR, MODEL_PATH, NUM_CLASSES from src.dataset import build_dataloaders, ensure_dataset_downloaded from src.flower_names import FLOWER_NAMES_CN, FLOWER_NAMES_EN, get_display_name from src.model import build_model class FlowerTrainer: 花卉识别模型训练器 支持进度回调供命令行与 PyQt6 界面共用 def __init__( self, epochs: int 10, batch_size: int 16, learning_rate: float 0.001, freeze_backbone: bool True, ): 初始化训练器 Args: epochs: 训练轮数 batch_size: 批大小 learning_rate: 学习率 freeze_backbone: 是否冻结 ResNet18 骨干 self.epochs epochs self.batch_size batch_size self.learning_rate learning_rate self.freeze_backbone freeze_backbone self.device torch.device(DEVICE) self.train_losses: List[float] [] self.val_accuracies: List[float] [] self._stop_requested False def request_stop(self) - None: 请求停止训练 self._stop_requested True staticmethod def _format_time() - str: 格式化当前时间 Returns: 形如 2026-11-02 17:25:17 的时间字符串 return datetime.now().strftime(%Y-%m-%d %H:%M:%S) def _evaluate(self, model: nn.Module, val_loader: DataLoader) - float: 在验证集上评估准确率 Args: model: 待评估模型 val_loader: 验证 DataLoader Returns: 验证集准确率0-1 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images images.to(self.device) labels labels.to(self.device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return correct / total if total 0 else 0.0 def _save_checkpoint(self, model: nn.Module, val_acc: float) - None: 保存模型权重与类别索引 Args: model: 训练完成的模型 val_acc: 最终验证准确率 MODEL_DIR.mkdir(parentsTrue, exist_okTrue) checkpoint { model_state_dict: model.state_dict(), num_classes: NUM_CLASSES, val_accuracy: val_acc, saved_at: self._format_time(), } torch.save(checkpoint, MODEL_PATH) class_index { str(i): { en: FLOWER_NAMES_EN[i] if i len(FLOWER_NAMES_EN) else fclass_{i}, cn: FLOWER_NAMES_CN[i] if i len(FLOWER_NAMES_CN) else f类别{i}, display: get_display_name(i), } for i in range(NUM_CLASSES) } with open(CLASS_INDEX_PATH, w, encodingutf-8) as file: json.dump(class_index, file, ensure_asciiFalse, indent2) def train( self, progress_callback: Optional[Callable[[Dict], None]] None, log_callback: Optional[Callable[[str], None]] None, ) - Dict: 执行完整训练流程 Args: progress_callback: 进度回调接收 epoch、loss、acc 等字典 log_callback: 日志回调接收带时间戳的日志字符串 Returns: 训练结果摘要字典 self._stop_requested False self.train_losses.clear() self.val_accuracies.clear() def emit_log(message: str) - None: 输出带时间戳的日志 line f[{self._format_time()}] {message} if log_callback: log_callback(line) emit_log(开始检查/下载 Flowers102 数据集...) if not ensure_dataset_downloaded( progress_callbacklambda percent, msg: emit_log(f[数据集 {percent}%] {msg}) ): emit_log(数据集下载失败请检查网络连接后重试。) return { epochs_done: 0, final_loss: 0.0, final_acc: 0.0, model_path: , stopped: True, error: dataset_download_failed, } emit_log(数据集就绪正在构建 DataLoader...) train_loader, val_loader build_dataloaders(batch_sizeself.batch_size) emit_log(f训练样本: {len(train_loader.dataset)}验证样本: {len(val_loader.dataset)}) model build_model(freeze_backboneself.freeze_backbone).to(self.device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam( filter(lambda p: p.requires_grad, model.parameters()), lrself.learning_rate, ) emit_log(模型初始化完成开始 CPU 训练...) for epoch in range(1, self.epochs 1): if self._stop_requested: emit_log(收到停止请求训练已中断。) break model.train() running_loss 0.0 batch_count 0 for images, labels in train_loader: if self._stop_requested: break images images.to(self.device) labels labels.to(self.device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() batch_count 1 avg_loss running_loss / max(batch_count, 1) val_acc self._evaluate(model, val_loader) self.train_losses.append(avg_loss) self.val_accuracies.append(val_acc) emit_log( fEpoch {epoch}/{self.epochs} - fLoss: {avg_loss:.4f}, Val Acc: {val_acc * 100:.2f}% ) if progress_callback: progress_callback({ epoch: epoch, total_epochs: self.epochs, loss: avg_loss, val_acc: val_acc, train_losses: list(self.train_losses), val_accuracies: list(self.val_accuracies), progress: int(epoch / self.epochs * 100), }) final_acc self.val_accuracies[-1] if self.val_accuracies else 0.0 if not self._stop_requested: self._save_checkpoint(model, final_acc) emit_log(f训练完成模型已保存至 {MODEL_PATH}) emit_log(f最终验证准确率: {final_acc * 100:.2f}%) return { epochs_done: len(self.train_losses), final_loss: self.train_losses[-1] if self.train_losses else 0.0, final_acc: final_acc, model_path: str(MODEL_PATH), stopped: self._stop_requested, } def train_from_cli( epochs: int 10, batch_size: int 16, learning_rate: float 0.001, ) - None: 命令行训练入口函数 Args: epochs: 训练轮数 batch_size: 批大小 learning_rate: 学习率 trainer FlowerTrainer( epochsepochs, batch_sizebatch_size, learning_ratelearning_rate, ) trainer.train(log_callbackprint)