大家好,我是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)