AlexNet实战:用PyTorch从零搭建花卉分类模型(附完整代码+数据集)

发布时间:2026/7/28 13:38:00

AlexNet实战:用PyTorch从零搭建花卉分类模型(附完整代码+数据集) AlexNet实战用PyTorch从零搭建花卉分类模型深度学习在计算机视觉领域的应用已经变得无处不在而图像分类作为最基础的任务之一仍然是初学者入门的最佳选择。本文将带你从零开始使用PyTorch框架实现经典的AlexNet模型并应用于花卉分类任务。不同于简单的API调用我们会深入模型架构的每个细节让你真正理解卷积神经网络的工作原理。1. 环境准备与数据预处理在开始构建模型之前我们需要准备好开发环境和数据集。PyTorch作为当前最流行的深度学习框架之一以其动态计算图和Pythonic的API设计赢得了大量开发者的青睐。1.1 安装必要的依赖首先确保你已经安装了Python 3.7或更高版本然后通过pip安装以下包pip install torch torchvision pillow matplotlib numpy tqdm对于GPU加速还需要安装对应版本的CUDA和cuDNN。可以使用以下命令检查PyTorch是否正确识别了你的GPUimport torch print(torch.cuda.is_available()) # 应该输出True如果有可用的GPU1.2 准备花卉数据集我们将使用一个包含5类花卉的公开数据集类别包括雏菊(daisy)蒲公英(dandelion)玫瑰(roses)向日葵(sunflowers)郁金香(tulips)数据集的组织结构应该如下flower_data/ train/ daisy/ image1.jpg image2.jpg ... dandelion/ roses/ sunflowers/ tulips/ val/ daisy/ dandelion/ roses/ sunflowers/ tulips/提示可以使用split_data.py脚本自动划分训练集和验证集确保验证集约占全部数据的10%-20%。1.3 数据增强与预处理在深度学习中数据预处理是至关重要的一步。对于图像分类任务我们通常需要进行以下操作from torchvision import transforms # 训练集的数据增强 train_transform transforms.Compose([ transforms.RandomResizedCrop(224), # 随机裁剪并缩放到224x224 transforms.RandomHorizontalFlip(), # 随机水平翻转 transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), # 颜色抖动 transforms.ToTensor(), # 转换为Tensor并归一化到[0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) # ImageNet标准化 ]) # 验证集的预处理不需要数据增强 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]) ])数据增强可以有效防止过拟合特别是在数据集较小的情况下。通过随机变换输入图像我们实际上是在创造新的训练样本。2. AlexNet模型架构详解AlexNet作为深度学习的里程碑在2012年ImageNet竞赛中以显著优势夺冠开启了深度学习在计算机视觉领域的新时代。虽然现在有更先进的模型但理解AlexNet仍然是学习CNN的重要一步。2.1 网络结构分析AlexNet的原始架构包含5个卷积层和3个全连接层由于当时GPU内存限制设计为在两个GPU上并行计算。在我们的实现中我们将所有参数减半以适应现代消费级GPU。完整的AlexNet架构如下表所示层类型参数配置输出尺寸说明输入-3×224×224RGB图像Conv14811×11, stride 4, pad 248×55×55使用大卷积核捕捉大范围特征ReLU-48×55×55非线性激活MaxPool13×3, stride 248×27×27下采样Conv21285×5, pad 2128×27×27中等感受野ReLU-128×27×27非线性激活MaxPool23×3, stride 2128×13×13下采样Conv31923×3, pad 1192×13×13小感受野增加深度ReLU-192×13×13非线性激活Conv41923×3, pad 1192×13×13小感受野增加深度ReLU-192×13×13非线性激活Conv51283×3, pad 1128×13×13小感受野ReLU-128×13×13非线性激活MaxPool33×3, stride 2128×6×6下采样Flatten-4608展平为向量Dropoutp0.54608防止过拟合FC14608→20482048全连接层ReLU-2048非线性激活Dropoutp0.52048防止过拟合FC22048→20482048全连接层ReLU-2048非线性激活FC32048→num_classesnum_classes输出层2.2 PyTorch实现下面是完整的AlexNet实现代码import torch.nn as nn import torch class AlexNet(nn.Module): def __init__(self, num_classes5, init_weightsTrue): super(AlexNet, self).__init__() self.features nn.Sequential( nn.Conv2d(3, 48, kernel_size11, stride4, padding2), # [3,224,224]→[48,55,55] nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), # [48,55,55]→[48,27,27] nn.Conv2d(48, 128, kernel_size5, padding2), # [48,27,27]→[128,27,27] nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), # [128,27,27]→[128,13,13] nn.Conv2d(128, 192, kernel_size3, padding1), # [128,13,13]→[192,13,13] nn.ReLU(inplaceTrue), nn.Conv2d(192, 192, kernel_size3, padding1), # [192,13,13]→[192,13,13] nn.ReLU(inplaceTrue), nn.Conv2d(192, 128, kernel_size3, padding1), # [192,13,13]→[128,13,13] nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size3, stride2), # [128,13,13]→[128,6,6] ) self.classifier nn.Sequential( nn.Dropout(p0.5), nn.Linear(128 * 6 * 6, 2048), nn.ReLU(inplaceTrue), nn.Dropout(p0.5), nn.Linear(2048, 2048), nn.ReLU(inplaceTrue), nn.Linear(2048, num_classes), ) if init_weights: self._initialize_weights() def forward(self, x): x self.features(x) x torch.flatten(x, start_dim1) x self.classifier(x) return x def _initialize_weights(self): for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, modefan_out, nonlinearityrelu) if m.bias is not None: nn.init.constant_(m.bias, 0) elif isinstance(m, nn.Linear): nn.init.normal_(m.weight, 0, 0.01) nn.init.constant_(m.bias, 0)关键点说明inplaceTrue的ReLU可以节省一些内存Kaiming初始化特别适合ReLU激活函数Dropout层只在训练时激活可以防止过拟合最后一层不需要激活函数因为我们将使用CrossEntropyLoss3. 模型训练与调优有了模型架构后我们需要设置训练流程。这部分将详细介绍如何高效地训练AlexNet模型。3.1 训练配置首先设置训练参数和优化器import torch.optim as optim from torch.utils.data import DataLoader # 初始化模型 model AlexNet(num_classes5, init_weightsTrue) model model.to(device) # 移动到GPU如果可用 # 定义损失函数和优化器 criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.0002) # 学习率调度器 scheduler optim.lr_scheduler.StepLR(optimizer, step_size5, gamma0.1)注意Adam优化器通常比原始论文中使用的SGD with momentum表现更好特别是对于初学者。3.2 训练循环完整的训练循环包括以下几个步骤前向传播计算输出计算损失反向传播计算梯度优化器更新权重定期在验证集上评估模型def train_model(model, criterion, optimizer, scheduler, num_epochs10): best_acc 0.0 for epoch in range(num_epochs): print(fEpoch {epoch}/{num_epochs - 1}) print(- * 10) # 每个epoch都有训练和验证阶段 for phase in [train, val]: if phase train: model.train() # 设置模型为训练模式 else: model.eval() # 设置模型为评估模式 running_loss 0.0 running_corrects 0 # 迭代数据 for inputs, labels in dataloaders[phase]: inputs inputs.to(device) labels labels.to(device) # 梯度清零 optimizer.zero_grad() # 前向传播 with torch.set_grad_enabled(phase train): outputs model(inputs) _, preds torch.max(outputs, 1) loss criterion(outputs, labels) # 只在训练阶段反向传播优化 if phase train: loss.backward() optimizer.step() # 统计 running_loss loss.item() * inputs.size(0) running_corrects torch.sum(preds labels.data) if phase train: scheduler.step() epoch_loss running_loss / dataset_sizes[phase] epoch_acc running_corrects.double() / dataset_sizes[phase] print(f{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}) # 深度复制模型 if phase val and epoch_acc best_acc: best_acc epoch_acc torch.save(model.state_dict(), best_model.pth) print() print(fBest val Acc: {best_acc:4f}) return model3.3 训练技巧与调优在实际训练中有几个关键技巧可以提升模型性能学习率调整初始学习率设置为0.0002每5个epoch乘以0.1早停(Early Stopping)如果验证集准确率连续几个epoch不提升可以提前终止训练模型检查点保存验证集上表现最好的模型梯度裁剪防止梯度爆炸# 添加梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 早停实现 patience 3 # 允许的连续不提升epoch数 no_improve 0 best_loss float(inf) for epoch in range(num_epochs): # ...训练代码... val_loss epoch_loss if val_loss best_loss: best_loss val_loss no_improve 0 torch.save(model.state_dict(), best_model.pth) else: no_improve 1 if no_improve patience: print(Early stopping!) break4. 模型评估与预测训练完成后我们需要评估模型性能并进行实际预测。4.1 评估模型性能使用混淆矩阵可以直观地展示模型在各个类别上的表现from sklearn.metrics import confusion_matrix import seaborn as sns import pandas as pd def plot_confusion_matrix(model, dataloader, class_names): model.eval() all_preds [] all_labels [] with torch.no_grad(): for inputs, labels in dataloader: inputs inputs.to(device) labels labels.to(device) outputs model(inputs) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) df_cm pd.DataFrame(cm, indexclass_names, columnsclass_names) plt.figure(figsize(10,7)) sns.heatmap(df_cm, annotTrue, fmtd, cmapBlues) plt.xlabel(Predicted) plt.ylabel(True) plt.show()4.2 单张图像预测下面是一个完整的预测流程可以用于单张图像的分类def predict_image(image_path, model, transform, class_names): # 加载图像 img Image.open(image_path) # 预处理 img_t transform(img) batch_t torch.unsqueeze(img_t, 0).to(device) # 预测 model.eval() with torch.no_grad(): output model(batch_t) # 获取预测结果 _, pred torch.max(output, 1) prob torch.nn.functional.softmax(output, dim1)[0] * 100 # 显示结果 plt.imshow(img) plt.title(fPredicted: {class_names[pred.item()]} ({prob[pred.item()]:.1f}%)) plt.axis(off) plt.show() # 打印所有类别概率 for i, (name, p) in enumerate(zip(class_names, prob)): print(f{name}: {p:.1f}%) return class_names[pred.item()]4.3 可视化中间特征理解CNN工作原理的一个好方法是可视化中间层的特征图def visualize_feature_maps(model, image_path, layer_index0): # 加载并预处理图像 img Image.open(image_path) 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]) ]) img_t transform(img).unsqueeze(0).to(device) # 获取指定层的输出 activation {} def get_activation(name): def hook(model, input, output): activation[name] output.detach() return hook target_layer list(model.features.children())[layer_index] handle target_layer.register_forward_hook(get_activation(fconv{layer_index1})) # 前向传播 with torch.no_grad(): output model(img_t) # 可视化特征图 act activation[fconv{layer_index1}].squeeze().cpu() fig, axarr plt.subplots(act.size(0)//8, 8, figsize(20, 20)) for idx in range(act.size(0)): ax axarr[idx//8, idx%8] ax.imshow(act[idx], cmapviridis) ax.axis(off) plt.tight_layout() plt.show() # 移除hook handle.remove()这个可视化可以帮助我们理解CNN每一层学习到了什么样的特征。通常浅层会学习边缘、颜色等低级特征而深层会学习更抽象的特征。

相关新闻