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

资讯详情

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

PyTorch实现AlexNet花卉图像分类:从数据准备到模型训练部署全流程

PyTorch实现AlexNet花卉图像分类:从数据准备到模型训练部署全流程 简介以AlexNet模型为核心的花卉分类实战项目面向深度学习初学者及图像分类开发者解决从数据准备、模型训练到结果预测的全流程实践难题并支持通过替换数据集快速迁移到其他分类任务。压缩包共2000个文件整体约270.61MB主要包含1995张花卉图片、4个Python脚本和1个JSON配置文件。其中model.py定义网络结构train.py负责加载数据并启动训练predict.py完成单图或批量的分类预测class_indices.json则记录类别与索引映射图片数据可直接用于训练与验证结构清晰便于二次修改。目前已有119人学习适合想通过项目实战理解AlexNet卷积层、全连接层、ReLU与Dropout机制的用户。按照资源内代码与目录组织读者可完成完整训练流程掌握保存与恢复模型参数的方法并将同一套方案应用到更多图像分类领域。 这批花真是把我折磨得不轻。前阵子接到个图像分类的需求数据集是常见的花卉图片要求先把整条流程跑通后面还要能无缝换成客户自己的数据。我第一反应就是拿AlexNet当基线。别嫌它老这网络放到今天做中小规模分类依然是块好用的试金石结构简单、显存占用不高、思路清晰模型出问题了好排查跑一版基线快得很。这篇文章就把这个项目从数据准备、网络搭建、训练评估到怎么把数据集替换成你自己的完整过一遍代码都是可以直接抄作业的级别。1. 项目整体设计与思路拆解1.1 为什么选AlexNet做花卉分类基线选AlexNet不是因为新潮恰恰是因为它足够经典。花卉分类这种任务的特点是类别之间差异细腻但图像本身结构相对固定背景复杂度和类别数都比ImageNet低一个量级。用ResNet、EfficientNet当然可以但模型复杂度上去了训练时间变长调参成本也高。对一个需要快速验证效果的基线项目来说AlexNet在精度和效率之间拿捏得刚刚好。另外它特别适合作为教学和工程起步的骨架。AlexNet整条前向传播路径很直白先是卷积层抓局部特征再通过全连接层做高阶语义组合最后接softmax输出类别概率。出了问题你一眼就能定位是特征提取环节还是分类器环节的事。换成自己的数据集时只需要改动最后一层的输出维度其余结构不用动这个特性在实际项目里非常省事。1.2 数据集方案与项目结构规划花卉分类的数据集我建议直接用两种方案一是公开数据集如Oxford 102 Flower或17 Category Flower二是自己拍或者爬整理出来的图片集。刚开始别贪多每个类别先凑100-200张图把流程跑通再说。数据集的组织形式直接决定代码复杂度我强烈建议统一用ImageFolder的标准目录结构flower_data/ ├── train/ │ ├── rose/ │ │ ├── 001.jpg │ │ └── 002.jpg │ ├── sunflower/ │ └── daisy/ └── val/ ├── rose/ └── sunflower/这种结构最大的好处是PyTorch的torchvision.datasets.ImageFolder可以直接读类别名从子目录名自动生成不需要手写label映射表。整个项目的目录规划大致是alexnet_flower/ ├── data/ # 存放数据集 ├── models/ # 网络定义 │ └── alexnet.py ├── train.py # 训练脚本 ├── predict.py # 单张图片预测 └── requirements.txt2. 数据集准备与预处理细节2.1 目录结构与类别映射原理ImageFolder的机制值得说透它会扫描根目录下的每个子文件夹按文件夹名字母顺序排序自动分配label索引。这个顺序容易踩坑比如daisy、rose、sunflower按字母排下来索引分别是0、1、2。如果后续你自己写预测脚本一定要保证类别索引映射和训练时一致最稳妥的做法是把dataset.class_to_idx保存成json文件预测时直接加载。另外数据清洗往往被忽略。我拿到手的数据集里经常混着损坏图片、重复图片和完全不相关的图。建议先写一段脚本用PIL.Image.open逐个尝试打开捕获异常把打不开的文件列出来删掉。这一步看着笨却能省下后面训练时反复报错的烦恼。2.2 数据增强与归一化参数选择AlexNet论文里输入是224x224但原版训练时会先resize到256然后随机裁剪224。这个策略我会保留因为它等于在训练时给模型看了原图不同位置的局部内容相当于免费扩充了数据。完整的数据增强配置我这样写from torchvision import datasets, transforms transform_train transforms.Compose([ transforms.Resize(256), transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), 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]) ]) transform_val 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]) ])归一化用的mean和std是ImageNet的标准值如果你换自己的数据集且图片风格差异很大建议在训练集上重新统计一遍。但现实情况是大多数自然图像用ImageNet的统计值都能正常工作。验证集只做CenterCrop不做随机增强这个区别很重要否则验证指标会忽高忽低失去参考意义。3. AlexNet核心结构与关键代码实现3.1 网络结构逐层拆解AlexNet整体是5个卷积层加3个全连接层。前两层卷积后面跟了局部响应归一化LRN和最大池化中间三层卷积直接相连最后接最大池化。这几个设计在当年都是很前卫的卷积核大小从11x11、5x5、3x3逐渐缩小对应从抓全局轮廓到抓局部纹理的过渡。ReLU激活函数解决深层网络梯度饱和问题训练速度快很多。Dropout只加在全连接层且概率设为0.5因为全连接层参数量巨大最容易过拟合。原版用了两块GPU并行训练现在单卡显存足够可以不考虑这个分支逻辑。动手实现时有个关键点很容易被忽视原版第一层卷积的感受野很大stride4这对于分辨率较低的图片会把细节一下子冲掉。所以如果你的数据集图片只有128x128左右建议把第一层stride改成2或者干脆把输入resize到224再送进去。我这次数据集图片比较大就保持了原版参数。3.2 PyTorch实现与动态类别数适配直接给一个完整可用的PyTorch版本import torch.nn as nn class AlexNet(nn.Module): def __init__(self, num_classes1000): super(AlexNet, self).__init__() self.features nn.Sequential( nn.Conv2d(3, 96, kernel_size11, stride4, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), nn.Conv2d(96, 256, kernel_size5, stride1, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), nn.Conv2d(256, 384, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(384, 384, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(384, 256, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), ) self.avgpool nn.AdaptiveAvgPool2d((6, 6)) self.classifier nn.Sequential( nn.Dropout(0.5), nn.Linear(256 * 6 * 6, 4096), nn.ReLU(inplaceTrue), nn.Dropout(0.5), nn.Linear(4096, 4096), nn.ReLU(inplaceTrue), nn.Linear(4096, num_classes), ) def forward(self, x): x self.features(x) x self.avgpool(x) x torch.flatten(x, 1) x self.classifier(x) return x这里有个细节值得展开原版AlexNet全连接层前接的是展平后的特征图尺寸必须是固定的。我用了AdaptiveAvgPool2d((6, 6))替代原版的直接展平好处是输入图片尺寸即使不是224只要接近也能自适应地池化成6x6再进全连接层。这样模型对输入尺寸的容忍度高了不少。num_classes参数就是为换数据集留的口子。实例化的时候直接从数据集的类别数读取num_classes len(train_dataset.classes) model AlexNet(num_classesnum_classes)这样无论你的数据集是3类还是102类模型都能自动适配不需要改网络结构。4. 训练流程、超参数配置与评估4.1 超参数配置思路训练超参数我实测下来有一组比较稳的配置优化器用SGDmomentum设为0.9weight_decay设为5e-4初始学习率0.001batch size看显存情况选32或64训练30-50个epoch。这组参数不是拍脑袋定的和AlexNet原论文一脉相承。学习率策略我建议用StepLR每10个epoch把学习率乘以0.1。实际训练中你会发现到后期loss下降很慢这时候把学习率降一档loss经常能再往下走一段。另外不要一开始就用Adam。Adam收敛快但容易收敛到泛化性能不那么好的点SGD配合动量虽然看着慢但最终精度往往更高尤其对于这种中小规模数据集。训练循环里有两个关键点务必注意第一训练模式下要调用model.train()验证模式下要调用model.eval()这会切换Dropout和BatchNorm的行为。第二验证阶段用torch.no_grad()包裹否则会额外占用显存还可能因为保存计算图导致内存暴涨。4.2 训练主循环与评估指标解读训练主循环给个精简版本import torch import torch.nn as nn from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model AlexNet(num_classesnum_classes).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.001, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) best_acc 0.0 for epoch in range(30): model.train() running_loss 0.0 for images, labels in train_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() * images.size(0) scheduler.step() epoch_loss running_loss / len(train_dataset) model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() val_acc 100.0 * correct / total print(fEpoch {epoch1:02d} | Loss: {epoch_loss:.4f} | Val Acc: {val_acc:.2f}%) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth)评估指标上除了准确率我建议再关注每个类别的召回率。花卉分类场景里类别不均衡很常见比如某个品种图片特别少整体准确率可能被多数类带高但这种模型对少数类几乎不可用。训练结束后写个小脚本把验证集的混淆矩阵打印出来一眼就能看出哪些类别之间互相混淆比如玫瑰和月季这种视觉上极接近的品种这时候就需要增加对应类别的样本量或者调整数据增强的强度。上一步模型保存要注意只保存state_dict()而不是整个模型这样后续加载时结构变了也能灵活适配。加载时要先实例化模型再load_state_dict。5. 换用自己的数据集方法与常见问题实录5.1 数据集替换的核心操作换数据集的流程其实已经被前面的代码设计好了核心就三步把你的图片按类别放到data/train/{类别名}/和data/val/{类别名}/目录下。确保所有图片格式统一jpg、png都行但最好统一一种省得处理通道数不一致的麻烦。运行脚本时确认num_classes自动变成你的类别数。不过有几类特殊情况需要特殊处理。如果图片数量特别少比如每个类别只有三五十张直接硬训AlexNet几乎必然过拟合。这种情况建议先用ImageNet上预训练好的AlexNet权重初始化只随机初始化最后一层然后以较小的学习率0.0001微调整个网络。PyTorch加载预训练权重的方式很简单注意要过滤掉最后一层import torchvision.models as models pretrained models.alexnet(weightsmodels.AlexNet_Weights.IMAGENET1K_V1) model_dict model.state_dict() pretrained_dict {k: v for k, v in pretrained.state_dict().items() if k in model_dict and classifier.6 not in k} model_dict.update(pretrained_dict) model.load_state_dict(model_dict)如果图片是灰度图比如一些老照片数据集加载时会报通道不匹配。解决办法是读图后用convert(RGB)转成三通道或者把第一层卷积的in_channels改成1并重新初始化该层参数。5.2 常见问题与排查技巧实录训练过程中我遇到的坑不少整理几个最常见的loss不降反升。先检查数据和标签的对应关系。比如ImageFolder按字母序分配label如果你的目录结构和预期不一致模型学到的就是错误映射关系。这种问题通常loss一开始就不正常不会降。再看学习率0.001不行就降到0.0001试试。训练集准确率高但验证集准确率低。这是典型的过拟合。优先增加Dropout强度或者把数据增强开猛一点比如加上RandomRotation(15)和随机擦除。如果是小数据集考虑用预训练权重微调别从头训练。显存不足。最简单粗暴的方法是batch size从64降到32或16。还可以把输入图片从224x224降到128x128AlexNet对输入尺寸的适应性比想象中强只是精度会稍有损失。另外num_workers不要开太大4或者8足够太大反而可能因为系统调度问题拖慢速度。训练到一半loss变成NaN。大概率是学习率过高导致梯度爆炸。先降学习率再检查数据里有没有异常值。还有一种情况是数据归一化参数写错输入变成很大的负数激活值异常。换成自己数据集后准确率只有百分之十几。这个大概率是数据集和预训练模型分布差异太大或者类别数远大于样本量。我的建议是先别追求精度把数据做了可视化检查随机抽几十张训练图片看有没有贴错标签、有没有加载成黑白、有没有resize变形。数据问题的优先级永远高于调参。6. 一点心得写在后头AlexNet跑花卉分类这个项目做完之后我最大的体会就是经典结构比花哨结构更容易定位问题。很多新手一上来就挑战Swin Transformer结果网络结构本身就把人绕晕了出了问题根本不知道是数据的事还是模型的事。从小而可靠的网络开始把数据流程、训练逻辑、评估方法吃透再迁移到复杂模型这条路快得多。最后再说一个实用小技巧训练脚本里把每个epoch的损失和准确率写入CSV文件训练完直接画学习曲线。这个东西的价值在于你能直观看到模型是否还在学习、学习率降的时机是否合适甚至可以和之后的实验做对比。很多时候调参不是靠感觉就是靠这些不起眼的记录。这个项目往下的扩展方向也很多比如把features部分换成ResNet或者MobileNet的骨干网络做对比实验再比如用Flask封装一个上传图片返回分类结果的服务。骨架已经打好了换起来都不难。本文还有配套的精品资源点击获取
返回列表