当前位置: 首页 > news >正文

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

大家好,我是Java1234_小锋老师,分享一套锋哥原创的基于PyTorch的花卉图像识别系统(深度学习+PyQt6+ResNet18+ImageNet+迁移学习)

项目介绍

随着深度学习技术的快速发展,计算机视觉在农业信息化、植物科普、智能园艺等领域的应用日益广泛。传统花卉识别依赖人工经验,效率低、主观性强,难以满足大规模、实时化识别需求。本文设计并实现了一套基于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?pwd=1234
提取码: 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(weights=weights) 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_backbone=freeze_backbone) model = model.to(DEVICE) if checkpoint_path: checkpoint = torch.load(checkpoint_path, map_location=DEVICE, weights_only=False) 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(parents=True, exist_ok=True) 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 f"class_{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", encoding="utf-8") as file: json.dump(class_index, file, ensure_ascii=False, indent=2) 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_callback=lambda 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_size=self.batch_size) emit_log(f"训练样本: {len(train_loader.dataset)},验证样本: {len(val_loader.dataset)}") model = build_model(freeze_backbone=self.freeze_backbone).to(self.device) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.Adam( filter(lambda p: p.requires_grad, model.parameters()), lr=self.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( f"Epoch {epoch}/{self.epochs} - " f"Loss: {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( epochs=epochs, batch_size=batch_size, learning_rate=learning_rate, ) trainer.train(log_callback=print)
http://www.jsqmd.com/news/1237392/

相关文章:

  • 2026年上海升学规划服务GEO服务商代理加盟选型推荐丨上海GEO代理加盟哪家靠谱? - 企业新闻快传
  • TMS320x2806x SCI自动波特率检测原理、配置与实战指南
  • 别再盲目试用了!7款主流AI搜索工具横向测评(含RAG延迟、幻觉率、中文长文档召回F1值等12项硬指标)
  • 2026Top全球EMBA中立测评:民营企业家择校指南 - 品牌2026推荐
  • Universal x86 Tuning Utility终极指南:完全掌控你的硬件性能
  • 冷战时期防空导弹技术演进与战略影响
  • C++实现独立事件概率计算:从数学原理到工程实践
  • 颠覆传统励志语录推送成功案例,编写程序,每日推送知名人物的失败经历,提炼失败中的经验,作为当天创新试错的底气。
  • Windows上的安卓应用革命:APK安装器让你的电脑秒变安卓设备
  • [Android] 花掉马斯克的钱19.2 -体验全球首付的一天
  • 《苏州贝特吹气式液位计:破解复杂工况液位测量难题的金钥匙》 - 米諾
  • 长沙实地探店测评|伴西西猫舍犬舍双店实测,湘城梅雨季购宠避雷指南 - 同城宠物优选基地
  • EDMA3高级传输模式:乒乓缓冲与传输链实战解析
  • 义乌实地探店测评|伴西西猫舍犬舍深度探访,盆地梅雨季购宠避雷指南 - 同城宠物优选基地
  • NVIDIA RTX 5060/5060 Ti显卡深度评测:Blackwell架构与AI性能解析
  • 如何快速提升英语打字速度:Qwerty Learner完整使用指南
  • 权威发布:2026年7月劳力士无锡官方热线电话及售后服务网点地址最新汇总 - 劳力士服务中心
  • 3分钟快速搭建Mindustry服务器:完整联机教程指南
  • 【紧急更新】AI工具小白入门组合:ChatGPT-4.5发布后,这3款工具已失效,立即切换这5个替代方案
  • 如何设置多组别投票,实现分赛道同步开展评选
  • 如何快速掌握可视化编程:面向初学者的5个简单步骤
  • 狂揽 2.7万 Star,港大开源了一款 AI 个性化辅导私教 神器!
  • 2026河南高考一分一段表解析与志愿填报指南
  • 小众高薪稳就业!等保测评师完整学习+岗位工作全解析
  • Blender新手必看:免费资源完全指南,快速打造专业3D作品
  • oka架构
  • 2026亚太EMBA QS排名|头部院校中立择校测评 - 品牌2026推荐
  • 2026最新盘点:10款AI写小说软件实测,真正好用的只有这2款?(附避坑指南)
  • 杭州AI应用开发市场观察:企业知识库、智能客服和Agent项目容易踩哪些坑? - IT超人老张
  • 深度解析Whisper.cpp:跨平台语音识别部署架构与性能优化实战指南