资讯动态

用PyTorch从零搭建AlexNet实现花卉图像分类的通用框架

发布时间:2026/9/20 14:07:29 来源:尧图企业网站定制
简介面向深度学习入门者与图像分类开发者这套实战资源基于经典卷积神经网络AlexNet提供一套完整可运行的花卉图像分类方案并支持自行更换数据集将模型迁移到自定义分类任务中适合教学实训、课程设计或竞赛练手。压缩包共含2000个文件以1995张JPEG花卉图片为主另附4个Python脚本与1个JSON配置文件总大小约270.61MB。脚本覆盖模型定义、训练、预测与类别索引映射四项核心环节配套数据输入目录与权重保存目录目录结构清晰适合直接运行和二次开发。目前已有119人学习使用。借助这套方案学习者可以完整体验从数据准备、模型训练到分类预测的全流程深入理解卷积层、全连接层、随机失活与线性修正单元等关键设计通过更换数据集还能将方法迁移至人脸识别、物体检测等其他图像分类任务实现举一反三。 用PyTorch从零搭建AlexNet对花卉数据集做分类而且整个代码结构在设计时就考虑了数据集的通用替换——你不需要改模型结构只需要把图片按文件夹放好类别就能自动识别。这篇文章我会把整个项目的来龙去脉、关键代码、训练细节、以及自己数据集怎么换上去全部讲清楚。这个项目适合谁说实话门槛不算高。你只要会Python基础语法用过PyTorch的基本张量操作就能跟着完整跑通。它能解决什么问题说白了就是一套“图像分类的通用工程模板”从数据加载、模型构建、训练验证到预测导出每个环节都封装成可复用的模块。你把花卉换成人脸、零件、场景核心代码几乎不用动。先讲我为什么选AlexNet。很多人一上来就上ResNet、EfficientNet动辄几百层听着很唬人。但AlexNet作为深度学习在图像识别领域真正“破圈”的架构结构简单、训练速度快、对显存友好特别适合作为理解CNN训练全流程的载体。它的参数量约6000万在224x224输入下一张RTX 3060级别显卡就能轻松跑起来。更重要的是它的结构足够经典——5层卷积加3层全连接每一层的作用和特征提取过程都清晰可辨出了问题也容易定位。做项目最怕的就是“代码能跑但对原理一知半解”最后模型崩了都不知道从哪排查。用AlexNet整个网络的行为可解释性很强哪一层干了什么一目了然。这才是选它做实战项目的真正价值。1. 项目整体设计一个能换数据集的通用分类框架1.1 核心设计思路我在动手写代码之前先把整个项目的运行流程过了一遍最终确定了一个非常朴素但好用的原则代码不同数据集耦合。什么意思就是代码里不写死任何类别名称、类别数量、图片路径所有跟数据相关的信息都通过文件夹结构来自动识别。PyTorch的torchvision.datasets.ImageFolder天然支持这种模式你只要把数据集按“类别文件夹”的方式组织好它就会自动扫描出所有子文件夹名作为类别标签。这个设计最大的好处就是不管你今天分类花卉明天分类猫狗后天分类工业缺陷图片代码一行都不用改。只替换数据文件夹然后重新跑训练和评估流程就行。项目整体分五个模块数据准备、数据增强、模型构建、训练验证、预测推理。每个模块对应一个独立脚本或类彼此之间通过清晰的接口衔接。1.2 为什么不用现成的训练脚本直接改有人可能觉得GitHub上现成的图像分类项目一大把直接拿来改改不就行了我试过很多次最后发现大部分现成脚本都有一个问题过度耦合特定数据集。比如有些项目把类别数目写死在模型定义里换数据集就得改模型代码有些项目的预处理方式完全照着ImageNet的规范来换到小数据集反而容易过拟合还有些项目的学习率、训练轮数、batch size都是针对特定数据规模调的直接套用很容易出现梯度爆炸或者欠拟合。所以这个实战项目我采用了“小而稳”的架构模型部分只负责网络结构数据集部分只负责数据加载训练部分只负责优化更新。三者之间的依赖关系降到最低这样每个模块都可以单独替换、单独测试。项目整体文件结构大概这样flower_classify/ ├── data/ │ ├── train/ │ │ ├── rose/ # 类别文件夹名字即标签 │ │ ├── tulip/ │ │ └── sunflower/ │ ├── val/ │ │ ├── rose/ │ │ ├── tulip/ │ │ └── sunflower/ │ └── test/ # 用于最终预测的图片 │ ├── test1.jpg │ └── test2.jpg ├── models/ │ └── alexnet.py # AlexNet网络定义 ├── utils/ │ ├── dataset.py # 数据集加载与增强 │ └── train.py # 训练与验证逻辑 ├── train.py # 训练入口 ├── predict.py # 单图预测入口 └── config.py # 全局配置参数2. 数据集准备从公开数据集到自制数据集的完整方案2.1 公开花卉数据集怎么选如果你不想自己收集图片公开的花卉数据集有现成的。最经典的是Oxford 102 Flower Dataset包含102类英国常见花卉每类有40到258张图片总共约8000多张。这个数据集的难点在于类别多、每类图片少天然适合测试模型的泛化能力和防过拟合能力。另一个选择是Flower Recognition DatasetKaggle上有包含5类花卉雏菊、蒲公英、玫瑰、向日葵、郁金香共约3700张图片。这个数据集更小更简单适合初学者快速跑通全流程一个训练周期也就十几分钟。我强烈建议第一次跑项目的人先用5类的Kaggle数据集把流程彻底跑通确认无误之后再换102类的Oxford数据集或者自己的数据。原因很简单小数据集训练快迭代实验的成本低你能在短时间内验证代码、调参、观察曲线建立起对模型训练的直觉。2.2 自制数据集的目录规范与逐项预检用自己收集的数据第一步就是把图片整理成ImageFolder要求的目录结构。data/ ├── train/ │ ├── 玫瑰/ │ │ ├── 1.jpg │ │ ├── 2.jpg │ │ └── ... │ ├── 向日葵/ │ └── 菊花/ └── val/ ├── 玫瑰/ └── 向日葵/需要注意几个细节类别文件夹的名称就是标签名中文、英文都行但同一个数据集内建议统一语言风格。每类图片数量尽量均衡如果某类图片特别多、某类特别少训练时模型会偏向数量多的类别。图片格式建议统一转成jpg或png避免混用导致读取异常。训练集和验证集不要有重复图片否则验证准确率会虚高你根本看不出真实泛化能力。准备完之后我习惯先写一段快速检查代码把每类的图片数量打印出来、随机可视化几张确认图片能正常读取、标签正确对应。这一步虽然简单但每次换数据集都帮我挡掉了至少一半的潜在问题。import os from collections import Counter data_dir data/train class_counts Counter() for cls in os.listdir(data_dir): cls_path os.path.join(data_dir, cls) if os.path.isdir(cls_path): class_counts[cls] len(os.listdir(cls_path)) print(f共 {len(class_counts)} 个类别总图片数 {sum(class_counts.values())}) for cls, cnt in class_counts.most_common(): print(f {cls}: {cnt} 张)2.3 数据增强小数据集的救命稻草花卉数据集通常不会很大如果你的自制数据集每类只有几十张图片直接拿原始图片去训练CNN过拟合几乎是必然的。过拟合的表现是训练准确率升到95%以上验证准确率却卡在60%左右来回震荡。解决办法是靠数据增强。数据增强的本质是“无中生有”——通过对原始图片做随机变换制造出多样化的训练样本让模型见过更多形态的同一对象。我常用的增强方案如下from torchvision import transforms train_transform transforms.Compose([ transforms.Resize(256), # 先将短边缩放到256 transforms.RandomCrop(224), # 随机裁剪224x224增强平移不变性 transforms.RandomHorizontalFlip(0.5), # 随机水平翻转增强镜像不变性 transforms.ColorJitter(0.2, 0.2, 0.2), # 随机调整亮度、对比度、饱和度 transforms.ToTensor(), # 转为Tensor像素值映射到[0,1] transforms.Normalize( mean[0.485, 0.456, 0.406], # ImageNet均值保持通用性 std[0.229, 0.224, 0.225] # ImageNet标准差 ) ])验证集和测试集不要做随机增强只需要Resize到256再CenterCrop到224然后转Tensor、归一化。随机增强用在验证集上会引入额外的随机性导致评估结果不稳定。这一点很多初学者容易踩坑我记得自己第一次做的时候训练集和验证集共用一套增强结果验证准确率每次跑都不一样还以为是代码bug排查了好久才发现是增强策略的问题。为什么要Resize到256再Crop到224AlexNet的原始输入是224x224但直接把任意尺寸的图片resize到224x224会导致宽高比失真尤其是花卉这种长宽比例不一致的图片会严重影响特征提取效果。先等比缩放到短边256再随机裁剪224x224既保留了宽高比信息又增加了裁剪位置的变化模型学到的特征更鲁棒。3. AlexNet网络结构拆解与PyTorch实现3.1 逐层结构解析AlexNet是2012年ImageNet竞赛的冠军它的历史意义这里不展开我重点讲清楚每层的作用。整个网络就是“卷积提取特征全连接分类”的经典组合Conv111x11卷积核96个步长4。这一层以较大的感受野快速降低空间尺寸提取边缘、颜色斑块等底层特征。输入224x224x3输出55x55x96。论文里原始输入是227x227现在多数实现用224x224都能正常工作。Conv25x5卷积核256个padding2。提取纹理、局部形状等中层特征。Conv3、Conv43x3卷积核384个padding1。继续加深特征抽象层次。Conv53x3卷积核256个padding1后接最大池化。输出6x6x256。FC1、FC24096个神经元的全连接层配合Dropout防止过拟合。FC3输出层神经元数量等于类别数比如花卉5类就是5。整个网络在结构设计上有个关键点层与层之间用ReLU激活函数而不是传统神经网络常用的tanh或sigmoid。ReLU的计算非常简单——大于0保留小于0置0。这个看似简单的操作解决了深层网络训练时梯度消失的问题让网络可以真正“深”下去。池化层使用的是最大池化取窗口内的最大值作为输出。最大池化的好处是保留了最显著的特征响应同时有一部分平移不变性和降采样作用。AlexNet论文里用的是3x3窗口、步长2的重叠池化现在很多实现改成2x2窗口、步长2效果差异不大但重叠池化能稍微保留多一些细节信息。3.2 PyTorch模型代码import torch.nn as nn class AlexNet(nn.Module): def __init__(self, num_classes5): super(AlexNet, self).__init__() self.features nn.Sequential( # Conv1 nn.Conv2d(3, 96, kernel_size11, stride4, padding0), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), # Conv2 nn.Conv2d(96, 256, kernel_size5, stride1, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), # Conv3 nn.Conv2d(256, 384, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), # Conv4 nn.Conv2d(384, 384, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), # Conv5 nn.Conv2d(384, 256, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), ) 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 x.view(x.size(0), -1) x self.classifier(x) return x几个细节说明一下num_classes通过参数传入这样换数据集时直接改config里的数字模型代码不用动。inplaceTrue表示ReLU直接在原Tensor上修改节省显存训练大batch时有一定帮助。两个Dropout放在三个全连接层之间位置是刻意的——FC1和FC2是全连接层的“重灾区”参数量极大最容易过拟合需要在它们后面都做随机失活。原始AlexNet在GPU显存有限的时代把网络拆成上下两半跑在两块GPU上现在单卡显存动不动8G、12G完全没必要复刻那种双分支结构。我现在写的就是“单卡简化版”在ImageNet上准确率比原始论文低一点但在花卉这种小数据集上几乎没有差别。4. 训练流程与关键参数配置4.1 完整训练代码实现训练逻辑我封装成一个标准的PyTorch训练循环。相比很多教程里那种几十行的“丐版”训练脚本这个版本加入了学习率调整、早停、模型保存等实际工程中必须考虑的模块。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder from tqdm import tqdm def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for images, labels in tqdm(dataloader, descTraining): 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) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc correct / total * 100 return epoch_loss, epoch_acc def validate(model, dataloader, criterion, device): model.eval() running_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in tqdm(dataloader, descValidating): images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) running_loss loss.item() * images.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc correct / total * 100 return epoch_loss, epoch_acc主训练入口这样写def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(fUsing device: {device}) train_dataset ImageFolder(rootdata/train, transformtrain_transform) val_dataset ImageFolder(rootdata/val, transformval_transform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue) num_classes len(train_dataset.classes) model AlexNet(num_classesnum_classes).to(device) criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay5e-4) scheduler optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.1, patience5, verboseTrue ) best_val_acc 0.0 epochs 80 for epoch in range(epochs): train_loss, train_acc train_one_epoch( model, train_loader, criterion, optimizer, device) val_loss, val_acc validate(model, val_loader, criterion, device) scheduler.step(val_loss) print(fEpoch [{epoch1}/{epochs}] fTrain Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% fVal Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%) if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_model.pth) print(f Best model saved with val acc {best_val_acc:.2f}%) print(fTraining finished. Best validation accuracy: {best_val_acc:.2f}%)4.2 超参数选择背后的逻辑相比直接抄参数我更想讲清楚每个超参数是怎么定下来的。优化器为什么选SGD而不是Adam这是很多初学者最容易困惑的地方。Adam的好处是收敛快、对学习率不敏感适合快速验证想法。但Adam的泛化性能在不少任务上不如带动量的SGD尤其当数据集比较小、训练轮数比较多的时候SGD加动量通常能收敛到更平坦的极值点泛化效果更好。我从头训练这个项目时如果时间充裕一定用SGD如果只是想快速看个效果就先用Adam跑20个epoch。学习率0.01怎么定的AlexNet原论文就是从0.01开始配合动量和权重衰减。从0.01开始基本不会出大问题但要注意如果换用更大的batch或者加了复杂的预处理可能需要把学习率调小。经验法则是学习率跟batch size正相关batch翻倍学习率大致要相应往上调。weight_decay和momentumweight_decay设为5e-4是图像分类任务的标准配置效果是让权重不会变得过大对过拟合有一定的抑制作用。momentum取0.9也是经典值它的作用是在梯度方向变化不大的维度上加速前进帮助SGD更快穿过平坦区域。Dropout 0.5这是AlexNet论文里的原始设定。Dropout只在训练时生效推理时自动关闭PyTorch的model.eval()会自动处理不需要手动干预。4.3 迁移学习用预训练权重到底香不香对于花卉分类这种小数据集我有一个非常推荐的策略变化如果你的数据集每类只有几十张图片真心建议加载ImageNet预训练权重再在你的数据集上微调。这比从头训练省时省力而且准确率通常高出一截。from torchvision import models model models.alexnet(weightsmodels.AlexNet_Weights.IMAGENET1K_V1) num_features model.classifier[6].in_features model.classifier[6] nn.Linear(num_features, num_classes)这里做的事情很简单加载官方在ImageNet上训练好的AlexNet权重然后只把最后一层全连接替换成你自己的类别数。因为前5层卷积已经学会了通用特征提取能力——边缘、纹理、形状、颜色组合这些特征是跨领域通用的花卉的“花瓣纹理”和ImageNet里“鸟的羽毛纹理”在底层特征上非常接近。微调时有个选择全量微调还是只训练最后一层全量微调适合每类图片在200到500张以上的数据集只训练最后一层适合数据量更小的场景前5层卷积全部冻结只更新最后的全连接层。冻结的实现方法for param in model.features.parameters(): param.requires_grad False只冻结特征提取层的参数优化器只会更新最后一层训练速度飞快显存占用也小。5. 训练结果评估与预测实战5.1 评估指标的解读训练完成之后除了看准确率我还建议看一下混淆矩阵。准确率只能告诉你整体正确率是多少但如果是102类的花卉数据集你很难从0.78的准确率里看出模型到底把哪几类搞混了。import numpy as np from sklearn.metrics import confusion_matrix import matplotlib.pyplot as plt import seaborn as sns def plot_confusion_matrix(model, dataloader, classes, device): model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in dataloader: images images.to(device) outputs model(images) _, predicted torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) plt.figure(figsize(12, 10)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclasses, yticklabelsclasses) plt.xlabel(Predicted) plt.ylabel(True) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150)举例来说我在跑5类花卉数据时混淆矩阵显示玫瑰和郁金香之间的混淆比较多。看了这批图片后发现两者在重叠时期的花型都是卷曲多层结构视觉上确实接近。这时候我已经不纠结“怎么提高准确率”了而是能清楚地定位到模型的“知识盲区”。5.2 单张图片预测流程训练好模型之后最常用的需求是拿一张新图片让模型分类。预测代码需要保持和训练预处理完全一致否则结果会乱七八糟。import torch from PIL import Image from torchvision import transforms from models.alexnet import AlexNet def predict_image(image_path, model_path, classes, devicecpu): 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] ) ]) image Image.open(image_path).convert(RGB) input_tensor transform(image).unsqueeze(0) input_tensor input_tensor.to(device) model AlexNet(num_classeslen(classes)) model.load_state_dict(torch.load(model_path, map_locationdevice)) model.to(device) model.eval() with torch.no_grad(): outputs model(input_tensor) probabilities torch.softmax(outputs, dim1) confidence, pred_idx torch.max(probabilities, 1) pred_class classes[pred_idx.item()] confidence confidence.item() * 100 print(f预测类别: {pred_class}, 置信度: {confidence:.2f}%) return pred_class, confidence if __name__ __main__: classes [daisy, dandelion, rose, sunflower, tulip] predict_image(data/test/test1.jpg, best_model.pth, classes)注意几个细节image.convert(RGB)这一步很关键。有些图片是RGBA模式或者灰度图如果不统一转换输入通道数会跟模型不匹配。torch.softmax把原始输出分数转成概率分布这样你不仅知道模型预测什么还知道它有多大把握。置信度过低比如低于40%时最好人工检查一下图片有可能是模型遇到没见过的新品种。model.eval()必须调用否则前面说的Dropout和BatchNorm在推理时会处于训练模式预测结果不稳定。6. 更换自己的数据集通用配置与避坑实录6.1 换数据集时只需要改动的地方整个项目的核心卖点是“可更换自己数据集”。如果你按照前面说的目录规范整理好数据那么换数据集时只需要改一个文件里的几个配置项# config.py TRAIN_DIR data/train VAL_DIR data/val TEST_DIR data/test BATCH_SIZE 64 EPOCHS 80 LEARNING_RATE 0.01 MODEL_SAVE_PATH best_model.pth NUM_CLASSES 0 # 不用手动填会自动从文件夹数量读取其余代码完全不用动。train_dataset.classes会自动读取文件夹名作为类别列表len(train_dataset.classes)会自动得到类别数量传给模型。换数据集的动作就是把数据文件夹替换掉运行train.py完事。我在实际使用中还遇到过一个需要处理的情况如果你的数据集图片分辨率都很高比如监控截图、数码相机原图建议在预处理里先压缩到合理尺寸。ImageFolder加载的时候会直接读取原始尺寸超大图片非常消耗内存和显存。我的经验是先把图片统一处理成短边不超过512像素再走Resize到256的流程速度和显存占用都会好很多。6.2 常见问题与排查技巧训练过程中有四个问题出现频率最高我把典型表现和排查办法整理成一个速查表现象可能原因排查与解决办法训练loss不降准确率保持在25%左右标签错乱检查数据文件夹是否放对位置每类图片是否对应正确的文件夹打印10个batch的标签核对训练准确率100%验证准确率很低过拟合加大数据增强强度增大Dropout比例减小模型容量或者增加数据量loss直接变成NaN学习率太大或数据有异常值调小学习率检查图片是否有损坏文件检查是否做了归一化显存不足OOMbatch_size太大或图片太大减小batch_size到16或8降低输入分辨率pin_memoryFalse除了表格里的问题还有一个非常隐蔽但出现概率极高的坑训练集和验证集的分布不一致。我做过一次花卉数据集训练集是从网上爬的完美光照条件下的花验证集是自己手机拍的日常光照图片结果训练准确率95%验证准确率只有60%。这不是模型出了问题而是两个集合的分布差异太大了。训练集里模型学到的特征是“光照充足、背景干净的花”真到了自然场景就认不出来。解决思路数据收集时尽量覆盖不同的光照条件、不同的拍摄角度、不同的背景环境让训练集的多样性匹配真实应用场景。数据增强也能部分缓解这个问题但治标不治本采集数据的多样性才是根本。6.3 进一步扩展的思路项目跑通之后你可以在多个方向继续扩展我按难度从低到高排个序把AlexNet换成VGG16、ResNet18对比不同结构的准确率和训练速度。加入学习率预热和余弦退火用更细的优化策略刷高准确率。用torchvision.transforms加更多数据增强手段比如随机旋转、随机裁剪缩放。集成Grad-CAM可视化把模型的注意力热力图画出来直观看到模型关注花的哪个部位。封装成Web服务用Flask或FastAPI做接口上传一张图片返回分类结果。我个人觉得对新手价值最大的是Grad-CAM可视化。它能回答一个关键问题“模型到底看到了什么”如果是通过花瓣颜色来分类你的热力图会集中在花中心如果模型错误地利用了背景信息热力图会集中在背景区域。这种可视化对优化数据集质量和模型调优非常有帮助。最后分享一个实操心得。训练深度学习模型不要只盯着最终准确率把训练过程的曲线记录下来包括训练loss、验证loss、训练准确率、验证准确率。观察曲线之间的差距能告诉你很多信息训练loss降了但验证loss涨了说明过拟合两个loss都居高不下说明模型容量不够或学习率不合适。这个记录曲线的习惯能让你的调试效率至少提高一倍。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价