十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

PyTorch花卉识别实战:小样本迁移学习与FastAPI部署全流程

PyTorch花卉识别实战:小样本迁移学习与FastAPI部署全流程 简介面向机器学习和计算机视觉初学者的花卉图像分类工具基于Python实现17种花卉的自动识别。项目涵盖数据预处理、特征提取、CNN模型构建、训练评估与部署等完整流程适合希望掌握图像分类实战方法的开发者参考学习。压缩包共2755个文件以2720张花卉JPG图像数据集为主体另有少量PNG图片、NumPy格式npy文件、Python源码及license/README等文档整体约251.53MB。目前已有927人学习浏览。资源内包含可直接运行的.py训练脚本、17类花卉图像数据、预处理与增强策略以及citations.bib和tex等引用文件便于理解从数据加载到模型推理的完整链路也可作为课程设计或算法对比的基线项目。1. 这个花卉识别项目到底在做什么拿到一份Flower-Recognition-master的资源包别急着去翻.gitignore或者.bib文件——先看图片命名image_0612.jpg、image_0398.jpg这种毫规律可言的编号其实已经暴露了项目的数据组织方式。按摘要里的描述这个任务有 17 个花卉类别每类 80 张训练图总共 1360 张样本。这个数据量放在深度学习里属于典型的「少量样本但类别均衡」场景你直接用随机初始化的 CNN 从头训练几乎必然过拟合验证集能过 70% 都算运气好。所以这个项目真正值得拆解的点不在「识别花卉」本身而在于它展示了小数据集图像分类的标准工程路径用预训练卷积网络VGG16、ResNet、InceptionV3 这类 ImageNet 模型冻结卷积基只训练顶部的全连接层完成特征迁移。对于 5 年以上经验的工程师重点不是跑通代码而是理解为什么这种方案在小数据场景下比端到端训练更稳、以及哪些环节容易翻车。这篇文章会用 PyTorch 风格代码把整条链路重新走一遍覆盖数据预处理、特征提取、训练调优到部署的完整闭环。2. 数据预处理与 PyTorch Dataset 封装先解决 ImageFolder 的坑2.1 为什么摘要里的数据组织形式直接决定了代码写法摘要虽然没有给出目录树但从常见的 Git 仓库结构推测大概率是train/类别名/图片.jpg或者17 个类别子文件夹的组织方式。这种布局天然适合torchvision.datasets.ImageFolder直接加载。但这里有一个坑ImageFolder 依赖文件夹命名顺序来生成类别索引roses和daisy的排序是按字符串 ASCII 码来的不是按文件顺序。如果你后面做混淆矩阵或者手动测试单张图索引对应关系搞错一切评估都失真。我一般会在读取后立即把dataset.class_to_idx打印出来存成 JSON避免反复靠记忆对齐。图像尺寸方面摘要里提到的 VGG16、ResNet、InceptionV3 输入尺寸不同VGG16 和 ResNet 用 224x224InceptionV3 用 299x299。如果要把这几个模型做对比实验就不能在 Dataset 里写死尺寸而是把 transform 改成可配置的根据选用的模型动态调整。以一个固定脚本跑不同模型改一行配置比改十行数据管道要优雅得多。2.2 标准 Dataset 封装与 transform 组合下面这段代码可以从零构建一个可用的 Dataset兼容 ImageFolder 无法处理的复杂目录结构同时把归一化和数据增强组合在一起。import torch from torch.utils.data import Dataset from torchvision import transforms, datasets import os class FlowerDataset(Dataset): def __init__(self, root_dir, input_size224, modetrain): self.dataset datasets.ImageFolder(root_dir) self.input_size input_size self.mode mode # 核心对每个类别下的文件名按数字排序保证与 image_0612.jpg 这类编号对应 self.samples [] for cls_idx, (cls_name, cls_path) in enumerate(self.dataset.classes): cls_dir os.path.join(self.dataset.root, cls_name) for img_file in sorted(os.listdir(cls_dir)): self.samples.append((os.path.join(cls_dir, img_file), cls_idx)) def __len__(self): return len(self.samples) def __getitem__(self, idx): from PIL import Image img_path, label self.samples[idx] image Image.open(img_path).convert(RGB) if self.mode train: # 数据增强随机旋转 15 度 水平翻转 随机裁剪提升泛化能力 transform transforms.Compose([ transforms.Resize((self.input_size 16, self.input_size 16)), transforms.RandomRotation(15), transforms.CenterCrop(self.input_size), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) else: # 验证集/测试集不做随机增强保持评估稳定性 transform transforms.Compose([ transforms.Resize((self.input_size, self.input_size)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) return transform(image), torch.tensor(label, dtypetorch.long)代码里的两个设计点值得注意第一Resize((input_size 16, input_size 16))配合CenterCrop可以让输入图像在网络入口处获得一定的随机平移效果这是比单纯 Resize 更稳的增强策略第二归一化使用的是 ImageNet 的平均值和标准差——因为我们接下来要加载的预训练模型是基于 ImageNet 统计量训练的输入分布必须对齐到它熟悉的分布区间否则迁移学习效果会大打折扣这是一个新手最常忽略的细节。2.3 训练集 / 验证集切分的边界问题1360 张图的体量80/20 切分意味着验证集只有 272 张每个类别只剩 16 张。这时候如果验证集和训练集来自同一个目录、同一次切分TTA测试时增强都不一定能弥补偶然性。我的建议是使用分层抽样保证每个类别在训练和验证中的比例一致使用sklearn.model_selection.StratifiedShuffleSplit比较稳妥。from sklearn.model_selection import StratifiedShuffleSplit # 取出所有标签用于分层切分 all_labels [label for _, label in dataset.samples] splitter StratifiedShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(splitter.split(all_labels, all_labels)) # 构造子集 from torch.utils.data import Subset train_dataset Subset(dataset, train_idx) val_dataset Subset(dataset, val_idx) print(f训练集: {len(train_dataset)} 张, 验证集: {len(val_dataset)} 张) # 期望输出形如训练集: 1088 张, 验证集: 272 张这里random_state固定为 42 是刻意为之——小数据集上不同的随机切分可能导致最终准确率波动 3 到 5 个百分点。固定随机种子是确保实验可复现的首要前提。注意上面的Subset用法不会复制图像数据只是索引切片内存友好但Subset对象没有classes属性后续要取类别名需要通过dataset.dataset.classes访问这是一个容易踩的暗坑。3. 预训练模型的特征封装与冻结策略识别该项目力推的迁移学习本质3.1 为什么是 VGG16 / ResNet / InceptionV3而不是自建网络摘要中明确提到了 VGG16、ResNet 和 InceptionV3。它们分别代表了三种设计哲学VGG16 是 3x3 卷积堆叠的经典范式结构简单直观但参数量高达 138M推理较慢ResNet 用残差连接解决了深层网络退化问题50 层版本在分类精度和计算量上取得了最佳平衡点InceptionV3 则通过多尺度卷积核并行提取特征参数效率更高。但在这个项目中选哪个模型的判断标准不是 ImageNet 上的 Top-1 准确率而是特征图尺寸和微调成本。17 类花卉与 ImageNet 的 1000 类相关性较高底层卷积核识别的边缘、纹理、颜色渐变通用性很强所以必须完整保留预训练权重。需要替换的只是最后全连接输出层的节点数从 1000 改成 17。实际工程里VGG16 在小数据集上往往是「最稳但最慢」的选择InceptionV3 精确度好但输入尺寸必须 299x299ResNet50 则是折中方案900 多万参数量加上 batch normalization 的预训练统计量收敛速度快。后续代码以 ResNet50 为主因为它的结构最容易被替换成未训练的顶层且不易出现先验崩坏。3.2 三种冻结级别不冻结、冻结卷积基、渐进解冻冻结策略直接影响训练速度与精度常见的三种做法效果差异明显。这里用代码演示并在注释中说明选择逻辑。import torchvision.models as models import torch.nn as nn def build_finetune_model(archresnet50, num_classes17, freeze_level1, pretrainedTrue): # 加载预训练权重pretrainedTrue 代表使用 ImageNet 上训练好的参数 if arch resnet50: model models.resnet50(pretrainedpretrained) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) elif arch vgg16: model models.vgg16(pretrainedpretrained) in_features model.classifier[-1].in_features model.classifier[-1] nn.Linear(in_features, num_classes) elif arch inceptionv3: model models.inception_v3(pretrainedpretrained, aux_logitsFalse) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) else: raise ValueError(fUnknown arch: {arch}) if freeze_level 1: # 只训练最后一层全连接冻结所有卷积特征层 for param in model.parameters(): param.requires_grad False for param in model.fc.parameters(): param.requires_grad True elif freeze_level 2: # 冻结整个特征提取器但解冻最后一个 stage 的残差块resnet50 的 layer4 for param in model.parameters(): param.requires_grad False for param in model.layer4.parameters(): param.requires_grad True for param in model.fc.parameters(): param.requires_grad True # freeze_level0 时全部参数可训练适合数据量大一两个数量级的场景 return model在实际训练中freeze_level1适合快速验证数据流和模型是否通freeze_level2是精度与速度的常见平衡点。为什么解冻最后一个 stage 有帮助因为 layer4 的输出特征最接近语义信息和花卉品种这种细粒度分类高度相关给它适当的梯度可以学习到更贴近任务的形状特征。但要预防的问题是解冻层数越多、学习率就需要越低否则预训练权重很快被冲掉——很多项目最终效果差不是网络不好而是解冻后的初始学习率仍然沿用默认的 1e-3导致底层灾难性遗忘。我一般把特征层学习率设为全连接层的 0.1 倍。3.3 类别不均衡与少样本场景下的 DataLoader 策略17 类每类 80 张类别完全均衡因此基本不需要做类别加权采样。但这个项目的瓶颈在于每类样本太少。常规训练一个 epoch 只有 1088 张图90 个 epoch 也只不过迭代 10 万张次远低于预训练模型原来见过的数据量。一种非常实用的数据增强手段是 MixUp 或 CutMix它们在小数据集分类上提升明显实现也简单。import numpy as np import random def mixup_data(x, y, alpha0.4): MixUp将两张图按比例混合标签也同步混合 lam np.random.beta(alpha, alpha) batch_size x.size(0) index torch.randperm(batch_size) mixed_x lam * x (1 - lam) * x[index, :] y_a, y_b y, y[index] return mixed_x, y_a, y_b, lamMixUp 的直觉是硬标签让模型对决策边界过度自信而混合后的软标签强迫模型学习线性的特征过渡在特征空间中形成更平滑的决策面。218 个 epoch 内它会显著抑制过拟合损失曲线也更稳定。但如果你要做部署推理阶段不需要 MixUp只保留验证时用的标准 transform 即可。这个注意到代码部署时经常被忽略导致推理结果与训练时验证结果有差异。4. 训练循环与评估指标监控损失变化、调优超参数的小细节4.1 优化器选择与学习率调度的工程配置Adam 是默认选择但在这个任务中SGD with momentum 往往在迁移学习场景下取得更高精度。原因在于 Adam 的自适应学习率会在后期出现收敛不稳定尤其在解冻卷积层时SGD 配合余弦退火学习率能更细致地打磨到最优解。以下是两种常用配置的对照。优化器初始学习率权重衰减常用场景注意事项Adam1e-4 ~ 3e-41e-4快速验证、小数据集预实验后期可能出现震荡SGD1e-3全连接层、1e-4特征层5e-4迁移学习微调需要更长训练时间和学习率调度针对这个项目我会用两段式策略先用freeze_level1跑 15 个 epoch训练全连接层此时用 Adam让分类头快速收敛之后切换freeze_level2解冻最后一个 stage用 SGD 配合MultiStepLR把学习率按里程碑降低。import torch.optim as optim from torch.optim.lr_scheduler import MultiStepLR def configure_optimizer(model, freeze_level): if freeze_level 1: optimizer optim.Adam(model.fc.parameters(), lr1e-3, weight_decay1e-4) scheduler MultiStepLR(optimizer, milestones[8, 12], gamma0.1) else: # 不同网络层分配不同学习率Grouped Parameters 是 PyTorch 的标准做法 optimizer optim.SGD([ {params: model.layer4.parameters(), lr: 1e-4}, {params: model.fc.parameters(), lr: 1e-3} ], momentum0.9, weight_decay5e-4) scheduler MultiStepLR(optimizer, milestones[10, 18], gamma0.1) return optimizer, schedulerMultiStepLR的优势是步进明确在第 8 和第 12 个 epoch 学习率降为原来的十分之一让训练进入到精细局部搜索阶段。对应的损失曲线会呈现「陡降-平稳-陡降」的阶梯状这是正常现象不需要担心。在验证集上中期出现的准确率浮动例如从 88% 掉到 86%不一定是过拟合也可能只是学习率偏高产生的震荡此时应当看 5 个 epoch 的滑动平均再做判断。4.2 训练循环中的 early stopping 和模型保存小数据集项目验证集占 20%训练时每个 epoch 的验证损失噪声较大。早期停止的判断基准不应该是最小验证损失而应该是「验证损失在连续 N 个 epoch 中没有创下新低」N 取 10 到 15 比较稳妥。另外保存模型时不能只存 state_dict还需要把class_to_idx一起序列化否则部署阶段会遇到「模型不知道类别顺序」的尴尬问题。def train_epoch(model, loader, criterion, optimizer, device): model.train() running_loss, correct, total 0.0, 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return running_loss / len(loader), 100.0 * correct / total def validate(model, loader, criterion, device): model.eval() running_loss, correct, total 0.0, 0, 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return running_loss / len(loader), 100.0 * correct / total上边的train_epoch和validate是两个独立函数各自通过model.train()和model.eval()切换模型状态。这一切换至关重要BatchNorm 层在训练和推理时行为不同遗忘model.eval()会导致验证准确率比实际偏低 2% 到 5%。后续 checkpoint 记录验证准确率和 epoch并在验证分数新高时覆盖保存最佳权重。import copy best_acc 0.0 best_model_wts copy.deepcopy(model.state_dict()) for epoch in range(1, 26): train_loss, train_acc train_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc validate(model, val_loader, criterion, device) scheduler.step() print(fEpoch {epoch}: train_acc{train_acc:.2f}%, val_acc{val_acc:.2f}%) if val_acc best_acc: best_acc val_acc best_model_wts copy.deepcopy(model.state_dict()) torch.save({ model_state_dict: model.state_dict(), class_to_idx: train_dataset.dataset.class_to_idx, best_acc: best_acc, arch: model.__class__.__name__ }, best_flower_model.pth)copy.deepcopy在 26 个 epoch 内会拷贝 26 次模型权重内存占用略高但换取的是安全性——如果不深拷贝后续任何改动当前模型权重的操作都会破坏存档。到这一步整个训练流程已经跑通。核心调参思路就是小数据 迁移学习下一切以验证集为准但不要迷信单次数值多看 3 个 epoch 的趋势。5. 把模型部署成 FastAPI 接口按需微调还是冻结特征层的最优取舍5.1 部署加载模型时的关键细节训练完的模型最终要提供给前端上传图片并返回预测结果。FastAPI 在这类轻量任务上游刃有余关键是把预处理流程和模型加载写得和训练时完全一致。如果预处理用 Pillow 读图后直接ToTensor但忘记归一化预测结果就会乱成一团。另一个细节是模型必须调用.eval()并整体放到 CPU 或 GPU 上否则推理结果可能不稳定。import io, json, torch from PIL import Image from torchvision import transforms from fastapi import FastAPI, UploadFile app FastAPI() device torch.device(cuda if torch.cuda.is_available() else cpu) checkpoint torch.load(best_flower_model.pth, map_locationdevice) model build_finetune_model(archresnet50, num_classes17, freeze_level0) model.load_state_dict(checkpoint[model_state_dict]) model.to(device).eval() # 与训练时保持一致的预处理流程 preprocess transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) app.post(/predict) async def predict(file: UploadFile): image_bytes await file.read() image Image.open(io.BytesIO(image_bytes)).convert(RGB) input_tensor preprocess(image).unsqueeze(0).to(device) with torch.no_grad(): logits model(input_tensor) probabilities torch.softmax(logits, dim1) # 把索引映射回真实类别名 idx_to_class {v: k for k, v in checkpoint[class_to_idx].items()} top5_idx torch.topk(probabilities, 5).indices.squeeze(0).tolist() results [ {class: idx_to_class[idx], confidence: probabilities.squeeze(0)[idx].item()} for idx in top5_idx ] return {results: results}这段代码的要点集中在idx_to_class的反查和torch.softmax的用法。softmax需要指定dim1表示针对类别维度进行归一化输出总和为 1之后才能用torch.topk取出置信度最高的前 5 个类别。返回前 5 个类别比只返回单一标签更有实用价值——尤其是相近的花色品种模型输出的二三名往往能暴露出类别间的混淆特征。5.2 分析模型未识别出来的情况类别混淆与特征可视化如果准确率不错但部署中发现「雏菊」经常被识别成「蒲公英」这时候有一个非常实用的检查手段是生成混淆矩阵。它能精确到类别对快速定位是数据问题还是类别本身视觉相似的问题。from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt # 加载全部验证集图片统计真实标签与预测标签 y_true, y_pred [], [] model.eval() with torch.no_grad(): for images, labels in tqdm(val_loader): images images.to(device) outputs model(images) _, predicted torch.max(outputs, 1) y_true.extend(labels.cpu().numpy()) y_pred.extend(predicted.cpu().numpy()) # 绘制混淆矩阵观察哪些类互相混淆、哪些类表现极佳 cm confusion_matrix(y_true, y_pred) disp ConfusionMatrixDisplay(confusion_matrixcm, display_labelslist(train_dataset.dataset.classes)) disp.plot(cmapBlues, xticks_rotation45) plt.title(17-Class Flower Classification Confusion Matrix) plt.savefig(confusion_matrix.png, bbox_inchestight, dpi150)从工程视角看混淆矩阵的价值在于告诉你「下一步加大哪一类的训练样本会获得最大收益」。比如两个科属相似的花互相误判盲目增加随机图片不如精准补充这两个类别的侧视、俯视、不同光照条件下的样本。如果某单个类别的准确率在 60% 以下优先检查数据采集是否包含了过多背景干扰而不是先改模型结构。到这里一个完整的花卉分类项目从数据到部署的路径已经全部落地。本文还有配套的精品资源点击获取
返回列表