☰
从AlexNet到ViT:PyTorch统一训练模板与模型部署实践
2026/9/26 11:30:45 网站建设 项目流程

1. 背景与核心概念

如果现在要评选过去十年影响最深远的深度学习模型,卷积神经网络(Convolutional Neural Network,CNN)一定是最有竞争力的候选之一。从 2012 年 AlexNet 在 ImageNet 大赛上一举夺冠开始,CNN 逐步成为图像分类、目标检测、语义分割、人脸识别等任务的默认底座。即使后来 Transformer 大行其道,CNN 在很多场景下仍然是训练稳定、推理高效、部署方便的首选。本文希望围绕一条主线展开:从 CNN 的基本结构出发,梳理 AlexNet、VGG、ResNet 和 ViT 这几类经典模型之间的关系,再给出一份可以一次训练、自由切换主干的 PyTorch 代码模板,最后聊一聊如何把训练好的模型导出、部署成本地推理服务。内容适合刚接触深度学习的同学,也适合想快速跑通实验结果的开发者。

先来补充一点必要的背景。CNN 本质上是一个“局部连接、权值共享”的神经网络,它不像全连接网络那样把上一层的每个神经元都与下一层的每个神经元连接,而是用一组小尺寸的卷积核在输入特征图上滑动,从而提取局部特征。这种设计有两个直接的好处:其一,参数数量大幅下降,模型更容易训练;其二,卷积核天然具有平移等变性,也就是说物体在图片中移动几个像素,提取到的特征依然类似。这两个特性使 CNN 在图像任务上远优于传统的全连接网络。为了控制特征图尺寸并提高非线性表达能力,CNN 还会配合池化层和激活函数使用,最终通过全连接层或者全局平均池化输出分类结果。

学习 CNN 有一个完整的训练闭环:数据集准备、数据预处理与增强、网络前向传播、损失计算、反向传播、参数更新、周期性验证、模型保存。很多人一开始只关注模型结构,却忽略了数据处理和训练策略,导致同样的代码在别人的电脑上能收敛,在自己这里就发散。下面我们会先把模型演进讲清楚,然后直接用一个统一代码模板把整个闭环串起来。在实际项目中,这套流程可以帮助你快速验证一个新的网络结构是否适合当前任务,也可以作为工程化改造的起点。

2. 经典模型演进:AlexNet/VGG/ResNet/ViT

2.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 ViT:Transformer进入视觉

ViT(Vision 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.py
  • config.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_size=224, patch_size=16, num_classes=10, dim=192, depth=6, heads=12, mlp_dim=384): 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_size=patch_size, stride=patch_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_model=dim, nhead=heads, dim_feedforward=mlp_dim, activation="gelu", batch_first=True, ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=depth) 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), dim=1) # (B, num_patches+1, 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_classes=10, use_pretrained=False): """根据名称返回模型实例。 Args: model_name: 支持 alexnet / vgg16 / resnet18 / resnet50 / simple_vit num_classes: 分类任务类别数 use_pretrained: 是否加载 ImageNet 预训练权重(ViT 示例不支持) """ if model_name == "alexnet": model = alexnet(weights=None) model.classifier[6] = nn.Linear(4096, num_classes) elif model_name == "vgg16": model = vgg16(weights=None) model.classifier[6] = nn.Linear(4096, num_classes) elif model_name == "resnet18": model = resnet18(weights=None) model.fc = nn.Linear(model.fc.in_features, num_classes) elif model_name == "resnet50": model = resnet50(weights=None) model.fc = nn.Linear(model.fc.in_features, num_classes) elif model_name == "simple_vit": model = SimpleViT(num_classes=num_classes) else: raise ValueError(f"Unknown model: {model_name}") if use_pretrained and model_name != "simple_vit": # 在实际项目中,可以通过 weights=预训练权重 加载,这里留空 pass return model

