PyTorch图像分类完整工程:数据清洗、ResNet50迁移学习与ONNX部署

发布时间:2026/9/16 19:23:11
PyTorch图像分类完整工程:数据清洗、ResNet50迁移学习与ONNX部署
简介基于PyTorch打造的深度学习物体分类系统完整工程包以图形化界面贯穿数据集构建、模型训练、测试评估与可视化展示全流程适合从事计算机视觉学习或毕业设计开发的读者快速上手。压缩包共含73个文件、约94.72MB涵盖24个Python脚本如数据清洗与划分、训练、预测、UI交互等、45张图像样本、1个预训练权重文件、依赖清单及说明文档目录模块化组织便于按需调用。目前已有778人学习/下载实践参考价值较高。系统中集成了ResNet等主流模型训练流程可通过窗口程序直观观察损失曲线、准确率及预测结果并方便地调整超参数配套的数据处理脚本、模型微调与测试代码覆盖了从数据准备到部署导出的主要环节可作为课程设计或入门级物体分类任务的工程范式。1. 从 PyTorch 到图形化界面一套能跑完整物体分类流程的工程代码在实际做图像分类落地时最耗精力的往往不是搭网络而是把散落的图片整理成数据集再把训练好的模型交给一个不太懂命令行的人去用。这个项目pytorch110_classification-master就把这条链路完整串了起来数据清洗、自动划分、基于 timm 的 ResNet50D 预训练模型微调、标准训练循环以及 window.py 提供的图形化交互界面。它面向两类人一类是刚接触 PyTorch 物体分类想照着一套能跑的代码理解训练全流程的学生另一类是已经在写代码、但想把分类工具快速变成内部可用产物的工程师。对后者而言值得细看的是预训练权重如何剥掉分类头、无验证集时怎么训练以及 ONNX 导出与 FLOPs 统计这些交付前的细节。下面按数据、模型、界面、部署的顺序拆开讲。2. 数据清洗与切分让图片文件夹变成模型能吃的 Dataset2.1 数据清理和自动整理data_clean.py 与 mv_imgs.py很多原始分类图片是从网上爬的或者现场拍的里面混着损坏文件、重复图、命名含空格的文件。data_clean.py 负责这一关。常见做法是先扫描目录用 PIL 打开校验完整性删掉没法解码的文件同时把扩展名统一成 .jpg。这里我建议在调用任何训练代码之前先跑一遍 clean否则后续 DataLoader 在遇到坏图时会直接中断整个 epoch。mv_imgs.py 则解决“不同类别放在一个目录”的问题。假设你从好几个文件夹收集了同一类物体需要把它们汇到统一的 train/cat 目录下可以这样执行python mv_imgs.py --src ./tmp_download --dst ./imgs --class_name cat --ext .jpg这条命令把 src 下所有 .jpg 文件移动到 ./imgs/cat/ 目录。如果同名文件存在它会自动追加短哈希后缀而不是覆盖。参数 --ext 用来限定扩展名避免把不想要的 .gif 动画也搬过去。我一般会把 --dst 底下的每个类别再建子目录比如飞机、汽车、船各一个文件夹这样后面 data_split.py 直接根据子目录名生成标签。这里有个实际建议不要直接用设备号当目录名比如 img001、img002否则类别列表和真实语义对不上。目录名最好是人类可读的英文或拼音并且保持一致例如 dog、cat、bus这样界面显示时就不用再做一次映射。2.2 data_split.py按比例切分训练集和验证集分类任务最忌讳手动分数据漏一张忘一张会导致模型评估虚高。data_split.py 的做法是扫描指定根目录把每个子文件夹作为一个类别然后按比例随机分配索引去生成 train.txt / val.txt而不是物理移动文件。这种方式在后续调试时回滚非常方便。python data_split.py --data_dir ./imgs --output ./split --train_ratio 0.8 --val_ratio 0.2 --seed 42参数含义--data_dir 是图片根目录每个子目录名就是类别标签--output 存放划分清单--train_ratio 和 --val_ratio 分别控制训练和验证比例留空的部分给测试集--seed 固定随机种子保证多次划分结果一致。我一般会把 seed 设为 42这样换机器、重跑时结果可复现。注意 train_ratio 加 val_ratio 最好小于 1留一部分独立测试集否则你最后评估模型时没有真正没见过的情况。2.3 DataLoader 中的预处理与数据增强有了文本清单还需要把它转成 PyTorch Dataset。项目里 torchutils.py 封装了这套逻辑核心 transform 如下train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), 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_transform 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]) ])RandomResizedCrop 的 scale 参数控制裁剪区域占原图比例数值越小模型看到的局部细节越多但太小会丢失全局信息我常用 0.71.0。ColorJitter 对光照变化的真实场景很关键比如工厂里不同时段拍摄的零件。val_transform 不用随机增强Resize 后中心裁剪是通用做法保证评测稳定性。Normalize 的均值和标准差直接采用 ImageNet 统计值因为后面要加载基于 ImageNet 预训练的 ResNet50D保持输入分布一致才能继承权重。除了 transformDataLoader 里还有几个容易忽略的参数我把常用配置列在下面参数推荐值说明batch_size16 / 32由显存决定显存不够时优先减小num_workers4 (Linux) / 0 (Windows)Windows 下多进程容易报错pin_memoryTrue用 GPU 训练时开启减小拷贝时间drop_lastTrue最后一个 batch 不足时丢弃稳定 BN 统计这里的 drop_last 看起来不起眼但如果不设 True最后一个 batch 只有几张图时BatchNorm 的均值和方差会剧烈抖动表现为 loss 曲线上出现周期性的尖峰。尤其是当你的训练集大小不能被 batch_size 整除时这个坑很隐蔽。3. 用 timm 预训练 ResNet50D 构建分类模型从加载权重到训练循环3.1 为什么选 ResNet50D 和 timm 库ResNet50 已经是老面孔但 D 变体在 stem 部分做了优化输入经过 3×3 卷积而不是 7×7同时在降采样的残差分支里补上 2×2 平均池化和 1×1 步长 2 卷积减小了信息丢失。timm 库把这些细节都封装好只需要一行代码就能创建带预训练权重的模型。本项目的 timm_models.py 就是基于 timm 的二次封装它支持直接指定模型名比如 resnet50d、efficientnet_b0改动结构时不污染业务代码。选择 timm 而不是 torchvision 还有一个现实原因torchvision 官方没有原生的 resnet50d 定义得自己改 stem而且它导出的权重文件是单一 state_dict接口不如 timm 统一。timm 里带 _ra2 后缀的权重使用了 RandAugment 技巧预训练迁移到小数据集上收敛更快。对于物体分类这种任务模型的 backbone 负责特征提取最后接一个全局平均池化和全连接层即可timm 的 create_model 已经把分类头也建好了。3.2 加载预训练权重 resnet50d_ra2-464e36ba.pth项目把预训练权重放在 pretrained/ 目录下文件名 resnet50d_ra2-464e36ba.pth 对应 timm 官方仓库的权重。加载时有个细节这个权重是完整模型的 state_dict包含最后 ImageNet 的 1000 类分类层而我们自定义数据集的类别数往往不是 1000。常见做法是先把模型创建出来再用 torch.load 加载权重并去掉最后分类层不匹配的键。import torch, timm model timm.create_model(resnet50d, pretrainedFalse, num_classes10) state_dict torch.load(pretrained/resnet50d_ra2-464e36ba.pth, map_locationcpu) # 过滤掉全连接层参数因为 num_classes 不同 state_dict {k: v for k, v in state_dict.items() if not k.startswith(fc.)} missing, unexpected model.load_state_dict(state_dict, strictFalse) print(missing keys:, [k for k in missing if not k.startswith(fc.)])这里 strictFalse 允许最后分类层缺失同时保留所有卷积层和 BatchNorm 层的预训练参数。missing 列表中通常只剩下 fc.weight 和 fc.bias说明其余部分全部加载成功。注意 map_location 设为 cpu 是为了避免在加载阶段就把模型放上 GPU后面训练循环里再用 .to(device) 移动防止多卡或低显存环境下的显存波动。如果你不慎把 state_dict 直接给 model.load_state_dict 而 strictTrue会看到 KeyError因为 fc 层的输出维度不匹配。3.3 训练主流程train.py 与 train_without_val.py 的差异项目里有两个训练入口。train.py 走标准训练验证流程每个 epoch 结束后在验证集上计算 Top-1 Accuracy用于挑选最佳模型train_without_val.py 则面向验证集不好收集的场景比如数据太少没法单独划分验证集。它每 N 个 step 打印一次 loss并把最后一个 checkpoint 直接保存。拿 train.py 举例核心训练循环做了三件事for epoch in range(start_epoch, epochs): model.train() 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() * inputs.size(0) # 每 epoch 后更新学习率 scheduler.step() # 验证并保存最优权重 save_best(model, val_loader, epoch, save_dir)train_without_val.py 的区别在循环里没有 val_loader 部分但它会维护一个滑动平均损失当 loss 连续几个 epoch 不下降时自动把学习率降为原来的 0.1。这个策略对快速迭代实验很有用避免反复人工盯日志。你在使用这个脚本时要注意它保存的最后 checkpoint 不一定是最优状态所以建议每隔固定步数额外存储一个带 step 编号的权重方便回溯。3.4 损失函数、优化器与学习率调整物体分类默认使用交叉熵损失PyTorch 里直接用torch.nn.CrossEntropyLoss()。如果你的类别很不均衡可以在构造时传入weight张量把少数类的权重设大。优化器我倾向于 AdamW权重衰减weight_decay设为 1e-4用 SGD momentum0.9 会需要更多 epoch 才能达到同样性能但在小数据集上最终精度往往稍高一点。项目默认训练 30 个 epochbatch size 32初始学习率 1e-3。学习率调整采用余弦退火CosineAnnealingLR其优势是前期下降慢、后期快速收敛比固定学习率更容易稳定在低损失区域。以下配置可以直接复用from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR optimizer AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max30, eta_min1e-6)T_max 应该与总 epoch 数一致这样最后一个 epoch 学习率会降到 eta_min。如果训练过程中 loss 变成 NaN优先检查 batch 里是否存在全黑图片或部分标签超出类别范围这比调学习率更有用。另外当验证集精度在某个 epoch 后突然下降往往是学习率踩到了过大值这时候把初始 lr 降到 3e-4 重新跑一轮会更稳。4. window.py 与 predict.py把训练好的模型封装成图形化分类工具4.1 图形界面的工程结构window.py 用 PySide2 或 Tkinter 这类常见桌面框架实现了主界面我拆下来发现它把“训练、验证、预测、模型导出”四个动作分别挂载到四个按钮上。负责训练的按钮在后台线程里调用 train.py 的入口函数避免主线程阻塞导致界面卡死预测按钮则绑定到 on_select_image 事件点击后弹出文件对话框选择图片随后调用 predict.py 的封装函数。这种解耦思路值得借鉴。界面层不直接 import 模型和 Datasets而是通过接口函数传递参数。比如 window.py 里的类别列表是从训练时保存的花名册classes.txt读出来的而不是在 UI 里硬编码这样换数据集时不需要改界面代码。如果你只想做推理演示可以直接把 window.py 里训练相关的按钮禁用保留一个“打开图片”和一个“开始识别”就够了。4.2 predict.py 的单图预测核心代码predict.py 封装了从图片路径到分类结果的完整推理流程。它先加载训练时保存的 checkpoint再重新构建与训练时一致的模型结构最后做预处理和推理。这里有个关键坑transform 必须和验证时保持一致如果你在训练时用了 RandomResizedCrop预测时一定不能沿用否则每张图裁剪位置随机结果不稳定。def predict_image(image_path, model, transform, class_names, devicecpu): from PIL import Image import torch.nn.functional as F img Image.open(image_path).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits model(img_tensor) probs F.softmax(logits, dim1)[0] top2 torch.topk(probs, 2) result [(class_names[i], probs[i].item()) for i in top2.indices] return result参数说明transform 传入的就是上一章 val_transform保证输入尺寸和归一化一致class_names 是一个 list其索引必须与模型输出通道一一对应。softmax 在 CPU 和 GPU 上结果略有误差但 top-2 基本不受影响。返回的 result 是一个列表里面每个元素包含类别名和置信度方便 UI 直接显示。注意这里的 Image.open 最好做一次 RGB 转换否则遇到带 alpha 通道的 PNG 会报尺寸错误。4.3 图形界面中的结果显示与阈值处理window.py 拿到 predict_image 的返回值后会把它渲染在右侧 label 上并用不同颜色标识置信度。我会在这里加一道阈值判断置信度低于 0.6 的结果显示为“疑似未知类别”避免模型对着背景图强行输出一个高置信度标签。这个阈值要放在模型输出之后不要在界面上做二次判断这样 predict.py 就能独立复用命令行也能调用。如果你要处理一批图片可以用循环批量调用 predict_image但每次调用都会重新走一次 forward。更高效的做法是把图片堆成一个 batch一次 forward 得到全部结果。界面演示通常不追求这个速度但当图片数量到达几百张时批处理能明显减少等待时间。4.4 常见运行坑设备、路径和类名顺序实际跑 window.py 时会遇到几个高频报错。第一个是RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor)原因是用 GPU 训练的权重在 CPU 上加载后没有把输入也放到 CPU反之亦然。解决方式是统一通过 device 参数控制model model.to(device)和img_tensor img_tensor.to(device)两者缺一不可。第二个坑是把模型文件和图片放在不同目录导致相对路径找不到权重我建议在界面初始化时用绝对路径拼接 checkpoint 地址不要依赖当前工作目录。第三点是类名顺序错乱。如果训练时用sorted(os.listdir(data_dir))生成类别列表那么验证和预测时也要用相同的排序规则。如果把训练时保存的 classes.txt 丢掉仅凭 checkpoint 文件无法知道类别对应关系。所以我的习惯是在训练结束后立即把 class_names 序列化成 json和权重一起存到同一目录。window.py 启动时会先检查这个 json 是否存在不存在就弹出提示而不是等到点预测才报错这样能尽早暴露配置问题。5. 导出 ONNX、统计 FLOPs以及交付前的验证技巧5.1 用 export_onnx.py 导出推理模型当模型训练完毕很多人会直接把 .pth 交给下游。但 PyTorch 权重依赖环境里的 Python 和 torch 版本而 ONNX 可以跨平台交换。export_onnx.py 的核心逻辑如下python export_onnx.py --checkpoint output/best_model.pth --output model.onnx --num_classes 10 --opset 11内部执行的是标准 torch.onnx.exportdummy_input torch.randn(1, 3, 224, 224) torch.onnx.export(model, dummy_input, model.onnx, opset_version11, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}})这里 dynamic_axes 把 batch 维度设为动态这样在同一份 ONNX 模型上既能推理单张图片也能一次推理多张图片。opset_version 选 11 是因为它兼容性最好从 PyTorch 1.10 到 2.x 都能正常加载。导出前需要保证 model 已加载了最佳权重并切换到 eval 模式否则 BatchNorm 的统计量是训练中的导出结果会有偏差。导出的 ONNX 可以用onnxruntime做一次推理对照输入一张相同的图片确认 softmax 输出与 PyTorch 原始结果误差小于 1e-4。5.2 get_flops.py量化模型计算量部署到边缘设备前要先算算 FLOPs 和参数量。get_flops.py 使用 thop 库它比 torchsummary 更轻量并且能统计卷积、全连接、BatchNorm 等各层贡献。命令和代码非常直接python get_flops.py --checkpoint output/best_model.pth --num_classes 10核心输出逻辑from thop import profile input_tensor torch.randn(1, 3, 224, 224) model.eval() flops, params profile(model, inputs(input_tensor,)) print(fFLOPs: {flops/1e9:.3f}G, Params: {params/1e6:.3f}M)这里不把 FLOPs 除以 2因为 thop 默认已经按乘加操作数算过一次。实际情况中 resnet50d 在 224×224 输入下约 5.4G FLOPs 和 26M 参数如果你的结果差太多优先检查是否把 input 尺寸写错或者模型没有加载预训练权重。FLOPs 这个数字不要只看绝对值还要结合部署设备的算力评估比如一块 Jetson Nano 的 INT8 算力约 0.5 TOPS跑 5.4G 的模型大约需要 11ms 的纯计算时间实际加上数据读取会更多。5.3 data_test.py 的批量评测与 Top-5 统计模型导出前还应该在测试集上重新评估一遍data_test.py 做的是批量读取 test_imgs 目录对每张图调用预测函数并统计准确率。它额外输出 top-1 和 top-5 两组指标因为很多分类任务的真实评价比 top-1 更宽容。代码里通过比较preds.topk(5)的索引与标签集合来判断 top-5 是否命中这比手动循环逐个判断要快得多。这里有一个容易忽略的细节data_test.py 默认使用和训练时相同的 class_names 顺序所以测试前它会重新加载训练输出目录下的 classes.json。如果你的测试图片目录里混入了不属于任何类别的负样本建议单独写一个--unknown_threshold参数当最大 softmax 值低于该阈值时视为未知类不计入准确率统计。5.4 最后留一个实用技巧用 classes.json 验证类名映射我吃过一次亏训练时类别列表来自os.listdir后来手动删过一个空文件夹导致分类索引整体前移一位预测结果全部错位。从那以后我在训练脚本里强制要求把类别列表写入 classes.json并且在预测脚本启动时先做一次自检把 classes.json 里的类别数与模型输出维度做比对数量不一致直接放弃加载。对于本项目你只需要在 window.py 初始化时打印一下 classes.json 的前三行确认与训练时的文件夹名完全一致即可。这个检查成本极低却能把“模型精度高但应用全错”这类最隐蔽的问题提前暴露出来。如果你还想更稳可以在 classify 时把每张图的 top-1 类别名也写进日志文件跑完测试后手工抽查几十条比只盯着平均准确率更有说服力。本文还有配套的精品资源点击获取