ARTICLE DETAIL

资讯详情

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

花类识别数据集实战:从解压、清洗到迁移学习训练的完整指南

花类识别数据集实战:从解压、清洗到迁移学习训练的完整指南 简介面向图像分类入门与实战练习的花卉识别数据集整套压缩包适用于机器学习和深度学习初学者快速开展五分类训练与验证。数据共包含4242张花朵图像划分为洋甘菊、郁金香、玫瑰、向日葵、蒲公英五个类别每类约800张图片来自数据流、谷歌图片与Yandex图片分辨率约320×240像素且比例不一更贴近真实采集场景可用作迁移学习、数据增强等进阶实验的练手数据。资源包约449.82MB文件记录共2000项主体为jpg花朵图片另含4个Python脚本、2个pyc与1个txt说明文件脚本可辅助完成数据读取、图片预处理与训练集/测试集划分适合直接嵌入分类实验流程。目前已有606人学习下载适合需要花卉图像数据开展分类实验、课程设计或算法练手的开发者。1. 花类识别数据集.zip别急着解压先想清楚你要拿它干什么花类识别数据集.zip听起来就是一个压缩包解压、扔进训练脚本、跑个准确率完事。但我在实际项目里见过太多人卡在第一步解压出来的东西跟想象完全不一样标注格式是花的目录结构是乱的图片里还混着水印和表情包。这个数据集的真正价值不在那几百兆图片里而在它逼你先把数据工程的基本功过一遍——目录怎么组织、标签怎么编码、样本怎么划分、脏数据怎么清洗这些才是花类识别模型能不能落地的分水岭。适合谁适合要做图像分类但不想从零爬数据的人也适合想拿一个干净基准验证自己数据管道的人。不适合谁不适合指望解压就能训练出生产级模型的人那是另一套工程。2. 解压与体检跑通训练前先花十分钟看清这个 zip 的真实结构2.1 解压并检查文件布局用一棵目录树定位标注格式拿到花类识别数据集.zip第一步不是写训练代码而是先建立一个无菌的检查环境。我一般会在 Linux 服务器或 WSL2 里操作因为后面所有统计命令都是 bash 风格的。解压命令很简单但有几个参数值得较真mkdir -p /data/flower cd /data/flower unzip -q ../花类识别数据集.zip -d raw/ ls -la raw/逻辑说明-q是安静模式避免成千上万条解压日志刷屏-d raw/指定解压目标目录而不是把 zip 里的内容直接摊在当前目录这样后续想删掉重来也干净。解压后先看ls -la确认顶层是单个文件夹还是散落一堆文件。接下来最关键的是绘制目录树。tree命令不是所有系统都有可以用find代替find raw/ -maxdepth 2 -type d | head -50 find raw/ -maxdepth 2 -type f | head -20参数说明-maxdepth 2只往下看两层避免把几百个类别文件夹全部打爆屏幕-type d只看目录-type f只看文件。这两个命令能在 10 秒内告诉你数据集是「train/val 按类别分子目录」的 ImageFolder 结构还是「所有图片平铺 一个 CSV 标注文件」的平面结构。常见的数据集结构有两种。第一种是经典的train/类别名/图片.jpg对应 PyTorch 的torchvision.datasets.ImageFolder几乎是零成本接入第二种是images/xxx.jpg labels.csvCSV 里写着文件名和类别 ID 的映射。如果你的 zip 里两者都有比如all/目录加labels.txt那就要注意了——这通常是别人从某个竞赛平台搬下来重新打包的标注格式可能带着平台特有的编号体系后面要重点核对。2.2 统计类别数与样本量算出每个类的图片数量分布结构看清之后立刻要做的是数量统计。花类识别数据集最怕的不是图片少而是类别极度不均衡。比如“玫瑰”有 5000 张“蒲公英”只有 100 张直接用原始分布训练模型的预测会严重偏向样本多的类别。# 统计每个类别目录下的图片数量 find raw/ -mindepth 2 -maxdepth 2 -type f -name *.jpg | \ sed s|/[^/]*$|| | sort | uniq -c | sort -nr | head -30逻辑说明-mindepth 2 -maxdepth 2限定只统计二级目录下的文件假设目录结构是raw/类别/图片.jpgsed s|/[^/]*$||把文件名部分去掉只保留目录路径uniq -c按目录聚合计数sort -nr按数量降序排列。这样一眼就能看出哪些类别是富样本、哪些是贫样本。如果标注是 CSV用 Python 统计更稳import pandas as pd df pd.read_csv(raw/labels.csv) print(df[label].value_counts().head(30))参数说明value_counts()返回每个类别的频次。这一步你会发现数据集名称里虽然有“花类识别”但具体到类别定义可能很宽泛——有的数据集把“花”分成几十个物种有的只分“玫瑰、向日葵、郁金香”这种粗粒度。这一点直接影响模型容量选择后面会讲。2.3 抽样看图和检查元数据排除损坏文件和标注错位统计完数量必须抽图看内容。这是整个流程里最容易被跳过但后果最严重的步骤。用 Python 写一个快速抽样脚本把每个类别的首尾各抽一张拼成 contact sheetimport os from PIL import Image import matplotlib.pyplot as plt root raw/ class_dirs sorted([d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]) fig, axes plt.subplots(len(class_dirs), 3, figsize(12, len(class_dirs) * 1.5)) for i, cls in enumerate(class_dirs): imgs sorted(os.listdir(os.path.join(root, cls)))[:3] for j, img_name in enumerate(imgs): img Image.open(os.path.join(root, cls, img_name)).convert(RGB) axes[i, j].imshow(img) axes[i, j].axis(off) axes[i, 0].set_ylabel(cls, fontsize8) plt.tight_layout() plt.savefig(contact_sheet.png, dpi100)逻辑说明sorted保证每个类别的抽样顺序一致convert(RGB)强制转成 RGB避免灰度图或带 alpha 通道的 PNG 在后面训练时引起通道数不匹配axes[i, 0].set_ylabel(cls)在每行左侧标注类别名。跑完这个脚本打开contact_sheet.png扫一眼2 分钟能发现大部分问题。这一眼看过去要确认三件事。第一图片内容是否和类别名匹配——有没有“玫瑰”目录里混着月季甚至玫瑰花束包装纸的第二图片是否来自同一分布——有没有大量网络截图、带水印的商品图、带 UI 界面的手机截图第三有没有明显损坏的图片——PIL 打开时报错或解码后全黑的。最后一项可以用一个快速脚本批量验证python -c from PIL import Image import os, sys rootraw/ bad [] for dirpath, _, files in os.walk(root): for f in files: try: with Image.open(os.path.join(dirpath, f)) as im: im.verify() except Exception as e: bad.append((os.path.join(dirpath, f), str(e))) print(损坏文件数:, len(bad)) for b in bad[:10]: print(b) 参数说明Image.verify()只检查文件完整性不加载像素数据速度很快注意 verify 之后不能直接用它做后续处理要重新Image.open()。这段脚本对几千张图也要不了几秒任何返回非空的输出都要记录并准备清洗。3. 从原始图片到可训练样本目录重排、标签编码与可复现的数据划分3.1 重排目录为 ImageFolder 标准结构解决 train/val 混合的问题绝大多数花类识别数据集在 zip 里不会贴心地替你分好 train/val/test。就算分了也可能只是按文件名前缀或者干脆全部混在一起。为了让后续训练代码能直接复用torchvision.datasets.ImageFolder我一般会先写一个重排脚本把数据集统一成data/train/类别/图片.jpg和data/val/类别/图片.jpg的结构。#!/bin/bash SRCraw/ DSTdata/ mkdir -p ${DST}/train ${DST}/val # 按 8:2 比例划分每个类别 for cls_dir in ${SRC}*/; do cls$(basename $cls_dir) mkdir -p ${DST}/train/$cls ${DST}/val/$cls imgs($cls_dir*.jpg) total${#imgs[]} val_count$((total * 2 / 10)) # 先排序保证可复现性 IFS$\n sorted($(sort ${imgs[*]})); unset IFS for ((i0; i${#sorted[]}; i)); do if (( i val_count )); then cp ${sorted[$i]} ${DST}/val/$cls/ else cp ${sorted[$i]} ${DST}/train/$cls/ fi done done echo 划分完成: $(find ${DST}/train -name *.jpg | wc -l) 张训练图, $(find ${DST}/val -name *.jpg | wc -l) 张验证图参数说明val_count$((total * 2 / 10))表示每个类别固定取前 20% 作为验证集而不是全局随机抽样这对类别不均衡的数据集至关重要——保证每个类在训练和验证里都出现。排序那步用IFS$\n sorted($(sort ...))是为了按文件名排序避免文件系统返回顺序不一致导致每次划分结果不同。为什么要复制而不是移动原文件因为原始 zip 可能还要留着做其他实验复制一份相当于制造“后悔药”。但这个脚本有个隐患它假设每张图都是.jpg后缀。如果数据集混着.png或.jpeg这里的*.jpg通配符会把它们漏掉。稳妥做法是用find -type f遍历所有图片扩展名或者统一用 Pillow 转换。到这里数据管道的第一个标准化产物就有了——一个能被 ImageFolder 直接加载的目录树。3.2 用 Python 脚本完成标签编码生成 id 到类别名的映射文件目录结构定下来后类别名本身就是标签但中文字段名或特殊字符在深度学习框架里容易惹麻烦。训练代码里通常用整数索引做 label类别名字符串只用于最终的 confusion matrix 展示。所以需要生成一个class_names.txtimport os train_root data/train classes sorted([d for d in os.listdir(train_root) if os.path.isdir(os.path.join(train_root, d))]) with open(class_names.txt, w, encodingutf-8) as f: for idx, cls in enumerate(classes): f.write(f{idx}\t{cls}\n) print(f共 {len(classes)} 个类别映射已保存到 class_names.txt)逻辑说明sorted()确保类别索引稳定不随文件系统遍历顺序变化idx从 0 开始与 PyTorch 的 CrossEntropyLoss 默认类别索引一致。这个文件是训练和推理共用的“单一事实来源”——模型输出整数推理脚本查这个文件转成中文类别名。这里有个容易忽略的点如果原始 zip 里的类别名带了空格或括号比如Rose (red)一定要在映射文件里原样保留同时确认目录名也一致。我见过有人手改映射文件但忘了改目录名训练到一半才发现 FileNotFoundError这个错位非常隐蔽。3.3 设定固定随机种子把数据加载器配置成可复现实验目录和标签都就位接下来要保证“同一份代码跑两次结果一致”。深度学习训练有大量随机性——数据 shuffle、权重初始化、数据增强的随机裁剪如果不固定随机种子前后两次实验就失去了可比性调参就变成玄学。import random, numpy as np, torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) # 让 cuDNN 使用确定性算法 torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False set_seed(1213)参数说明torch.backends.cudnn.deterministic True会让 cuDNN 选择确定性卷积算法虽然可能慢一点但结果可复现benchmark False关闭自动调优否则同一个模型在不同批次大小下可能选不同算法。1213这个种子值没有特殊含义但一旦固定下来就不要轻易改否则前面所有调参记录都作废。数据加载器方面DataLoader的shuffleTrue依赖内部随机数生成器PyTorch 提供generatortorch.Generator().manual_seed(seed)做精细控制。多进程 worker 的随机性更难完全复现但这已经足够了——训练和验证的划分是确定的这就保证了准确率的波动主要来自模型本身而不是数据管道。4. 训练一个花类识别基线模型选对网络、设对参数、跑通最小闭环4.1 为什么迁移学习是花类识别最快见效的路径花类识别属于细粒度图像分类的入门级变体——类别之间有区分度但不像“鸟种识别”那样需要数羽毛纹理。处理这种任务从零训练一个 CNN 是最坏的选择因为数据集规模通常只有几千到几万张而花类图片的纹理、颜色、形态特征恰恰是 ImageNet 预训练模型已经学过的底层模式。常见做法是采用迁移学习加载 ImageNet 预训练的 ResNet 或 EfficientNet冻结前几层只微调最后几层和分类头。我一般会优先试resnet18或resnet34原因很实际——训练周期短显存占用小迭代快而且花类识别任务通常不需要 ResNet50 那种容量就能达到 95% 以上准确率。如果数据集类内差异很大比如同一种花有盛开、含苞、枯败等多种状态再考虑升级到 EfficientNet-B3 或 ConvNeXt-Tiny。模型的容量选择要匹配数据规模。resnet18有 1100 万参数左右适合 5000-20000 张图的规模如果每个类别只有一两百张用更小的mobilenet_v3_small反而更稳不容易过拟合。训练花类数据集过拟合是最大的敌人数据增强比换大模型更有效。4.2 用 PyTorch 写一个可运行的花类训练脚本核心参数逐一说明下面这个脚本是花类识别训练的最小闭环自带验证集评估和 checkpoints 保存。我尽量保持通用你可以直接替换数据路径跑起来import torch import torch.nn as nn from torch.utils.data import DataLoader, random_split from torchvision import datasets, transforms, models # ---------- 超参数 ---------- BATCH_SIZE 32 EPOCHS 30 LR 1e-3 # 微调阶段用 1e-3全网络微调用 1e-4 到 3e-4 MOMENTUM 0.9 WEIGHT_DECAY 1e-4 SEED 1213 # ---------- 数据增强 ---------- train_transforms transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transforms transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # ---------- 数据集 ---------- train_dataset datasets.ImageFolder(data/train, transformtrain_transforms) val_dataset datasets.ImageFolder(data/val, transformval_transforms) train_loader DataLoader(train_dataset, batch_sizeBATCH_SIZE, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_sizeBATCH_SIZE, shuffleFalse, num_workers4, pin_memoryTrue) # ---------- 模型 ---------- model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) num_ftrs model.fc.in_features model.fc nn.Linear(num_ftrs, len(train_dataset.classes)) model model.to(cuda if torch.cuda.is_available() else cpu) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lrLR, momentumMOMENTUM, weight_decayWEIGHT_DECAY) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.5)参数说明RandomResizedCrop的scale(0.6, 1.0)限制裁剪比例不低于 60%避免花的主体被切得只剩一角——花类图片的判别信息集中在花心而不是整张图的背景RandomRotation(15)的角度不要超过 15 度旋转太多会让花瓣结构失真ColorJitter的饱和度扰动对花类特别有效因为户外拍摄时同一朵花在不同光照下的饱和度差异极大。SGD 配动量是迁移学习的稳妥组合Adam 适合从零训练但微调阶段 SGD 的泛化更好。StepLR每 10 轮降一半学习率适合 30 轮的训练节奏。4.3 训练循环里加验证集评估跨轮保存最优模型训练脚本不能只跑训练必须每轮结束都在验证集上评估并且把最优模型单独存出来。否则你跑完 30 轮最后保存的那一个 checkpoint 可能在某轮已经过拟合了best_acc 0.0 for epoch in range(EPOCHS): model.train() running_loss 0.0 for inputs, labels in train_loader: inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() # 验证 model.eval() correct 0 total 0 with torch.no_grad(): for inputs, labels in val_loader: inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() val_acc 100.0 * correct / total print(fEpoch {epoch1}/{EPOCHS} | Loss: {running_loss/len(train_loader):.4f} | Val Acc: {val_acc:.2f}%) # 保存最优模型 if val_acc best_acc: best_acc val_acc torch.save({ epoch: epoch 1, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), val_acc: val_acc, class_names: train_dataset.classes }, best_flower_model.pth) print(f 已保存新最优模型验证准确率 {val_acc:.2f}%)逻辑说明model.train()和model.eval()切换训练/验证模式——eval 模式关闭 dropout 和 BatchNorm 的统计更新如果不切验证集结果会虚高torch.no_grad()关闭梯度计算显存占用大幅下降而且验证不需要反传torch.max(outputs.data, 1)取预测类别索引和前面编码的整数对齐。保存 checkpoints 时把class_names也存进去比单独维护一个文本文件更保险推理时只要加载 checkpoint 就能知道类别映射。这个训练循环在单张 RTX 3090 或 4090 上resnet18 32 batch size 大约每个 epoch 5-10 分钟取决于图片总数30 轮大概 3-5 小时。如果你的机器只有 CPU 或者只有 4GB 显存的卡把 batch size 降到 16模型换成mobilenet_v3_small时间可以压缩到 1-2 小时准确率一般也有 90% 以上。4.4 类别不平衡时的策略加权损失函数的引入点如果前面统计时发现某个类别样本极少比如少于 50 张那么直接算 CrossEntropyLoss 会让模型把稀有类全部忽略。此时要给 loss 加上类别权重权重一般取样本数的倒数或倒数的平方根import numpy as np from collections import Counter sample_counts Counter([label for _, label in train_dataset.samples]) total_samples len(train_dataset.samples) weights torch.tensor( [total_samples / sample_counts[i] for i in range(len(train_dataset.classes))], dtypetorch.float32 ).to(device) criterion nn.CrossEntropyLoss(weightweights)参数说明样本总数 / 该类别样本数是最常用的加权方式——样本多的类别权重小样本少的类别权重大让 loss 在稀有类上贡献更大。注意train_dataset.samples是 ImageFolder 内存放的(路径, 标签)列表可以直接取标签做统计。加了权重之后训练曲线会有明显波动验证准确率反而可能下降一点但混淆矩阵里稀有类的召回率会明显提升。5. 花类识别避坑手册解压到训练全流程的 5 个真实踩坑点5.1 图片文件名包含中文或空格shutil 复制时突然报错现象按类别目录重排时shutil.copy抛FileNotFoundError但文件明明存在。排查后发现问题出在文件名里夹杂了全角空格和中文括号shell 脚本里路径没有正确引用。原因zip 文件里的文件名是 UTF-8 编码中文文件名本身没问题但如果在 bash 里用未加引号的变量传递路径空格会被当成分隔符拆词。Python 端的os.listdir能识别但是某些第三方库统计时对特殊字符处理不当。解决所有路径变量一律加双引号cp ${SRC}/${cls}/${img} ${DST}/train/${cls}/Python 脚本里统一用pathlib.Path处理路径不要手动拼接字符串。清洗阶段最有效的手段是把所有文件名重命名为纯 ASCII比如类名序号.jpg丑但绝对安全。5.2 验证集准确率 95%但真实场景一测就翻车现象模型在清洗过的验证集上表现很好95%换到手机随手拍的花上准确率掉到 60%。反复检查代码也没有 bug。原因数据集里的图片分布单一——背景干净、主体居中、光线均匀。真实场景里花是嵌在复杂背景里的有遮挡、有阴影、有虚化。验证集的分布和部署场景不一致准确率只是“自我感觉良好”。解决训练时增强背景扰动RandomResizedCrop的scale下限调到 0.3并加入RandomErasing模拟遮挡另建一个 50-100 张的“野采”测试集从网上下载不同拍摄风格的花图单独做一次评估这才是真实水平的反映。如果野采测试集准确率低于 80%说明数据集本身和你的落地场景有 gap要回头补充数据。5.3 训练时 loss 直接变成 nan前两步就崩现象loss 在第一个 batch 后就变成nan验证集准确率一直是 0。检查数据没问题代码也没有明显错误。原因最常见是学习率过大或者输入图片里有全黑的损坏图导致 BatchNorm 计算方差为 0。resnet18预训练模型在 ImageNet 归一化参数下如果输入值域不对比如忘记ToTensor()导致输入还是 0-255 的整数梯度会瞬间爆炸。解决先确认预处理管道里ToTensor()和Normalize按顺序执行再检查有没有全黑或全白图片用前面 5.2 的 verify 脚本再跑一遍最后把学习率降到1e-4试试——如果 loss 恢复正常说明原来的 LR 对这个 batch size 太大了。血泪经验nan问题 80% 出在数据端不是模型端。5.4 同样代码在两个环境跑验证集准确率差 3 个百分点现象同一个脚本在 A 服务器上训练完验证集准确率 93%在 B 服务器上复现只有 90%。代码一样、数据一样、超参一样检查了随机种子也一样。原因PyTorch 版本差异导致ResNet18_Weights.IMAGENET1K_V1的下载权重可能有微小差异cuDNN 版本不同导致卷积算法的浮点舍入不同数据加载器的num_workers不同导致 shuffle 顺序不同。解决训练模型前把torch.__version__、torchvision.__version__、CUDA 版本记录下来写进训练日志开头固定随机种子的同时把DataLoader的generator也固定如果要求极高可以用torch.use_deterministic_algorithms(True)强制所有算子确定性但这会禁用某些不支持确定性的算子慎用。真正要复现的是实验结论不是逐比特一致。5.5 数据集里的“花”和你要识别的“花”不是一回事现象用花类识别数据集.zip 训练后模型对月季、蔷薇、玫瑰的区分度很差但对菊花、向日葵、郁金香的识别率很高。看起来是数据集本身的类别定义和你的需求错位。原因很多公开花类数据集是基于牛津 102 类花卉Oxford-102或类似竞赛数据集构建的类别划分遵循植物学分类而你的业务诉求可能是按商品名分类——比如“玫瑰”和“月季”在植物学上是不同物种但电商场景里它们可能归在同一类。解决先看class_names.txt里的类别列表对照你的实际需求逐项核对。如果默认类别定义与你需求不符合并相近类别、重命名标签都是可行的——直接在data/train目录下用软链接或复制重命名即可不需要重写训练代码。最怕的是不看类别就拿去用训出来才发现标签体系对不上。6. 让模型真正可用混淆矩阵定位易混类ONNX 导出提速部署训练完拿到一个 90% 的模型并不意味着项目结束了。我一般会再做两件收尾的事一是用混淆矩阵找出模型分不清的类别二是把 PyTorch 模型导出成 ONNX让推理速度提升一个档次。混淆矩阵的实现非常直接核心就是收集预测结果import seaborn as sns import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, classification_report all_preds [] all_labels [] model.eval() with torch.no_grad(): for inputs, labels in val_loader: inputs inputs.to(device) outputs model(inputs) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) class_names train_dataset.classes plt.figure(figsize(12, 10)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclass_names, yticklabelsclass_names) plt.xlabel(Predicted) plt.ylabel(True) plt.xticks(rotation45, haright) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150) print(classification_report(all_labels, all_preds, target_namesclass_names))逻辑说明classification_report会输出每个类别的 precision、recall、F1-score比只看准确率信息量大得多。比如“菊花”和“蒲公英”在照片里都是白色放射状花瓣如果这两类的混淆矩阵数值明显偏高说明模型抓取的特征还不够分可以考虑针对性增加这两类的样本或者裁掉图片里过多的背景区域。关于 ONNX 导出PyTorch 自带torch.onnx.export但有几个参数直接影响部署时的表现dummy_input torch.randn(1, 3, 224, 224).to(cpu) model_simple model.to(cpu).eval() torch.onnx.export( model_simple, dummy_input, flower_model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}}, opset_version13 )参数说明dynamic_axes把 batch 维设为动态这样导出的 ONNX 模型可以接受任意 batch size 的输入部署到生产环境时不需要固定为 1。opset_version13比较保守兼容性最好。导出后建议用onnxruntime验证输出是否和 PyTorch 一致python -c import onnxruntime as ort import numpy as np sess ort.InferenceSession(flower_model.onnx) input_name sess.get_inputs()[0].name ort_out sess.run(None, {input_name: np.random.randn(1, 3, 224, 224).astype(np.float32)}) print(ONNX 输出维度:, ort_out[0].shape) 到这一步花类识别数据集的价值才真正闭环——你手里有一个能跑通训练、能定位错误、能导出部署的完整链路。最后说一个我的习惯每次训练完我会把epochs、batch_size、lr、最终验证准确率、数据增强配置记成一个experiment.log存进项目目录。这个文件在三个月后回看时比代码本身更值钱——没有哪个参数组合是全能的记录了什么能跑比知道什么理论最优更实在。希望帮你避开我踩过的坑。本文还有配套的精品资源点击获取
返回列表