ARTICLE DETAIL

资讯详情

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

PyTorch统一训练模板:CNN与ViT图像分类模型工程化实战

PyTorch统一训练模板:CNN与ViT图像分类模型工程化实战 学深度学习图像分类最容易遇到的一个坑不是模型看不懂而是每个模型对应一套独立的训练代码。上周还在用 torchvision 读 AlexNet这周导师让换成 ResNet网上找到的代码数据预处理是一套写法训练循环又是另一种封装能跑通但换个数据集就报错。你以为时间花在研究模型差异上实际上全消耗在对齐数据、改训练逻辑、调试维度不匹配这些重复劳动里。更扎心的是模型定义本身通常只有几十行真正的工程工作——数据管线、训练循环、验证与保存、部署导出——才是占据 80% 工作量的部分。把这部分做成一份统一模板让 AlexNet、VGG、ResNet、ViT 共用同一条数据管道和同一个训练入口才是从入门走向工程化的关键一步。这篇文章会给你一份可直接运行的代码模板用模型注册表统一创建四个模型用同一套数据加载、训练、验证、保存逻辑最后给出 ONNX 部署导出的通用路径。读完你可以把它当脚手架替换成自己的数据集和网络结构不必每个新模型都重新造轮子。文中代码基于 PyTorch 编写整体思路与框架无关TensorFlow 用户同样可以借鉴。1. 先看清问题为什么训练代码比模型定义更容易拖垮你很多初学者误以为“读懂网络结构就能写训练代码”。实际上一个分类任务的完整训练流程由以下部分组成数据读取与预处理resize、归一化、增强策略模型实例化与权重初始化损失函数、优化器、学习率调度器配置训练循环前向、反向、梯度更新、日志输出验证逻辑关闭梯度、计算准确率模型保存与加载部署导出ONNX、TensorRT 等模型的网络结构只是其中一环而且是最容易复用的一环。因为 torchvision 等库已经帮你实现了主流分类网络的完整定义你需要改的往往只是最后的分类头。真正的差异集中在数据管线、训练策略和验证方式上——这些恰恰是网上代码质量最参差不齐的地方。如果再叠加另一个现实因素AlexNet 输入是 227×227VGG 是 224×224ResNet 和 ViT 也有各自的预处理习惯。每个模型对应一种数据增强组合手写多个独立脚本必然是灾难。统一模板的价值就在这里模型可以换数据管线只有一份训练验证循环只有一份命令行参数切换即可。换模型从“改代码”变成“改参数”。2. 四个模型的核心原理与演进脉络要真正用好模板得先理解这四个模型在解决什么问题以及它们的特征表达方式有什么不同。这里不打算展开到论文逐行推导而是抓住每个模型最关键的创新点。2.1 CNN 的基本组件卷积神经网络CNN的核心思想是用卷积核在图像上滑动提取局部特征。卷积层负责特征提取池化层负责下采样压缩空间尺寸全连接层负责在最后做分类决策。相比全连接网络直接铺平成向量CNN 保留了图像的二维空间结构参数量也大幅减少。2.2 AlexNet深度学习登场的开山之作AlexNet 在 2012 年 ImageNet 竞赛以巨大优势夺冠让深度学习从此成为视觉领域的主流路线。它的核心贡献是使用 ReLU 激活函数缓解梯度消失使用 Dropout 抑制过拟合并通过两块 GPU 并行训练来支撑更大的网络容量。结构上它由 5 个卷积层和 3 个全连接层组成输入尺寸为 227×227。从今天的视角看AlexNet 的结构不算复杂但它证明了“更深更大的网络 GPU 训练 正则化技巧”这套组合的威力。学习它的意义在于理解现代 CNN 的骨架雏形。2.3 VGG用 3×3 卷积把“深”做到极致VGG 的核心判断很朴素用多个 3×3 小卷积核堆叠替代一个大的卷积核。两个 3×3 卷积的感受野等于一个 5×5 卷积但参数量更少且中间多了一次非线性变换表达能力更强。VGG16 和 VGG19 在相当长一段时间里是特征提取的默认选择。VGG 的问题也很明显全连接层参数量巨大一个 VGG16 的模型文件超过 500MB训练和部署成本都不低。模板中保留它是为了让你对比“结构简单直接”和“参数冗余严重”这两个工程现实。2.4 ResNet残差连接解决退化问题按理说网络越深效果越好但实验发现当网络深到一定程度后训练误差反而上升这就是“退化问题”。ResNet 给出的解法是残差连接让每个块学习输入与输出之间的残差即输出 输入 卷积变换结果。这个设计的物理意义很直接即使新增的层没有学到有效特征它至少可以学成恒等映射让深层网络的性能不低于浅层网络。同时跳跃连接让梯度可以更顺畅地回传到浅层缓解了梯度消失。ResNet 让网络可以安全地堆到 50 层、101 层甚至更深是 CNN 发展史上的关键转折。如今你看到的 ResNet18、ResNet50 几乎成了视觉任务的默认底座。2.5 ViT从 CNN 换成自注意力的路线切换Vision TransformerViT在 2020 年提出把 NLP 领域的 Transformer 结构直接搬到图像上。做法是把图像切成固定大小的 patch例如 16×16每个 patch 展平后映射为 token 向量再加上位置编码表示空间位置最后送入 Transformer Encoder 用自注意力建模全局关系。ViT 的特别之处在于它一开始就没有依赖卷积的局部归纳偏置而是完全靠数据学习哪些区域需要交互。代价是需要大量训练数据在 ImageNet 这种规模的数据集上ViT 才能发挥出超越 ResNet 的能力。如果数据量不够ViT 的效果往往不如同规模的 CNN。2.6 特征粒度视角粗粒度与细粒度的不同处理方式理解这四个模型还可以从“特征粒度”的角度切入。CNN 的特征是逐层抽象的浅层卷积感受野小学到的是边缘、纹理这类细粒度特征深层卷积感受野大学到的是物体部件、语义类别这类粗粒度特征。ResNet 的残差连接还有一个隐性作用——保留浅层的细粒度细节避免深层抽象过程中信息过度丢失。ViT 的粒度逻辑完全不同。patch embedding 直接把图像切成 16×16 的块每个 token 天然就是“粗粒度的局部区域”自注意力再让这些 token 全局交互。它缺少了 CNN 那种逐层从细到粗的渐变过程。因此在实际项目中检测任务常借助特征金字塔结构例如 FPN把高层粗粒度语义特征和低层细粒度位置特征融合使用Transformer 模型也会额外设计多尺度模块来补足细粒度信息。位置编码则是 ViT 里补偿空间信息的关键手段——自注意力本身是不感知顺序的没有位置编码模型无法知道 patch 来自图像的哪个区域。表格总结四个模型的核心差异模型核心机制特征粒度特点主要瓶颈AlexNet深度卷积网络 ReLU Dropout浅层细粒度深层粗粒度参数多结构笨重VGG3×3 小卷积堆叠逐层抽象简单直接全连接层参数冗余ResNet残差连接缓解退化保留细粒度细节支持深层堆叠设计残差块需要经验ViTpatch 切分 自注意力 位置编码全局交互粗粒度 token 起步数据量要求高训练慢3. 环境准备与实验约定开始写代码前先统一环境。本模板使用 PyTorch因为它同时覆盖了模型库、数据加载和训练生态是目前学习成本最低的深度学习框架。3.1 版本与依赖推荐环境如下Python 3.8 及以上PyTorch 2.x1.13 也可运行新版 API 兼容性更好torchvision需与 PyTorch 版本匹配用于加载 CIFAR-10 和预置模型CUDA GPU 可选没有 GPU 也能运行只是训练时间会明显变长安装命令pip install torch torchvision如果需要 GPU 版本请前往 PyTorch 官网根据你的 CUDA 版本生成安装命令。这里不写死具体命令是因为不同机器的 CUDA 环境差异较大写错反而会装不上。3.2 数据集与目录组织模板使用 CIFAR-10 作为演示数据集原因是它足够小下载快交叉验证方便而且能直接验证数据管线是否正确。运行代码时torchvision 会自动下载到本地无需手工准备。建议的项目目录结构如下image-classification-template/ ├── models.py # 模型注册表与模型工厂 ├── train.py # 统一训练脚本 ├── export_onnx.py # 部署导出脚本 ├── data/ # 数据集目录自动创建 └── checkpoints/ # 模型权重保存目录自动创建把代码拆成独立文件而不是把所有内容堆在一个脚本里是为了后续扩展更舒服。模型、训练、导出各自的修改不会互相影响。4. 统一训练模板的设计与实现这一节是文章的核心。先把模板的整体设计讲清楚再给出完整代码。设计上采用“模型注册表 统一训练循环”的模式这是中小型项目里性价比最高的架构。4.1 模型注册表用一份代码管理多个模型模型注册表的思路很简单用一个字典保存模型名称和对应构建函数的映射关系新增模型时只需要注册一个新函数主程序不需要修改。# 文件路径models.py import torch.nn as nn import torchvision.models as models MODEL_REGISTRY {} def register_model(name): def decorator(builder): MODEL_REGISTRY[name] builder return builder return decorator register_model(alexnet) def build_alexnet(num_classes): model models.alexnet(weightsNone) # AlexNet 最后一个分类层是 classifier[6] model.classifier[6] nn.Linear(4096, num_classes) return model register_model(vgg16) def build_vgg16(num_classes): model models.vgg16(weightsNone) model.classifier[6] nn.Linear(4096, num_classes) return model register_model(resnet18) def build_resnet18(num_classes): model models.resnet18(weightsNone) # ResNet 的分类头属性名为 fc model.fc nn.Linear(model.fc.in_features, num_classes) return model register_model(vit) def build_vit(num_classes): model models.vit_b_16(weightsNone) # 不同 torchvision 版本中 ViT 分类头属性名可能有差异 # 不确定时先 print(model) 查看结构再决定替换哪一层 model.heads.head nn.Linear(model.heads.head.in_features, num_classes) return model def build_model(name, num_classes10): if name not in MODEL_REGISTRY: raise ValueError( fUnsupported model: {name}, available: {list(MODEL_REGISTRY.keys())} ) return MODEL_REGISTRY[name](num_classes)这里需要注意一个现实问题不同 torchvision 版本里模型的分类头属性名并不完全相同。ResNet 是fcAlexNet 和 VGG 是classifier[6]ViT 在多数新版本里是heads.head。如果你使用的版本不同运行后会报维度不匹配错误此时先打印模型结构再定位需要替换的分类层即可。这就是“注册表模式”的优势——问题被隔离在构建函数内部不会污染训练代码。4.2 数据加载训练集与验证集的 transform 策略训练集和验证集的数据处理策略必须不同。训练集需要随机裁剪、随机翻转等数据增强验证集只需要统一 resize 和中心裁剪保证评测结果稳定可复现。两者共享同一套 Normalize 参数这是 ImageNet 预训练的标准均值标准差。# 文件路径train.py第一部分数据加载 import argparse import os import time import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from models import build_model def get_transform(resize_size224, trainTrue): if train: return transforms.Compose([ transforms.RandomResizedCrop(resize_size), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize( mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225], ), ]) return transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(resize_size), transforms.ToTensor(), transforms.Normalize( mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225], ), ]) def load_data(batch_size64, data_dir./data): train_ds datasets.CIFAR10( rootdata_dir, trainTrue, downloadTrue, transformget_transform(trainTrue), ) val_ds datasets.CIFAR10( rootdata_dir, trainFalse, downloadTrue, transformget_transform(trainFalse), ) train_loader DataLoader( train_ds, batch_sizebatch_size, shuffleTrue, num_workers4, pin_memoryTrue, ) val_loader DataLoader( val_ds, batch_sizebatch_size, shuffleFalse, num_workers4, pin_memoryTrue, ) return train_loader, val_loaderCIFAR-10 原始图片尺寸是 32×32而这里统一 resize 到 224×224是为了让四个模型的输入尺寸保持一致。代价是训练速度变慢但可以避免因输入尺寸不同带来的各种维度错误。如果你只想快速验证代码可以把resize_size改成 64 或 128模板依然能运行只是精度会有变化。4.3 训练循环与验证逻辑训练循环是所有模型共用的包含前向传播、计算损失、反向传播、参数更新四个步骤。验证时关闭梯度只统计损失和准确率防止 BatchNorm 等层的行为被验证过程影响。# 文件路径train.py第二部分训练与验证 def train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in loader: images, labels images.to(device), 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 torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) return total_loss / total, correct / total torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() total_loss, correct, total 0.0, 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) return total_loss / total, correct / total小技巧torch.no_grad()装饰器告诉 PyTorch 不需要在验证阶段构建计算图可以显著减少显存占用和验证时间。新手最容易犯的错误是验证时忘记model.eval()导致 Dropout 和 BatchNorm 的行为与训练状态不一致验证精度出现异常波动。4.4 命令行入口主函数设计成命令行工具通过参数控制模型类型、训练轮数、学习率等。这样切换模型不需要改代码也方便做超参数对比实验。# 文件路径train.py第三部分主程序 def main(): parser argparse.ArgumentParser() parser.add_argument( --model, typestr, defaultresnet18, choices[alexnet, vgg16, resnet18, vit], ) parser.add_argument(--epochs, typeint, default20) parser.add_argument(--batch-size, typeint, default64) parser.add_argument(--lr, typefloat, default1e-3) parser.add_argument(--data-dir, typestr, default./data) parser.add_argument(--ckpt-dir, typestr, default./checkpoints) parser.add_argument( --device, typestr, defaultcuda if torch.cuda.is_available() else cpu, ) args parser.parse_args() os.makedirs(args.ckpt_dir, exist_okTrue) device torch.device(args.device) train_loader, val_loader load_data(args.batch_size, args.data_dir) model build_model(args.model, num_classes10).to(device) criterion nn.CrossEntropyLoss() optimizer optim.AdamW( model.parameters(), lrargs.lr, weight_decay5e-4, ) scheduler optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxargs.epochs, ) best_acc 0.0 for epoch in range(1, args.epochs 1): start time.time() train_loss, train_acc train_one_epoch( model, train_loader, criterion, optimizer, device, ) val_loss, val_acc evaluate( model, val_loader, criterion, device, ) scheduler.step() elapsed time.time() - start print( fEpoch [{epoch:03d}/{args.epochs}] ftrain_loss{train_loss:.4f} train_acc{train_acc * 100:.2f}% fval_loss{val_loss:.4f} val_acc{val_acc * 100:.2f}% ftime{elapsed:.1f}s ) if val_acc best_acc: best_acc val_acc torch.save( model.state_dict(), os.path.join(args.ckpt_dir, fbest_{args.model}.pth), ) print(fBest val_acc: {best_acc * 100:.2f}%) if __name__ __main__: main()这里有几个细节值得解释。优化器选择 AdamW 而不是 Adam因为 AdamW 将权重衰减从动量项中解耦是目前 Transformer 类模型的默认选择同时也能兼容 CNN不需要为不同模型更换优化器。学习率调度使用 CosineAnnealingLR让学习率按余弦曲线从初始值下降至接近 0是一种简单且稳健的策略。保存模型时只保存state_dict而非整个模型对象这样文件更小且加载时不依赖模型类的定义位置。5. 将四个模型接入统一模板模板写好后四个模型的接入过程其实已经在模型注册表中完成了。这里再逐个拆解方便你理解每个模型需要修改哪些部分以后换自己的网络时也能按同样的思路接入。5.1 AlexNet 的接入AlexNet 在 torchvision 中由features卷积段和classifier全连接段组成。我们只替换最后一层classifier[6]把 4096 维映射到分类数。如果你新增的模型是自研网络直接在models.py里新增一个register_model(mynet)装饰的构建函数即可训练主程序完全不用动。5.2 VGG 的接入VGG 和 AlexNet 的分类头结构类似替换方式相同。值得一提的是 torchvision 的vgg16默认使用 BatchNorm 版本真正的名称是vgg16_bn而vgg16是不带 BatchNorm 的原始版本。建议你对比跑一下两个版本观察 BatchNorm 对收敛速度和最终精度的影响这是一个很有价值的小实验。5.3 ResNet 的接入ResNet 的分类头是fc单层线性层。替换时使用model.fc.in_features获取输入维度而不是硬编码 512 或 2048。不同深度的 ResNet 最后的特征维度不同ResNet18 是 512ResNet50 是 2048。用in_features自适应获取代码就不会因为换了模型深度而报错。5.4 ViT 的接入与 Patch Embedding 原理ViT 的分类头在 torchvision 新版本中位于model.heads.head。如果你运行时报错找不到该属性大概率是版本差异先打印模型结构再替换。理解 ViT 的关键在于 Patch Embedding。下面这段代码展示了 ViT 如何把图像切成 patch 并生成 token 序列这是 ViT 和 CNN 最本质的区别# 简化版 PatchEmbedding用于理解 ViT 的输入处理 import torch import torch.nn as nn class PatchEmbedding(nn.Module): def __init__(self, in_channels3, embed_dim768, patch_size16, num_patches196): super().__init__() # 一个 stridepatch_size 的卷积等价于切 patch 并映射成向量 self.proj nn.Conv2d( in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size, ) # 分类 token用来汇总全局信息 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) # 位置编码补偿自注意力不感知顺序的问题 self.pos_embed nn.Parameter( torch.zeros(1, num_patches 1, embed_dim) ) def forward(self, x): # 输入 x: [B, 3, 224, 224] x self.proj(x) # 输出 [B, embed_dim, 14, 14] x x.flatten(2).transpose(1, 2) # 输出 [B, 196, embed_dim] # 拼接 cls_token序列变为 197 个 token x torch.cat([self.cls_token.expand(x.size(0), -1, -1), x], dim1) # 加入位置编码 x x self.pos_embed return x224×224 的图像patch_size16切分后是 14×14196 个 patch。每个 patch 展平映射成 embed_dim例如 768维向量组成长度为 196 的 token 序列再加一个分类 token总共 197 个 token 输入 Transformer Encoder。位置编码是这里最容易被忽略的部分——自注意力本身不包含顺序信息没有位置编码模型无法知道 token 之间的空间位置关系。5.5 进阶使用预训练权重做迁移学习上面的模板默认weightsNone也就是随机初始化。这种做法的优点是代码简单、不受下载预训练权重的时间影响缺点是在小数据集上收敛慢、精度低。实际工程项目中更推荐加载 ImageNet 预训练权重再做迁移学习。只需把构建函数里的weightsNone改成weightsmodels.ResNet18_Weights.IMAGENET1K_V1即可其他代码完全不动。迁移学习通常只需要原来十分之一到五分之一的学习率数据增强可以更加激进。6. 运行、验证与效果判断代码写完后接下来要解决的问题是怎么运行、怎么判断训练是否正常、效果差应该从哪里开始调。6.1 训练命令切换模型只需要修改--model参数# 训练 ResNet18 python train.py --model resnet18 --epochs 20 --batch-size 64 # 训练 AlexNet python train.py --model alexnet --epochs 20 --batch-size 64 # 训练 ViTCPU 上会很慢建议有 GPU 再运行 python train.py --model vit --epochs 20 --batch-size 32 --lr 5e-4第一次运行会自动下载 CIFAR-10 数据集时间取决于网络。之后再次运行会直接读取本地文件不会重复下载。6.2 预期输出与日志解读正常训练时终端会逐轮输出类似下面的日志Epoch [001/020] train_loss1.8342 train_acc37.25% val_loss1.7214 val_acc46.11% time12.3s Epoch [002/020] train_loss1.4217 train_acc49.83% val_loss1.5862 val_acc54.34% time12.1s Epoch [003/020] train_loss1.2184 train_acc57.02% val_loss1.4517 val_acc58.79% time12.2s不要过多关注具体数值因为随机初始化训练 CIFAR-10 的精度受多种因素影响。你需要关注的是趋势train_loss 是否持续下降train_acc 和 val_acc 是否同步上升。只要符合这个趋势就说明模型正在正常学习。6.3 如何判断是否正常收敛判断标准有三条训练损失稳定下降没有出现大幅震荡。验证准确率随训练轮数逐渐上升而不是长时间在原地徘徊。训练集准确率和验证集准确率的差距没有越拉越大。如果两者差距过大说明过拟合已经开始需要增加正则化手段或数据增强。随机初始化在 CIFAR-10 上训练 20 轮精度会明显低于 ImageNet 预训练迁移的效果这是正常现象。要追求更高精度最有效的做法是加载预训练权重而不是盲目增加训练轮数。6.4 训练变慢或精度异常的调整方向如果训练过程明显偏慢优先检查是否真的在使用 GPU。在train.py中--device默认自动选择 CUDA如果显存不足或 CUDA 不可用会自动回退到 CPU。你可以手动指定--device cuda确认 GPU 是否参与计算并通过nvidia-smi查看显存占用。如果精度长期不提升按照这个顺序排查先确认数据预处理是否和模型匹配例如训练集和验证集 transform 是否一致再看学习率是否过大或过小AdamW 默认学习率 1e-3 对大多数 CNN 可用ViT 通常需要更低的学习率最后检查损失函数是否正确单标签多分类统一使用 CrossEntropyLoss它内部已包含 Softmax不要在模型输出层再手动加 Softmax否则梯度计算会出问题。7. 常见问题与排查方法下面是这个模板运行过程中最常遇到的问题均以表格形式整理方便收藏后对照排查。问题现象可能原因排查方式解决方案运行报错维度不匹配分类头没有正确替换打印模型结构print(model)检查最后一层输出维度根据实际类别数重新替换分类头用in_features自适应获取输入维度找不到model.heads.head属性torchvision 版本差异ViT 分类头命名不同print(model)查看 heads 结构按实际属性名替换分类层CIFAR-10 下载失败网络问题或下载源不可达检查网络查看报错信息手动下载数据集放到data/目录后重试训练速度极慢没有使用 GPU 或 batch_size 过大执行nvidia-smi看 GPU 状态查看命令行 device 参数使用 GPU 训练或调低 batch_size 直至显存能容纳验证精度远低于训练精度验证时忘记model.eval()检查验证函数是否设置为 eval 模式在验证循环前调用model.eval()训练前调用model.train()损失出现 NaN学习率过大或数据异常观察第几个 epoch 出现 NaN调低学习率检查输入数据是否包含异常值ViT 训练很慢ViT 参数量大无预训练收敛慢查看参数量sum(p.numel() for p in model.parameters())降低输入分辨率使用预训练权重或降低 batch_sizebatch_size 报显存不足模型参数量大且 batch_size 设置过高查看报错的 CUDA out of memory 信息调节 batch_size 或降低输入分辨率这里特别提醒一个容易忽略的点调试时优先使用小数据集和小模型。先用resnet18跑通流程再切换成vit不要一上来就在大模型上报错后同时排查数据和代码这样效率极低。把“先跑通再跑好”作为固定习惯能省下大量排查时间。8. 从训练到部署的工程建议训练只是模型生命周期的一部分。真正的生产环境还需要完成导出、验证、部署和监控。这一节给出从 PyTorch 模型导出到 ONNX 的完整流程以及部署端必须注意的工程细节。8.1 导出 ONNX 的通用流程ONNX 是开放神经网络交换格式可以让 PyTorch、TensorFlow 等多个框架的模型在统一的中间表示上运行是目前部署最通用的模型格式之一。导出代码只需在训练好的权重上执行一次前向PyTorch 会记录计算图并转换为 ONNX。# 文件路径export_onnx.py import torch from models import build_model model build_model(resnet18, num_classes10) model.load_state_dict( torch.load(checkpoints/best_resnet18.pth, map_locationcpu) ) model.eval() dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, resnet18.onnx, input_names[input], output_names[output], dynamic_axes{ input: {0: batch}, output: {0: batch}, }, opset_version17, ) print(ONNX export done)dynamic_axes定义了动态维度。上面把 batch 维标记为动态意味着推理时一次可以输入 1 张图片也可以输入 32 张图片模型结构不会变化。如果不设置动态轴导出的模型会固定 batch 大小为 1部署灵活性会大大降低。8.2 部署端注意事项ONNX 文件导出后强烈建议先在本地加载并验证输出是否与 PyTorch 原模型一致再做上线操作。验证方法是准备同一张输入图片分别用 PyTorch 和 ONNX Runtime 推理比较输出结果。推理框架如 ONNX Runtime 和 TensorRT 在算子支持上有差异个别网络层可能转换失败务必在实际部署环境完整测试一遍。部署端另一个高频踩坑点是输入预处理不一致。训练阶段使用 Resize(256) CenterCrop(224) Normalize 的统计量这些参数必须完整复刻到部署代码里。很多线上模型效果偏差不是模型训练有问题而是部署端的图片预处理与训练时不一致导致输入分布漂移。8.3 工程规范与安全提醒以下几条是生产环境的基本要求模型文件纳入版本管理记录训练参数、数据版本和评估指标方便回溯。权重文件要备份保存时建议同时保存state_dict和训练超参数配置文件便于复现。在测试集上评估后再导出不要只看验证集指标防止验证集上过拟合。涉及敏感数据或生产系统时遵循最小权限原则模型服务使用独立账号和受限环境运行。使用预训练权重前检查模型许可证是否符合你的业务场景尤其是商业用途。9. 总结与后续学习方向这份模板解决的核心问题是把训练流程从“人肉适配每个模型”变成“注册模型 统一训练”让模型转换的成本降到一个命令行参数。四个模型的接入过程也反过来帮助你理解它们的本质差异AlexNet 确立 CNN 的骨架VGG 验证小卷积核堆叠的思路ResNet 用残差连接让深度成为优势ViT 则用自注意力跳出卷积的局部假设走向全局建模。接下来建议你做三件事。第一用这份模板对自己的数据集跑通一次完整流程不要停留在复制代码第二重点研究 ResNet 的残差块和 ViT 的注意力机制它们代表了两种完全不同的特征表达范式第三把 ONNX 导出加入工作流在部署环境里完整验证一遍输入输出一致性。建议收藏这份模板备用后面做分类、迁移学习、模型对比实验都用得上。
返回列表