从AlexNet到ViT:PyTorch统一训练模板与模型部署实践
发布时间:2026/9/26 11:30:47来源:尧图网络
1. 背景与核心概念如果现在要评选过去十年影响最深远的深度学习模型卷积神经网络Convolutional Neural NetworkCNN一定是最有竞争力的候选之一。从 2012 年 AlexNet 在 ImageNet 大赛上一举夺冠开始CNN 逐步成为图像分类、目标检测、语义分割、人脸识别等任务的默认底座。即使后来 Transformer 大行其道CNN 在很多场景下仍然是训练稳定、推理高效、部署方便的首选。本文希望围绕一条主线展开从 CNN 的基本结构出发梳理 AlexNet、VGG、ResNet 和 ViT 这几类经典模型之间的关系再给出一份可以一次训练、自由切换主干的 PyTorch 代码模板最后聊一聊如何把训练好的模型导出、部署成本地推理服务。内容适合刚接触深度学习的同学也适合想快速跑通实验结果的开发者。先来补充一点必要的背景。CNN 本质上是一个“局部连接、权值共享”的神经网络它不像全连接网络那样把上一层的每个神经元都与下一层的每个神经元连接而是用一组小尺寸的卷积核在输入特征图上滑动从而提取局部特征。这种设计有两个直接的好处其一参数数量大幅下降模型更容易训练其二卷积核天然具有平移等变性也就是说物体在图片中移动几个像素提取到的特征依然类似。这两个特性使 CNN 在图像任务上远优于传统的全连接网络。为了控制特征图尺寸并提高非线性表达能力CNN 还会配合池化层和激活函数使用最终通过全连接层或者全局平均池化输出分类结果。学习 CNN 有一个完整的训练闭环数据集准备、数据预处理与增强、网络前向传播、损失计算、反向传播、参数更新、周期性验证、模型保存。很多人一开始只关注模型结构却忽略了数据处理和训练策略导致同样的代码在别人的电脑上能收敛在自己这里就发散。下面我们会先把模型演进讲清楚然后直接用一个统一代码模板把整个闭环串起来。在实际项目中这套流程可以帮助你快速验证一个新的网络结构是否适合当前任务也可以作为工程化改造的起点。2. 经典模型演进AlexNet/VGG/ResNet/ViT2.1 AlexNet深度学习引爆点AlexNet 出现在 2012 年是第一个在 ImageNet 大规模图像识别竞赛中取得碾压性成绩的 CNN。它的核心贡献不是某个单独的数学技巧而是把深层 CNN 的工程细节整合到了一起使用 ReLU 作为激活函数缓解梯度消失使用 Dropout 减少过拟合使用重叠池化增强特征同时借助 GPU 并行训练加速。从今天的眼光看AlexNet 的 5 层卷积加 3 层全连接并不算深但它证明了只要数据量足够大、算力足够强深度模型是可以被有效训练的。理解 AlexNet重点是理解它奠定了“卷积提取特征 全连接分类”的基本范式。在代码层面AlexNet 的输入尺寸通常是 224×224 的 RGB 图像输出 1000 类概率。它的第一个卷积层采用 11×11 的较大卷积核步长为 4目的是在图像尺度还比较大的时候快速降低空间分辨率之后的卷积层逐渐过渡到 5×5 和 3×3。这种从粗到细的特征提取思路在后来的 VGG 中被进一步简化。对于新手来说AlexNet 最大的价值是帮助你建立“网络是一层一层拼接起来”的空间直觉。2.2 VGG更深的卷积堆叠VGG 的核心思想非常朴素与其设计各种尺寸复杂的卷积核不如全部使用 3×3 小卷积核通过堆叠更多层来增加感受野。两个 3×3 卷积叠加其有效感受野等于一个 5×5 卷积三个 3×3 卷积叠加则约等于一个 7×7 卷积。但小卷积核叠加的参数量更少非线性更强训练也更容易。VGG 通常有 VGG16 和 VGG19 两种常见配置分别对应 16 层和 19 层带权重层。它的缺点是全连接层参数非常多模型体积偏大但这并不妨碍它成为许多迁移学习任务中的经典主干。从工程角度看VGG 的模块化设计值得借鉴卷积层、ReLU、池化层被封装成多个 block重复堆叠。我们在写统一训练模板时也可以采用这种模块化思路把数据加载、模型构建、训练逻辑拆开让切换模型只改一个配置参数。2.3 ResNet残差学习打破退化当网络深度继续增加时会出现一个反直觉的现象训练集上的误差不降反升。这不是过拟合而是优化困难导致的“退化问题”。ResNet 的解决方案是引入残差连接也就是让某一层的输出变为 F(x) x其中 x 是该层的输入F(x) 是需要学习的残差映射。这样即使 F(x) 学不到什么网络也至少能保持恒等映射不会比浅层网络更差。ResNet 的残差块通常包含两个或三个卷积层配合 Batch Normalization 和 ReLU。这个设计让上百层甚至上千层的网络都能稳定训练。ResNet 是当今最常用的 CNN 主干之一ResNet18、ResNet34、ResNet50 在工业界和学术界都非常普遍。它的残差连接思想也被后来的很多模型吸收包括后面要讲的 ViT 中的跳跃连接。对于统一代码模板来说ResNet 适合作为默认的“可靠选择”因为它的训练稳定性最好。2.4 ViTTransformer进入视觉ViTVision Transformer不是 CNN但它和 CNN 在视觉任务中处于同一生态位因此在讨论模型演进时几乎绕不开它。ViT 把图像切成固定大小的 Patch例如 16×16 像素的方块每个 Patch 展平后通过一个线性映射得到向量再加入位置编码然后送入标准的 Transformer Encoder 中。由于 Transformer 的注意力机制能够建模全局依赖ViT 在大规模数据集上可以取得比 CNN 更好的效果但它非常依赖数据量和训练策略。在中小规模数据集上如果缺少预训练权重ViT 通常不如 ResNet 好训练。从工程角度看ViT 的输入形状和 CNN 不同。CNN 接受的是 (B, C, H, W) 的张量而 ViT 在 Patch Embedding 之后会把张量变换成 (B, N, D) 的 token 序列。因此统一训练模板需要针对不同模型做输入适配。最常见的做法是判断模型类型如果是以 ViT 为代表的 Transformer 模型就把图像从 (B, C, H, W) 展平成 patch 后变成序列如果是 CNN则保留四维张量直接过卷积层。下面给出的模板中我们会用一个 build_model 函数来屏蔽这种差异。2.5 模型选择建议模型核心特点适合场景训练难度AlexNet结构简单概念经典学习入门、小型数据集较低VGG统一小卷积核结构规整迁移学习、特征提取中等ResNet残差连接深度稳定通用视觉任务生产首选低ViT全局注意力依赖大数据大数据集、预训练微调较高如果是第一次跑通训练流程建议先使用 ResNet18因为它收敛快且对超参数不敏感。如果是为了复现论文中的对比实验则需要把 AlexNet、VGG、ResNet、ViT 都加到统一模板里通过命令行参数自由切换。3. 环境准备与项目结构3.1 环境依赖本文的代码基于 PyTorch它目前是学术界和工业界使用最广泛的深度学习框架之一。你需要准备一个可以运行 PyTorch 的 Python 环境。建议使用 Python 3.8 以上版本安装以下依赖pip install torch torchvision tqdm numpy tensorboard如果你的机器有 NVIDIA GPU还需要安装对应版本的 CUDA 和 cuDNN并确认 PyTorch 版本能够识别 GPU。可以用一段简单命令验证python -c import torch; print(torch.cuda.is_available()); print(torch.__version__)如果输出True说明环境已经支持 GPU 训练。如果输出False则后面代码会回退到 CPU 模式训练速度会慢很多但流程不受影响。版本号不需要刻意固定本文示例以常见环境为例重点是演示完整思路你可以根据自己项目实际调整版本。3.2 项目文件结构为了不让代码挤成一团建议按下面结构组织项目cnn_training_template/ ├── config.yaml ├── data_loader.py ├── models.py ├── train.py ├── predict.py └── deploy/ ├── export_model.py └── serve.pyconfig.yaml统一的配置入口包括数据集路径、模型类型、超参数等。data_loader.py封装数据预处理和 DataLoader。models.py模型工厂支持 AlexNet/VGG/ResNet/ViT。train.py训练、验证、保存脚本。predict.py本地加载模型做单张图片推理。deploy/export_model.py导出 TorchScript 或 ONNX。deploy/serve.py一个简单的 HTTP 推理服务。这样的结构把“数据、模型、训练、部署”四个环节分离后续替换数据集或者换模型时不需要改动全部代码。4. 统一代码模板设计4.1 模型工厂一键切换主干网络我们首先实现models.py。为了让模板尽可能简洁这里不手工实现 AlexNet 的完整结构而是使用torchvision.models中提供的官方实现。对于 ViT为了降低对 torchvision 版本的依赖我们实现一个简化版 ViT包含 Patch Embedding、位置编码和一个 Transformer Encoder 层。这个简化版可以用于教学和快速验证不去追求大规模 SOTA 效果。# 文件路径models.py import torch import torch.nn as nn from torchvision.models import alexnet, vgg16, resnet18, resnet50 class SimpleViT(nn.Module): 简化版 Vision Transformer仅用于说明 ViT 的前向流程。 def __init__(self, image_size224, patch_size16, num_classes10, dim192, depth6, heads12, mlp_dim384): super().__init__() assert image_size % patch_size 0 num_patches (image_size // patch_size) ** 2 patch_dim patch_size * patch_size * 3 self.patch_embed nn.Conv2d(3, dim, kernel_sizepatch_size, stridepatch_size) self.pos_embed nn.Parameter(torch.randn(1, num_patches 1, dim)) self.cls_token nn.Parameter(torch.randn(1, 1, dim)) encoder_layer nn.TransformerEncoderLayer( d_modeldim, nheadheads, dim_feedforwardmlp_dim, activationgelu, batch_firstTrue, ) self.transformer nn.TransformerEncoder(encoder_layer, num_layersdepth) self.norm nn.LayerNorm(dim) self.head nn.Linear(dim, num_classes) def forward(self, x): # x: (B, C, H, W) B x.size(0) x self.patch_embed(x) # (B, dim, H/patch, W/patch) x x.flatten(2).transpose(1, 2) # (B, num_patches, dim) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat((cls_tokens, x), dim1) # (B, num_patches1, dim) x x self.pos_embed x self.transformer(x) x self.norm(x[:, 0]) return self.head(x) def build_model(model_name, num_classes10, use_pretrainedFalse): 根据名称返回模型实例。 Args: model_name: 支持 alexnet / vgg16 / resnet18 / resnet50 / simple_vit num_classes: 分类任务类别数 use_pretrained: 是否加载 ImageNet 预训练权重ViT 示例不支持 if model_name alexnet: model alexnet(weightsNone) model.classifier[6] nn.Linear(4096, num_classes) elif model_name vgg16: model vgg16(weightsNone) model.classifier[6] nn.Linear(4096, num_classes) elif model_name resnet18: model resnet18(weightsNone) model.fc nn.Linear(model.fc.in_features, num_classes) elif model_name resnet50: model resnet50(weightsNone) model.fc nn.Linear(model.fc.in_features, num_classes) elif model_name simple_vit: model SimpleViT(num_classesnum_classes) else: raise ValueError(fUnknown model: {model_name}) if use_pretrained and model_name ! simple_vit: # 在实际项目中可以通过 weights预训练权重 加载这里留空 pass return model这里要注意几个细节。第一alexnet(weightsNone)表示随机初始化如果你希望加载 ImageNet 预训练权重可以改用alexnet(weightsDEFAULT)但需要确保 torchvision 版本足够新。第二替换最后一层分类头时classifier[6]对应 AlexNet 和 VGG 的最后一层全连接model.fc对应 ResNet 的全连接层这是由 torchvision 源码结构决定的。第三SimpleViT 的实现省略了很多训练技巧比如学习率预热、随机深度等在真正的生产项目中建议使用官方实现或成熟库。模型工厂的好处是让训练脚本与具体模型解耦。后续想加入新的网络结构只需要在build_model中增加一个分支并保证它的 forward 输入输出格式一致。4.2 数据加载与增强接下来实现data_loader.py。这里以 CIFAR-10 数据集为例因为 CIFAR-10 规模小、类别清晰适合跑通全流程。如果你有自己的图像分类数据集只需要把ImageFolder指向对应目录并调整均值方差即可。# 文件路径data_loader.py import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms def get_transforms(image_size224, trainTrue): 返回训练/验证的预处理流程。 if train: return transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.4, 0.4, 0.4), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) else: return transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) def get_dataloader(data_root./data, batch_size64, num_workers4, image_size224, use_cifar10True): train_transform get_transforms(image_size, trainTrue) val_transform get_transforms(image_size, trainFalse) if use_cifar10: train_dataset datasets.CIFAR10( rootdata_root, trainTrue, downloadTrue, transformtrain_transform) val_dataset datasets.CIFAR10( rootdata_root, trainFalse, downloadTrue, transformval_transform) else: train_dataset datasets.ImageFolder( rootf{data_root}/train, transformtrain_transform) val_dataset datasets.ImageFolder( rootf{data_root}/val, transformval_transform) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue, num_workersnum_workers, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_sizebatch_size, shuffleFalse, num_workersnum_workers, pin_memoryTrue) return train_loader, val_loader关于Normalize的参数很多同学会好奇为什么使用0.485, 0.456, 0.406这一组数字。这不是随便定的而是 ImageNet 数据集的 RGB 通道均值。如果你的任务不是 ImageNet 风格的自然图像比如医学影像、遥感图像就需要重新统计自己数据集的均值和标准差。训练时使用错误的 Normalize 参数会导致模型难以收敛这是非常常见的坑。4.3 训练循环核心训练脚本是train.py。它负责读取配置、构建模型、加载数据、执行多轮训练并在每一轮结束时验证和保存模型。为了保证代码可读性我们把训练步骤和验证步骤拆成两个函数。# 文件路径train.py import os import time import yaml import torch import torch.nn as nn from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm from data_loader import get_dataloader from models import build_model def train_one_epoch(model, loader, criterion, optimizer, device, epoch): model.train() total_loss 0.0 correct 0 total 0 pbar tqdm(loader, descfEpoch {epoch} [Train]) for images, labels in pbar: images images.to(device) labels labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total images.size(0) pbar.set_postfix(lossloss.item()) return total_loss / total, correct / total def validate(model, loader, criterion, device): model.eval() total_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in tqdm(loader, desc[Validate]): images images.to(device) labels labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total images.size(0) return total_loss / total, correct / total def main(): with open(config.yaml, r, encodingutf-8) as f: cfg yaml.safe_load(f) device torch.device(cuda if torch.cuda.is_available() else cpu) torch.manual_seed(cfg[seed]) train_loader, val_loader get_dataloader( data_rootcfg[data_root], batch_sizecfg[batch_size], num_workerscfg[num_workers], image_sizecfg[image_size], use_cifar10cfg.get(use_cifar10, True), ) model build_model( model_namecfg[model], num_classescfg[num_classes], use_pretrainedcfg.get(use_pretrained, False), ).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lrcfg[lr]) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxcfg[epochs]) writer SummaryWriter(log_dircfg[log_dir]) best_acc 0.0 os.makedirs(cfg[save_dir], exist_okTrue) for epoch in range(1, cfg[epochs] 1): start time.time() train_loss, train_acc train_one_epoch( model, train_loader, criterion, optimizer, device, epoch) val_loss, val_acc validate(model, val_loader, criterion, device) scheduler.step() writer.add_scalar(Loss/train, train_loss, epoch) writer.add_scalar(Loss/val, val_loss, epoch) writer.add_scalar(Acc/train, train_acc, epoch) writer.add_scalar(Acc/val, val_acc, epoch) print(fEpoch {epoch}: ftrain_loss{train_loss:.4f}, train_acc{train_acc:.4f}, fval_loss{val_loss:.4f}, val_acc{val_acc:.4f}, ftime{time.time()-start:.2f}s) if val_acc best_acc: best_acc val_acc checkpoint_path os.path.join(cfg[save_dir], f{cfg[model]}_best.pth) torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), val_acc: val_acc, }, checkpoint_path) print(fBest model saved to {checkpoint_path}) writer.close() if __name__ __main__: main()这段代码有几个工程细节值得说明。tqdm用来显示进度条SummaryWriter把训练指标写到 TensorBoard方便观察曲线。模型保存时没有只保存state_dict而是把优化器状态、当前 epoch、验证准确率一起打包成字典这样以后想恢复训练时可以直接加载。CosineAnnealingLR是常用的学习率调度策略它让学习率按照余弦曲线从初始值降到接近 0在 ImageNet 训练中被证明很有效。4.4 验证与模型保存你可能注意到上面的训练脚本中验证函数已经集成在validate里而模型保存由main中的if val_acc best_acc控制。为什么要保存最优模型而不是最后一轮模型因为深度学习训练过程中验证集准确率通常会在某些 epoch 达到峰值随后可能出现轻微过拟合。保存最优模型可以保证最终拿到的权重在验证集上表现最好。同时建议定期保存中间 checkpoint例如每 10 个 epoch 保存一次以便训练中断后可以恢复。篇幅所限上面代码只展示了最优保存逻辑你可以在循环内额外加一行torch.save(model.state_dict(), fcheckpoint_epoch_{epoch}.pth)4.5 训练入口与配置文件统一模板的入口是train.py但所有可调参数都放在config.yaml中。这样做的好处是不同实验之间只需要复制一份 YAML 文件修改模型名称和超参数而不需要改动代码。下面是一个示例配置# 文件路径config.yaml data_root: ./data save_dir: ./checkpoints log_dir: ./runs model: resnet18 # 可选: alexnet / vgg16 / resnet18 / resnet50 / simple_vit num_classes: 10 use_pretrained: false image_size: 224 batch_size: 64 num_workers: 4 epochs: 30 lr: 0.0003 seed: 42 use_cifar10: true如果你的数据不是 CIFAR-10而是自定义数据集将use_cifar10改为false并把数据按照data/train和data/val的子文件夹分类存放每个子文件夹名对应一个类别ImageFolder就能自动读取。这个配置思路在实际项目中非常常用。5. 完整实战MNIST/CIFAR10训练全流程5.1 配置说明上面的代码默认使用 CIFAR-10。CIFAR-10 包含 10 个类别每张图片尺寸为 32×32但我们的预处理流程会将图片 Resize 到 224×224。这在实际图片分类中很常见不同来源的图片分辨率不同统一 resize 到模型输入尺寸是必要的。如果你希望使用 MNIST 数据集只需要做三处调整第一MNIST 是单通道灰度图而 AlexNet/VGG/ResNet 的官方实现接受三通道输入需要把单通道图复制成三通道第二类别数改为 10 不变第三归一化参数应该使用 MNIST 的均值和标准差而不是 ImageNet 的。下面给出一个可选的 MNIST DataLoader 实现片段# 文件路径data_loader.py 中的可选 MNIST 分支 def get_mnist_dataloader(data_root./data, batch_size64, num_workers4): transform transforms.Compose([ transforms.Resize((224, 224)), transforms.Grayscale(num_output_channels3), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) train_dataset datasets.MNIST( rootdata_root, trainTrue, downloadTrue, transformtransform) val_dataset datasets.MNIST( rootdata_root, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue, num_workersnum_workers) val_loader DataLoader(val_dataset, batch_sizebatch_size, shuffleFalse, num_workersnum_workers) return train_loader, val_loaderGrayscale(num_output_channels3)会在保留灰度信息的同时将一张单通道图扩展成三通道这样就能直接输入到标准的 ResNet 中。MNIST 的均值 0.1307 和标准差 0.3081 是官方提供的常用数值。如果你不把 MNIST 图片 resize 到 224×224而是希望保持原始尺寸那么大多数 CNN 会无法直接工作因为 AlexNet 和 VGG 的最前面通常有下采样层对任意尺寸也可以接受但与预训练权重的输入约束不一致所以建议统一大小。5.2 运行训练在项目根目录执行python train.py如果一切正常你会在控制台看到类似下面的进度输出Epoch 1: train_loss1.9023, train_acc0.3298, val_loss1.5342, val_acc0.4561, time45.23s Epoch 2: train_loss1.3128, train_acc0.5350, val_loss1.0812, val_acc0.6050, time45.11s ... Epoch 30: train_loss0.0865, train_acc0.9741, val_loss0.2154, val_acc0.9412, time44.87s Best model saved to ./checkpoints/resnet18_best.pth注意上述数值只是一个示例实际结果会因设备、随机种子、数据增强和超参数不同而波动。如果你是第一次在 CPU 上运行训练时间会明显更长建议调小epochs和batch_size快速验证流程。5.3 结果分析与可视化训练结束后可以使用 TensorBoard 查看损失曲线和准确率曲线tensorboard --logdir./runs在浏览器中打开 TensorBoard 地址你可以看到训练集和验证集的 Loss 曲线。如果训练 Loss 持续下降但验证 Loss 上升说明发生了过拟合需要增加数据增强、增大 Dropout 或者提前停止。如果两个 Loss 都降不下去大概率是学习率设置不当或者模型结构不适合当前任务。6. 部署与推理从权重到服务6.1 模型导出训练好的模型最终要脱离训练环境运行。PyTorch 提供了两种常用的导出方式TorchScript 和 ONNX。TorchScript 可以让模型在一个不依赖 Python 原生代码的环境中被torch.jit加载ONNX 则可以转换到 ONNX Runtime、TensorRT 等推理引擎。这两种方式各有优劣我们这里以 ONNX 导出为例因为它更通用部署方不一定需要 PyTorch 环境。# 文件路径deploy/export_model.py import torch import yaml from models import build_model def main(): with open(config.yaml, r, encodingutf-8) as f: cfg yaml.safe_load(f) model build_model(cfg[model], num_classescfg[num_classes]) checkpoint torch.load(f{cfg[save_dir]}/{cfg[model]}_best.pth, map_locationcpu) model.load_state_dict(checkpoint[model_state_dict]) model.eval() dummy_input torch.randn(1, 3, cfg[image_size], cfg[image_size]) torch.onnx.export( model, dummy_input, f{cfg[model]}.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version12, ) print(ONNX model saved.) if __name__ __main__: main()导出 ONNX 时dynamic_axes允许 batch 维度是动态的这样部署端可以在单条样本和批量样本之间切换。opset_version需要根据你的 ONNX Runtime 版本来选择这里使用的是 12属于较常见的兼容选择。6.2 本地推理脚本导出 ONNX 后可以用 ONNX Runtime 或 PyTorch 原生的方式加载模型做推理。这里先提供一个纯 PyTorch 的推理脚本适合快速验证单张图片# 文件路径predict.py import torch from PIL import Image from torchvision import transforms from models import build_model def predict(image_path, model_name, checkpoint_path, num_classes10): device torch.device(cuda if torch.cuda.is_available() else cpu) model build_model(model_name, num_classesnum_classes).to(device) checkpoint torch.load(checkpoint_path, map_locationdevice) model.load_state_dict(checkpoint[model_state_dict]) model.eval() transform transforms.Compose([ transforms.Resize((224, 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) tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): output model(tensor) prob torch.softmax(output, dim1) pred prob.argmax(dim1).item() confidence prob[0, pred].item() return pred, confidence if __name__ __main__: result, conf predict( image_pathtest.jpg, model_nameresnet18, checkpoint_path./checkpoints/resnet18_best.pth, ) print(fPredicted class: {result}, confidence: {conf:.4f})这个脚本很简单但已经包含了一个生产级推理脚本必须的三个原则加载权重、设置 eval 模式、在torch.no_grad()下前向传播。如果你使用 ONNX Runtime可以加载导出的.onnx文件而不需要导入训练时的模型类。6.3 一个简单的HTTP推理服务最后我们部署成一个可供其他服务调用的 HTTP 接口。使用 Python 自带的http.server或者 Flask 都可以这里用标准库和 urllib 保持最小依赖。实际生产环境建议使用 FastAPI 或 Flask但核心逻辑是一样的。# 文件路径deploy/serve.py import io import json from http.server import HTTPServer, BaseHTTPRequestHandler from PIL import Image import torch from torchvision import transforms # 在真实项目中模型应该初始化为全局变量避免每次请求重复加载 model None device cpu def load_model(): global model model torch.jit.load(resnet18_traced.pt) model.eval() model.to(device) def prepare_image(image_bytes): transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) image Image.open(io.BytesIO(image_bytes)).convert(RGB) return transform(image).unsqueeze(0).to(device) class Handler(BaseHTTPRequestHandler): def do_POST(self): content_length int(self.headers[Content-Length]) image_bytes self.rfile.read(content_length) tensor prepare_image(image_bytes) with torch.no_grad(): outputs model(tensor) prob torch.softmax(outputs, dim1) pred prob.argmax(dim1).item() conf prob[0, pred].item() result json.dumps({class: pred, confidence: conf}) self.send_response(200) self.send_header(Content-Type, application/json) self.end_headers() self.wfile.write(result.encode(utf-8)) if __name__ __main__: load_model() server HTTPServer((0.0.0.0, 8000), Handler) print(Server started on port 8000) server.serve_forever()需要说明的是serve.py中使用torch.jit.load加载 TorchScript 模型而不是直接加载state_dict。这是因为服务端通常不需要关心模型的具体实现只负责前向计算。如果你想在服务端保留 python 模型定义也可以像predict.py那样加载 checkpoint但每次启动服务都会重新构建模型结构长期维护成本更高。在实际项目中TorchScript 或者 ONNX 是更合理的部署形态。部署服务后你可以用 curl 发送一张图片测试curl -X POST -H Content-Type: application/octet-stream --data-binary test.jpg http://127.0.0.1:8000/如果返回{class: 3, confidence: 0.9871}说明服务链路已经跑通。后续可以把这个服务放在内网或者接入 API 网关供业务系统调用。7. 常见问题与排查思路7.1 训练不收敛问题现象常见原因解决思路Loss 一直不下降学习率太大或太小尝试从 3e-4 开始按数量级调整验证集准确率低数据预处理与训练不一致检查 Normalize 参数和 Resize 逻辑出现 NaN Loss学习率过大或权重初始化不当降低学习率检查是否除以 0训练 Loss 下降但验证 Loss 上升过拟合增加数据增强、Dropout使用早停如果你刚开始训练可以先跑 5 个 epoch 观察趋势。如果 Loss 在初期没有下降优先怀疑学习率。学习率太大优化过程会在损失曲面震荡学习率太小则收敛极慢。一个有效的策略是使用学习率预热和余弦退火这也是我们在训练脚本中使用CosineAnnealingLR的原因。7.2 显存不足报错信息通常是CUDA out of memory。解决思路依次是减小batch_size降低图片输入分辨率使用gradient accumulation在多个小 batch 上累积梯度后再更新参数换用更小的模型例如从 ResNet50 改成 ResNet18。注意减小 batch size 后可能需要同步调整学习率因为 batch size 变小意味着每个 step 的梯度噪声变大为了让训练稳定可以适当降低学习率。# 梯度累积示例每 4 个 batch 更新一次 accumulation_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): outputs model(images.to(device)) loss criterion(outputs, labels.to(device)) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()7.3 模型加载报错如果加载 checkpoint 时报错size mismatch for fc.weight说明你保存模型时的类别数和当前模型类别数不一致。这种问题多发生在你换了一个数据集训练却用旧类别的预训练头继续加载。解决方法是确保构建模型时传入的num_classes等于训练时的类别数。如果是加载预训练模型做迁移学习你可能会覆盖最后一层那么旧 checkpoint 中最后一层的权重不匹配是正常的可以忽略或只加载前缀匹配的参数。7.4 部署后推理速度慢ONNX 模型在 CPU 上运行时如果速度和 PyTorch 差不多可以考虑量化或更换推理引擎。也可以检查是否取消了梯度部署推理时一定要在torch.no_grad()下运行。如果使用 GPU 部署还需要保证输入张量被放在 GPU 上。对于 ViT 这类 Transformer 模型它的计算量和输入分辨率是平方关系如果输入图片从 224 改为 448速度会下降数倍。在不改变模型结构的前提下可以通过 ONNX Runtime 的优化、动态形状、算子融合等手段提高速度。8. 最佳实践与工程建议8.1 数据与训练分离一个成熟的训练项目应该把数据读取、预处理、模型定义和训练逻辑完全分离。一旦某个环节变更比如从 CIFAR-10 切换到自己的业务数据集你不需要重写训练脚本只需要修改数据和配置。上面的模板已经体现了这一点但建议你在业务代码中做得更彻底数据增强策略单独配置模型结构用注册机制管理损失函数和优化器也从配置读取。8.2 配置管理使用 YAML 或 JSON 管理配置比硬编码在代码里更安全。一个实验对应一个配置目录下面可以包含config.yaml、model_architecture.py、训练日志和 checkpoint。复现实验时直接看配置目录就能知道当时的超参数和数据路径。特别注意任何涉及生产环境训练的操作都要遵循最小权限原则不要在共享服务器上随意覆盖别人的配置。8.3 日志与监控训练过程要保留三类信息训练指标、运行日志、模型元信息。训练指标包括 loss、accuracy、learning rate运行日志包括启动时间、GPU 温度、显存占用模型元信息包括模型版本、训练数据版本、日期。这些信息可以帮助你在模型上线后快速定位问题。如果是在企业内网部署请务必遵守公司的安全规范不要将敏感数据轻易写入公开日志。8.4 安全与权限模型部署到生产环境后接口可能被非预期调用。服务端应该增加身份认证和访问频率限制避免被刷接口。模型文件本身可能包含数据集信息如果数据集涉及用户隐私需要对模型做加密和访问控制。涉及数据库或系统权限变更时必须经过合法授权并在测试环境验证。无论训练还是部署都不要将账号密码、密钥等明文写入代码或配置文件。8.5 性能优化当模型需要服务大量并发请求时单线程的 Python HTTP 服务显然不够。优化的方向包括使用异步框架如 FastAPI使用 ONNX Runtime 或 TensorRT 替换 PyTorch 推理开多个 worker 进程将模型放在 GPU 显存中避免重复加载对输入图片做尺寸压缩和格式统一。性能优化要结合自己的瓶颈来做不要盲目堆机器先用 profiling 工具定位耗时阶段。9. 总结与学习路线到这里你已经跟着一份统一代码模板走完了从 CNN 入门到模型部署的全流程。我们梳理了 AlexNet、VGG、ResNet 和 ViT 的设计思路实现了一个可以一键切换主干的模型工厂完成了数据加载、训练循环、验证保存和 ONNX 导出最后部署了一个最简单的 HTTP 推理服务。无论你以后是在做毕业设计、算法竞赛还是企业级视觉项目这套流程都可以作为起步模板。下一步的学习路线取决于你的目标。如果你想深入理解 CNN 的工作原理建议手动实现一遍 ResNet 残差块的 forward 和 backward观察每个张量形状的变化如果你想在业务中落地建议尝试修改config.yaml中的模型名称跑通 AlexNet、VGG、ResNet、ViT 四种模型的对比实验体会不同结构在收敛速度、准确率和模型体积上的差异如果你想进一步了解大规模部署可以从 ONNX Runtime 的服务化框架开始学习如何用异步接口接收图片并进行批量推理。技术路上没有捷径但一个好的代码模板能让你在重复劳动中节约大量时间。把上面的代码跑一遍然后开始改造它你会发现自己的成长比想象中更快。
网站建设高端定制企业官网