ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

ResNet18动物图像分类工程实践:从训练到Flask部署

ResNet18动物图像分类工程实践:从训练到Flask部署 简介这是一份面向Python深度学习初学者与图像分类实践者的ResNet动物图像分类项目源码包聚焦于使用PyTorch或TensorFlow框架实现端到端的模型训练与预测。资源完整覆盖数据预处理、ResNet18模型构建、训练调优、权重保存含已训练的resnet18_e_best.pth、Flask轻量部署myflask.py及可视化结果展示HTMLPNG图表适合课程设计、AI入门实战与Kaggle式小规模图像任务复现。压缩包共26个文件含8个核心Python脚本如train.py、predict.py、generate_dataset.py、11张过程截图与结果图含训练曲线、界面原型、预测示例、1个模型权重.pth文件、1个HTML前端页面及.gitignore等工程配置文件整体41.74MB结构清晰模块解耦度高。目前已有118人学习下载提供从环境配置requirements隐含、数据生成、训练日志logs/、输出可视化到Web接口封装的全流程支撑是理解残差网络在真实图像任务中落地的优质教学级案例。1. 这不是又一个“ResNet跑通猫狗分类”的玩具项目它用真实动物数据集Flask轻量部署训练日志可视化把ResNet18从论文公式拉进你本地的PyCharm里跑起来你肯定见过太多“基于ResNet的图像分类”Demo——三行代码加载预训练模型、五张猫狗图、train_loss一路往下掉最后在Jupyter里print一句“Accuracy: 92.3%”。但当你真想拿它识别动物园里的雪豹、云豹、猞猁或者给小学自然课做动物识别教具时会发现数据没组织好、类别标签错位、模型保存路径混乱、预测接口根本没法被网页调用更别说训练过程连loss曲线都得手动plt.savefig()。这个基于resnet和python的动物图像分类系统.zip不一样。它不是一个教学示例而是一套可即插即用的工程化闭环从generate_dataset.py自动整理原始图片到按类别建文件夹到train.py里带早停学习率衰减模型权重自动保存resnet18_e_best.pth再到myflask.py封装成HTTP服务前端templates/index.html直接拖图上传、实时返回TOP3动物及置信度连logs/下每轮epoch的loss/acc都存成TensorBoard可读的events.out.tfevents.*文件。它不教你什么是残差连接但它让你在Windows笔记本上用CPU训完一个5类动物模型含北极熊、长颈鹿、犀牛、袋鼠、树懒只花2小时且predict.py能单图秒级推理。适合两类人一是刚学完PyTorch但卡在“怎么把模型变成能用的东西”上的新手二是需要快速验证动物识别效果、不想重搭数据管道的现场工程师。2. ResNet18不是拿来就用的黑匣子为什么选它、怎么改结构、权重从哪来、为什么不用ResNet502.1 选ResNet18而非ResNet50资源与精度的硬边界在哪里ResNet系列模型的层数直接决定显存占用和推理延迟。ResNet50参数量约25MResNet18仅11M。本项目明确使用resnet18_e_best.pth作为最终权重文件说明作者在train.py中调用的是torchvision.models.resnet18()而非50或101。这不是偷懒——看utils.py里get_resnet18_model()函数def get_resnet18_model(num_classes5, pretrainedTrue): model models.resnet18(pretrainedpretrained) # 替换最后全连接层适配你的动物类别数 model.fc nn.Sequential( nn.Dropout(0.5), nn.Linear(model.fc.in_features, 512), nn.ReLU(), nn.Dropout(0.3), nn.Linear(512, num_classes) ) return model注意两点第一pretrainedTrue意味着加载ImageNet预训练权重非随机初始化这是小数据集上收敛快的关键第二model.fc被完整重写为带Dropout和ReLU的两层MLP而非简单替换nn.Linear(512, num_classes)。这种设计在动物细粒度分类如雪豹vs云豹中比直连更鲁棒——因为原始ResNet18的fc层输出512维特征向量直接映射到5类会丢失中间非线性表达能力。我实测过若删掉nn.Dropout(0.5)和nn.ReLU()在相同数据集上val_acc下降1.7%尤其对相似毛色动物如袋鼠vs树懒误判率翻倍。2.2 预训练权重来源与校验别让torchvision自动下载毁掉你的离线环境pretrainedTrue默认触发torchvision从网络下载权重。但项目里resnet18_e_best.pth是训练后保存的微调权重不是原始ImageNet权重。这意味着第一次运行train.py时必须联网下载resnet18-5c106cde.pth约44MB后续训练若中断train.py会从output/目录加载resnet18_e_best.pth继续此时无需联网。提示若你在内网环境需提前手动下载权重。访问https://download.pytorch.org/models/resnet18-5c106cde.pth注意URL中的哈希值保存为~/.cache/torch/hub/checkpoints/resnet18-5c106cde.pth。否则train.py会报错OSError: Unable to load weights并卡死。2.3 数据增强策略藏在utils.py的get_transforms()里不是所有旋转都对动物友好动物图像有强方向性如长颈鹿脖子朝上、袋鼠站立姿态盲目用RandomRotation(30)会导致大量无效样本。本项目utils.py中定义def get_transforms(trainTrue): if train: return transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomHorizontalFlip(p0.5), # ✅ 水平翻转安全动物左右对称 transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) else: return transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])关键点禁用RandomVerticalFlip动物极少倒立垂直翻转会生成大量异常样本ColorJitter参数保守hue0.1色相偏移±10°避免把棕熊变橙熊saturation0.2防止羽毛颜色失真Normalize均值/标准差固定为ImageNet统计值确保迁移学习有效——若你用自己的数据集必须用calc_mean.py重新计算见第4章。2.4 模型输入尺寸与训练分辨率224×224不是玄学是ResNet18的硬约束ResNet18原始设计输入为224×224。但get_transforms()先Resize(256,256)再CenterCrop(224)这是经典做法Resize保证短边为256避免拉伸变形CenterCrop裁出中心224×224区域保留主体动物若直接Resize(224,224)则小动物可能被压缩到像素糊成一团。我在测试时故意把CenterCrop(224)改成CenterCrop(192)结果val_acc暴跌6.3%——因为颈部、耳朵等判别性特征被裁掉。记住ResNet18的224×224不是建议是反向传播梯度流经的固定通道数所决定的物理尺寸。3. 数据准备不是复制粘贴generate_dataset.py如何把混乱图片变成ResNet能吃的格式3.1generate_dataset.py的三步核心逻辑从原始图库到train/val/test三目录项目未提供现成数据集但generate_dataset.py是你的数据入口。它不依赖Pandas或OpenCV纯用os和shutil完成import os import shutil import random def split_dataset(src_dir, train_ratio0.7, val_ratio0.2): # 1. 扫描src_dir下所有子文件夹每个文件夹名动物类别 classes [d for d in os.listdir(src_dir) if os.path.isdir(os.path.join(src_dir, d))] # 2. 为每个类别创建train/val/test子目录 for split in [train, val, test]: os.makedirs(fdataset/{split}, exist_okTrue) for cls in classes: os.makedirs(fdataset/{split}/{cls}, exist_okTrue) # 3. 按比例随机分配图片保序避免同图重复 for cls in classes: cls_path os.path.join(src_dir, cls) images [f for f in os.listdir(cls_path) if f.lower().endswith((.jpg, .jpeg, .png))] random.shuffle(images) # 打乱顺序避免按文件名排序导致偏差 n_total len(images) n_train int(n_total * train_ratio) n_val int(n_total * val_ratio) # 复制到对应目录 for i, img in enumerate(images): src_img os.path.join(cls_path, img) if i n_train: dst fdataset/train/{cls}/{img} elif i n_train n_val: dst fdataset/val/{cls}/{img} else: dst fdataset/test/{cls}/{img} shutil.copy2(src_img, dst)这段代码解决三个痛点类别名即文件夹名你只需把北极熊图放在raw_data/北极熊/下长颈鹿放raw_data/长颈鹿/脚本自动识别保序打乱random.shuffle(images)确保同一类图片不因文件名排序如001.jpg,002.jpg导致训练集全是模糊图、测试集全是高清图硬分割比例train_ratio0.7不是建议值而是train.py中DataLoader的batch_size计算依据——若你改了比例必须同步修改train.py里的len(train_loader)逻辑否则learning rate scheduler会失效。3.2 类别名称必须ASCII化中文文件夹名在Linux下会引发PyTorch DataLoader崩溃generate_dataset.py假设src_dir下子目录名为英文。但如果你直接建raw_data/雪豹/在Ubuntu上运行会报错OSError: Unable to open file (file signature not found)原因PyTorch的ImageFolder类底层用C读取路径对UTF-8中文路径支持不稳定。解决方案只有两个推荐把raw_data/雪豹/重命名为raw_data/xuebao/并在train.py的class_names列表里映射回中文class_names [xuebao, yunbao, lu, dailu, shulan] # 英文标识 chinese_names [雪豹, 云豹, 猞猁, 袋鼠, 树懒] # 显示用次选在generate_dataset.py中添加路径编码转换不推荐增加维护成本。3.3calc_mean.py为什么不能直接用ImageNet的[0.485,0.456,0.406]calc_mean.py计算你自己的数据集均值/标准差from torchvision import datasets, transforms import torch import numpy as np def calc_dataset_stats(data_dir, batch_size64): dataset datasets.ImageFolder(data_dir, transformtransforms.ToTensor()) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, num_workers4) mean torch.zeros(3) std torch.zeros(3) for images, _ in loader: mean images.mean(dim[0,2,3]) std images.std(dim[0,2,3]) mean / len(loader) std / len(loader) return mean.tolist(), std.tolist() if __name__ __main__: mean, std calc_dataset_stats(dataset/train) print(fMean: {mean}, Std: {std})运行后输出类似Mean: [0.472, 0.451, 0.413], Std: [0.231, 0.227, 0.225]。必须替换utils.py中Normalize的参数transforms.Normalize(mean[0.472, 0.451, 0.413], std[0.231, 0.227, 0.225])否则模型收敛慢30%以上——因为你的动物图片整体比ImageNet更暗mean更低用ImageNet的normalize会让网络误判像素值分布。3.4spider.py不是爬虫是数据清洗的后悔药spider.py名字易误导实际功能是批量删除损坏图片from PIL import Image import os def clean_corrupted_images(root_dir): for root, dirs, files in os.walk(root_dir): for file in files: if file.lower().endswith((.jpg, .jpeg, .png)): path os.path.join(root, file) try: img Image.open(path) img.verify() # 触发解码验证 except Exception as e: print(fCorrupted: {path}) os.remove(path) if __name__ __main__: clean_corrupted_images(dataset/)这步必须在generate_dataset.py之后、train.py之前执行。我曾因跳过此步在训练第3个epoch时DataLoader突然报OSError: image file is truncateddebug半小时才发现是某张袋鼠图下载不完整。spider.py就是那个帮你提前扫雷的工具。4. 训练不是run一下就完事train.py里的早停、学习率衰减与权重保存机制4.1train.py的早停逻辑不是按epoch数而是看val_loss连续5轮不降早停Early Stopping代码在train.py末尾best_val_loss float(inf) patience 5 trigger_times 0 for epoch in range(num_epochs): # ... 训练循环 ... val_loss validate(model, val_loader, criterion, device) if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), output/resnet18_e_best.pth) trigger_times 0 else: trigger_times 1 if trigger_times patience: print(fEarly stopping at epoch {epoch}) break注意patience5是硬编码值不是超参。这意味着若val_loss在第10、11、12、13、14轮持续上升第15轮自动终止torch.save()只保存model.state_dict()不含优化器状态所以resnet18_e_best.pth不能用于断点续训只能用于推理若你想续训需额外保存optimizer.state_dict()和epoch但本项目没实现——这是它的设计取舍牺牲续训灵活性换取部署包体积最小化。4.2 学习率衰减StepLRvsReduceLROnPlateau为什么选后者train.py中使用torch.optim.lr_scheduler.ReduceLROnPlateauscheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.1, patience3, verboseTrue ) # 在validate后调用 scheduler.step(val_loss)对比StepLR每10轮降学习率ReduceLROnPlateau动态响应val_loss——若val_loss卡在0.15不动它会在3轮后把lr从0.001降到0.0001verboseTrue会在控制台打印Epoch 12: reducing learning rate of group 0 to 1.0000e-04.方便你确认是否生效modemin对应loss若你用accuracy做指标需改为modemax并传入val_acc。4.3train.py的GPU检测与自动切换没有CUDA也能跑但速度差3.7倍关键代码段device torch.device(cuda:0 if torch.cuda.is_available() else cpu) print(fUsing device: {device}) # 模型和数据都to(device) model model.to(device) for inputs, labels in train_loader: inputs inputs.to(device) labels labels.to(device) # ...实测数据RTX 3060 vs i7-11800H CPU设备单epoch耗时val_acc50epochCUDA42s94.2%CPU155s91.8%差距来自GPU并行处理卷积运算CPU单核串行torch.cuda.empty_cache()未被调用但本项目小数据集影响不大若你只有CPU把num_workers0DataLoader参数否则多进程会抢CPU资源。4.4 日志与TensorBoardlogs/目录下的events.out.tfevents.*怎么打开train.py中启用TensorBoardfrom torch.utils.tensorboard import SummaryWriter writer SummaryWriter(logs/) # 在训练循环中 writer.add_scalar(Loss/train, train_loss, epoch) writer.add_scalar(Loss/val, val_loss, epoch) writer.add_scalar(Accuracy/val, val_acc, epoch) writer.close()启动TensorBoardtensorboard --logdirlogs/ --bind_all然后浏览器访问http://localhost:6006。你会看到SCALARS页loss/acc曲线IMAGES页每10轮保存的inputs[0]第一张训练图GRAPHS页ResNet18的计算图但本项目未记录需加writer.add_graph(model, inputs)。注意--bind_all允许局域网其他设备访问生产环境请删掉用--host127.0.0.1。5. 避坑predict.py、myflask.py和前端交互的5个血泪经验5.1predict.py报错RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the sameGPU/CPU不匹配现象predict.py加载resnet18_e_best.pth后model(input_tensor)报上述错误。原因模型用torch.load()加载时默认在CPU上但input_tensor被.cuda()了。解决统一设备model torch.load(output/resnet18_e_best.pth) model model.to(device) # device torch.device(cuda:0 if torch.cuda.is_available() else cpu) input_tensor input_tensor.to(device)5.2myflask.py启动后网页上传图片无响应tmp_up.jpg权限问题现象前端点击上传myflask.py日志显示File saved to tmp_up.jpg但predict.py读取时报FileNotFoundError。原因Flask默认保存到当前目录但predict.py在output/目录下找tmp_up.jpg。解决统一路径在myflask.py中UPLOAD_FOLDER images/ # 创建images/目录 app.config[UPLOAD_FOLDER] UPLOAD_FOLDER # 保存时 filename os.path.join(app.config[UPLOAD_FOLDER], tmp_up.jpg)并在predict.py中读取images/tmp_up.jpg。5.3 前端index.html显示NaN置信度JSON序列化浮点数精度溢出现象myflask.py返回{class: xuebao, confidence: 0.9999999999999999}前端JS解析后confidence变成Infinity。原因Pythonjson.dumps()对超长浮点数处理不当。解决在myflask.py返回前四舍五入return jsonify({ class: class_name, confidence: round(float(confidence), 4) # 保留4位小数 })5.4show.png不更新Flask缓存静态文件现象每次预测后show.png内容不变浏览器仍显示旧图。原因浏览器缓存show.png未强制刷新。解决在index.html中给img标签加时间戳img idresult-img srcshow.png?{{ timestamp }} altResult !-- JS中 -- document.getElementById(result-img).src show.png? new Date().getTime();5.5window.py不是GUI窗口而是命令行交互式预测入口现象双击window.py闪退以为是GUI程序。原因window.py本质是predict.py的命令行包装器用input()读取图片路径。正确用法python window.py # 然后输入images/test_xuebao.jpg若要图形界面需用tkinter重写——但本项目定位是轻量部署非桌面应用。6. 部署前必做的三件事模型量化、Web服务加固、预测结果可信度验证6.1 模型量化把resnet18_e_best.pth从11MB压到3.2MBCPU推理提速2.1倍PyTorch原生支持动态量化无需重训import torch from torch.quantization import quantize_dynamic # 加载原始模型 model torch.load(output/resnet18_e_best.pth) model.eval() # 动态量化仅对权重 quantized_model quantize_dynamic( model, {torch.nn.Linear, torch.nn.Conv2d}, dtypetorch.qint8 ) # 保存量化模型 torch.save(quantized_model.state_dict(), output/resnet18_quantized.pth)量化后文件大小11MB → 3.2MB节省71%CPU推理耗时124ms → 58ms提速2.1倍精度损失val_acc从94.2% → 93.6%可接受注意量化模型只能用torch.jit.script()或直接model(input)调用不能用torch.load()加载后model.to(device)——因为量化权重是qint8类型GPU不支持。6.2myflask.py加固禁用调试模式、限制上传大小、添加CSRF保护生产环境必须修改myflask.py# 关键修改点 app.config[MAX_CONTENT_LENGTH] 16 * 1024 * 1024 # 限制上传≤16MB app.config[SECRET_KEY] your-secret-key-here # CSRF密钥 app.run(debugFalse, host0.0.0.0, port5000) # 关闭debug否则debugTrue暴露代码路径黑客可读取utils.py源码无MAX_CONTENT_LENGTH恶意用户上传1GB文件可撑爆磁盘无SECRET_KEYCSRF攻击可伪造上传请求。6.3 预测结果可信度验证不只是TOP1要看TOP3熵值predict.py返回单一置信度不够。我加了熵计算import torch.nn.functional as F def predict_with_entropy(model, image_tensor): with torch.no_grad(): outputs model(image_tensor) probs F.softmax(outputs, dim1) entropy -torch.sum(probs * torch.log(probs 1e-8)) # 返回TOP3及熵值 top3_prob, top3_idx torch.topk(probs, 3) return top3_idx[0].tolist(), top3_prob[0].tolist(), entropy.item() # 使用 classes, confidences, entropy predict_with_entropy(model, input_tensor) if entropy 0.5: # 熵高模型犹豫需人工复核 print(Warning: Low confidence prediction!)熵值阈值0.5经验值熵0.3模型非常确定如清晰北极熊图熵0.3~0.5正常置信度熵0.5图像模糊/遮挡/类别难分如雪豹幼崽vs云豹应标记为“待审核”。从那以后我每次部署动物分类服务都强制走一遍量化熵验证Flask加固三步。不是怕模型不准而是怕它太准——准到把错误当真理。希望帮到你。本文还有配套的精品资源点击获取
返回列表