这里要注意几个细节。第一,alexnet(weights=None)表示随机初始化,如果你希望加载 ImageNet 预训练权重,可以改用alexnet(weights="DEFAULT"),但需要确保 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_size=224, train=True): """返回训练/验证的预处理流程。""" 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_size=64, num_workers=4, image_size=224, use_cifar10=True): train_transform = get_transforms(image_size, train=True) val_transform = get_transforms(image_size, train=False) if use_cifar10: train_dataset = datasets.CIFAR10( root=data_root, train=True, download=True, transform=train_transform) val_dataset = datasets.CIFAR10( root=data_root, train=False, download=True, transform=val_transform) else: train_dataset = datasets.ImageFolder( root=f"{data_root}/train", transform=train_transform) val_dataset = datasets.ImageFolder( root=f"{data_root}/val", transform=val_transform) train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True) 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, desc=f"Epoch {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(loss=loss.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", encoding="utf-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_root=cfg["data_root"], batch_size=cfg["batch_size"], num_workers=cfg["num_workers"], image_size=cfg["image_size"], use_cifar10=cfg.get("use_cifar10", True), ) model = build_model( model_name=cfg["model"], num_classes=cfg["num_classes"], use_pretrained=cfg.get("use_pretrained", False), ).to(device) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["lr"]) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=cfg["epochs"]) writer = SummaryWriter(log_dir=cfg["log_dir"]) best_acc = 0.0 os.makedirs(cfg["save_dir"], exist_ok=True) 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(f"Epoch {epoch}: " f"train_loss={train_loss:.4f}, train_acc={train_acc:.4f}, " f"val_loss={val_loss:.4f}, val_acc={val_acc:.4f}, " f"time={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(f"Best 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(), f"checkpoint_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_size=64, num_workers=4): transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.Grayscale(num_output_channels=3), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) train_dataset = datasets.MNIST( root=data_root, train=True, download=True, transform=transform) val_dataset = datasets.MNIST( root=data_root, train=False, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers) val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers) return train_loader, val_loader

Grayscale(num_output_channels=3)会在保留灰度信息的同时将一张单通道图扩展成三通道,这样就能直接输入到标准的 ResNet 中。MNIST 的均值 0.1307 和标准差 0.3081 是官方提供的常用数值。如果你不把 MNIST 图片 resize 到 224×224,而是希望保持原始尺寸,那么大多数 CNN 会无法直接工作,因为 AlexNet 和 VGG 的最前面通常有下采样层,对任意尺寸也可以接受,但与预训练权重的输入约束不一致,所以建议统一大小。

5.2 运行训练

在项目根目录执行:

python train.py

如果一切正常,你会在控制台看到类似下面的进度输出:

Epoch 1: train_loss=1.9023, train_acc=0.3298, val_loss=1.5342, val_acc=0.4561, time=45.23s Epoch 2: train_loss=1.3128, train_acc=0.5350, val_loss=1.0812, val_acc=0.6050, time=45.11s ... Epoch 30: train_loss=0.0865, train_acc=0.9741, val_loss=0.2154, val_acc=0.9412, time=44.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", encoding="utf-8") as f: cfg = yaml.safe_load(f) model = build_model(cfg["model"], num_classes=cfg["num_classes"]) checkpoint = torch.load(f"{cfg['save_dir']}/{cfg['model']}_best.pth", map_location="cpu") 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_version=12, ) 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_classes=10): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = build_model(model_name, num_classes=num_classes).to(device) checkpoint = torch.load(checkpoint_path, map_location=device) 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, dim=1) pred = prob.argmax(dim=1).item() confidence = prob[0, pred].item() return pred, confidence if __name__ == "__main__": result, conf = predict( image_path="test.jpg", model_name="resnet18", checkpoint_path="./checkpoints/resnet18_best.pth", ) print(f"Predicted 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, dim=1) pred = prob.argmax(dim=1).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 的服务化框架开始,学习如何用异步接口接收图片并进行批量推理。技术路上没有捷径,但一个好的代码模板能让你在重复劳动中节约大量时间。把上面的代码跑一遍,然后开始改造它,你会发现自己的成长比想象中更快。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询