ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

SeaFormer轻量Transformer图像分类实战:轴向注意力与PyTorch实现

SeaFormer轻量Transformer图像分类实战:轴向注意力与PyTorch实现 简介SeaFormer图像分类实战资料包聚焦轻量级Transformer在移动端图像分类任务中的应用面向有一定PyTorch基础、希望掌握完整训练流程的中级开发者。资源以SeaFormer_T等轻量模型为例配套可运行的训练与测试代码覆盖CutOut、MixUp、CutMix等数据增强手段以及DP多显卡训练、混合精度、梯度裁剪、EMA、余弦退火等训练技巧。包体共2451个文件以2436张可视化结果图包括损失曲线、ACC曲线、Grad-CAM热力图为主体另有8个Python脚本、权重文件、类别映射JSON和TAR压缩包整体约768MB。已有1014人学习适合复现论文、迁移图像分类任务或参考训练管线时使用。读者可直接运行脚本完成训练、验证与测试并获得测评报告和可视化结果省去从零搭建的时间。1. 为什么用SeaFormer做图像分类轻量Transformer里的低开销选择当transformer图像分类模型在ImageNet榜单上把精度越推越高的时候我所在的工业视觉团队反而把目光收回到部署这件事上。手机、边缘盒子、工业相机后面的小主机算力没有数据中心那么大方SeaFormer这类轻量Transformer就成了分类任务落地的主力方案。它把全局注意力拆成横向和纵向两次计算又用一层小卷积兜住局部信息所以不像ViT那样动辄上亿参数也能在森林图像分类、零件缺陷分类这类自定义数据集上跑出比同规模CNN更稳的精度。这篇实战笔记会从网络结构讲起直接落到一份可运行的PyTorch训练流程和几个我在项目里踩过的高频坑。2. 把SeaFormer拆开看轴向注意力、卷积增强和一次完整前向2.1 为什么图像分类模型要算两次一维注意力Transformer做图像分类的核心操作是自注意力要把特征图每个位置和所有其他位置做相似度计算。输入分辨率固定在224的时候最后一层特征图如果还有14×14序列长度就是196计算量还能接受可一旦图像换成大图或者特征图分辨率保持到28×28序列长度变成784全局注意力的矩阵乘法就会急剧膨胀。对边缘部署场景来说这多出来的几分之一秒都可能让整个产线节拍变慢。SeaFormer的解决办法是把二维注意力拆成两步先在高度方向上做一次全局自注意力再在宽度方向上做一次全局自注意力也就是轴向注意力。这样每个位置仍然能看到整张图的信息只不过把一次N×N的计算换成了两次N×√N量级的计算。用大白话讲原来是所有人一起开会现在改成按行开会、按列开会会开两次信息交换也没少多少。第二个关键点是分类任务和分割任务对特征的需求不太一样。图像分类更看重全局语义但也不能丢了边缘、纹理这些局部细节。纯Transformer结构常常在早期阶段就做patch embedding把图像切成一个个不重叠的小块局部纹理信息容易被截断。SeaFormer在结构里加入了一个卷积增强分支专门负责补回这部分局部表达。这个分支不负责把特征图变大只负责在原有通道空间里做局部信息融合这种混合结构也是它能在精度和速度之间取得平衡的原因。还有一些实现会把相对位置偏置加进注意力矩阵我在小数据集上用下来是负优化。全局注意力需要位置编码是因为patch的sequence是一维拍扁的轴向注意力在height和width两个方向分别做天然保留二维结构感不额外编码也能让模型感知到上下、左右关系。这个特性对图像分类是加分项省掉位置编码也就少了一处部署时要处理的动态shape逻辑。说回选型。如果你正在对比图像分类算法常见备选方案有三条线一条是MobileNetV3这种纯卷积部署简单但精度到后期靠堆深度才能涨一条是ViT/MobileViT这类Transformer全局建模能力强但工程化要处理的位置编码、归一化层比CNN多第三条就是SeaFormer这条轴向注意力加卷积分支的路线它把全局注意力拆细参数量少在边缘推理框架里又比标准Multi-Head Attention更容易被优化。实测下来在同类FLOPs下它通常比MobileNet高1到2个点和MobileViT相近但推理时延更低。下表是我在边缘设备上选型时的一个对比口径按真实项目里最常见的几类需求写数值是相对比较不是跑分精确值。方案全局建模方式部署复杂度精度水准更推荐的使用场景MobileNetV3全程卷积极低同FLOPs下偏低数据量小、工期紧、对精度要求中等ViT-Tiny全局多头注意力中需要足够数据支撑预训练权重齐全的通用分类MobileViT局部全局混合中偏高精度好但算子碎有成熟部署团队、算子能合入框架SeaFormer轴向注意力卷积增强低同FLOPs下稳且高移动端、边缘盒子、自定义数据集2.2 最小可运行的SeaFormer块代码与参数说明我一般会自己维护一个可跑的版本思想与论文对齐细节按工程简化。下面是能直接塞进训练脚本的核心block。import torch import torch.nn as nn import torch.nn.functional as F class DropPath(nn.Module): 随机丢弃整条残差路径训练时用推理时恒等。 def __init__(self, p0.1): super().__init__() self.p p def forward(self, x): if not self.training or self.p 0: return x keep_prob 1 - self.p mask x.new_empty(x.shape[0], 1, 1, 1).bernoulli_(keep_prob) return x * mask / keep_prob class AxialAttention(nn.Module): 单方向轴向注意力axish 时沿高度做全局自注意力axisw 时沿宽度做。 def __init__(self, dim, num_heads8, axish): super().__init__() self.num_heads num_heads self.axis axis self.scale (dim // num_heads) ** -0.5 self.norm nn.LayerNorm(dim) self.qkv nn.Linear(dim, dim * 3, biasFalse) self.proj nn.Linear(dim, dim) def forward(self, x): # 输入 x: (B, C, H, W) B, C, H, W x.shape if self.axis h: # 把高看成序列长度宽拼进 batch 维度 x x.permute(0, 3, 2, 1).reshape(B * W, H, C) S, N H, B * W else: # 把宽看成序列长度高拼进 batch 维度 x x.permute(0, 2, 3, 1).reshape(B * H, W, C) S, N W, B * H q, k, v self.qkv(self.norm(x)).chunk(3, dim-1) q q.reshape(N, S, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3) k k.reshape(N, S, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3) v v.reshape(N, S, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) out attn v out out.transpose(1, 2).reshape(N, S, C) out self.proj(out) if self.axis h: return out.reshape(B, W, H, C).permute(0, 3, 2, 1) return out.reshape(B, H, W, C).permute(0, 3, 1, 2) class SqueezeConvBranch(nn.Module): 卷积增强分支1x1 扩张通道depthwise 提局部1x1 压回原通道。 def __init__(self, dim, expand_ratio4): super().__init__() hidden dim * expand_ratio self.pointwise nn.Conv2d(dim, hidden, 1) self.depthwise nn.Conv2d(hidden, hidden, 3, padding1, groupshidden) self.reduce nn.Conv2d(hidden, dim, 1) def forward(self, x): identity x x F.gelu(self.depthwise(F.gelu(self.pointwise(x)))) x self.reduce(x) return x identity class SeaFormerBlock(nn.Module): 一个完整 block卷积增强 横向轴向注意力 纵向轴向注意力。 def __init__(self, dim, num_heads8, drop_path0.1): super().__init__() self.conv_branch SqueezeConvBranch(dim, expand_ratio4) self.attn_h AxialAttention(dim, num_heads, axish) self.attn_w AxialAttention(dim, num_heads, axisw) self.drop_path DropPath(drop_path) def forward(self, x): x x self.drop_path(self.conv_branch(x)) x x self.drop_path(self.attn_h(x)) x x self.drop_path(self.attn_w(x)) return x这里有三个参数值得盯着调dim是通道数太小会让注意力学不到长程关系太大在边缘设备上内存会先吃紧num_heads我习惯在stage3之后设成8早期stage设4注意力头太碎在小数据集上反而不稳drop_path是残差结构的随机丢弃概率小数据集用0.1足够拿ImageNet做预训练再微调时可以降到0.05。代码里最需要注意的地方是轴向注意力的shape变化。输入是(B,C,H,W)height方向注意力先把H变成序列长度W拼到batch上算完再还原。这里permute和reshape的顺序搞反结果不会报错但注意力会作用在错误方向上训练损失呈一条直线。我第一次实现时就是看论文里的图想当然写折腾了两天才发现是两个维度的还原顺序错了。2.3 前向尺寸推算与一次验证把block组装成完整网络之前先推算一下224×224输入经过每一步的尺寸这一步能筛掉半数结构错误。stem里两个stride2的卷积会把空间尺寸压到56×56stage2末尾是28×28stage3末尾是14×14stage4末尾是7×7。轴向注意力只出现在后两个阶段每个block里先做高注意力再做宽注意力两次attention的序列长度分别是14或7量级比全局注意力小一个维度。我习惯在写完整模型之前先跑一个最小前向确认logits尺寸符合预期if __name__ __main__: from seaformer_blocks import SeaFormerBlock block SeaFormerBlock(dim64, num_heads4) x torch.randn(1, 64, 56, 56) y block(x) print(y.shape) # torch.Size([1, 64, 56, 56])输出shape没有变化说明残差结构写对了如果输出少了一半或者维度对不上问题大概率出在轴向注意力还原时用的reshape参数上而不是forward的下一行。这一层验证花不了十秒钟但能省掉后续整个训练排错的力气。3. 用SeaFormer跑通图像分类训练数据、模型和训练配置3.1 图像分类数据集下载与目录规范不管是从公开数据集下载ImageNet子集还是自己收集森林图像分类这类自定义场景数据我都会先把数据整理成PyTorch ImageFolder能用的目录结构。大类在下一层小类在再下一层。train和val分开val里每个类至少要保留20张以上否则训练曲线看着很好换一批真实图片就露馅。data/ ├── train/ │ ├── forest_needle/ │ ├── forest_broadleaf/ │ └── grassland/ └── val/ ├── forest_needle/ ├── forest_broadleaf/ └── grassland/图像分类数据集下载回来经常是一个tar包里面套了好几层目录。我习惯先确认图片数量再跑一遍坏图检查坏图会在训练中途直接让DataLoader崩掉报错信息还特别隐蔽。from PIL import Image import os for root, _, files in os.walk(./data): for f in files: if f.lower().endswith((.jpg, .jpeg, .png)): try: Image.open(os.path.join(root, f)).verify() except Exception: print(bad image:, os.path.join(root, f))坏图检查跑完再用下面的方式加载数据from torchvision.datasets import ImageFolder from torchvision import transforms from torch.utils.data import DataLoader normalize transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) train_tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), normalize, ]) val_tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), normalize, ]) train_set ImageFolder(./data/train, transformtrain_tf) val_set ImageFolder(./data/val, transformval_tf) train_loader DataLoader(train_set, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue) val_loader DataLoader(val_set, batch_size128, shuffleFalse, num_workers8, pin_memoryTrue) print(类别数:, len(train_set.classes)) print(训练集:, len(train_set), 验证集:, len(val_set))这段代码里比较讲究的是RandomResizedCrop(scale(0.7, 1.0))。如果数据集里目标的尺寸比较一致比如工业零件这个范围可以收到(0.8, 1.0)避免每次都裁出一个很小区域如果做森林图像分类这类目标尺度变化大的地面图像用默认(0.08, 1.0)也行但训练周期得拉长。Resize(256)再CenterCrop(224)是验证阶段最经典也最稳的配置别直接Resize((224, 224))那样会把图像长宽比压变形验证集精度会少0.3到0.5个点。3.2 搭建SeaFormer分类模型分类头与前向在2.2的block基础上我通常直接组织成一个四阶段的分类网络。前两个stage可以不放注意力只堆卷积分支因为早期分辨率高轴向注意力的计算量依然不小到stage3再插入注意力块这样全局信息在高语义层被交换效率最高。class SeaFormerTiny(nn.Module): def __init__(self, num_classes1000): super().__init__() self.stem nn.Sequential( nn.Conv2d(3, 64, 3, stride2, padding1), nn.BatchNorm2d(64), nn.GELU(), nn.Conv2d(64, 64, 3, stride2, padding1), nn.BatchNorm2d(64), nn.GELU(), ) self.stage1 nn.Sequential( SqueezeConvBranch(64, expand_ratio4), SqueezeConvBranch(64, expand_ratio4), ) self.stage2 nn.Sequential( nn.Conv2d(64, 128, 3, stride2, padding1), SeaFormerBlock(128, num_heads4, drop_path0.1), SeaFormerBlock(128, num_heads4, drop_path0.1), ) self.stage3 nn.Sequential( nn.Conv2d(128, 256, 3, stride2, padding1), SeaFormerBlock(256, num_heads8, drop_path0.1), SeaFormerBlock(256, num_heads8, drop_path0.1), SeaFormerBlock(256, num_heads8, drop_path0.1), SeaFormerBlock(256, num_heads8, drop_path0.1), ) self.stage4 nn.Sequential( nn.Conv2d(256, 512, 3, stride2, padding1), SeaFormerBlock(512, num_heads8, drop_path0.1), SeaFormerBlock(512, num_heads8, drop_path0.1), ) self.head nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(512, num_classes), ) def forward(self, x): x self.stem(x) x self.stage1(x) x self.stage2(x) x self.stage3(x) x self.stage4(x) return self.head(x)forward里没有额外做激活因为分类头后面直接接CrossEntropyLoss交叉熵内部会把logits转成概率在head里提前做softmax反而会损害数值稳定性。nn.Conv2d作为下采样没有配BN这是故意的我见过的项目里在stage边界加BN有时会导致相邻stage输出的量纲不一致注意力分数一波动训练就变得敏感。如果你在自己数据上发现深层loss震荡再考虑在每个下采样后面补一个BN。这个轻量版参数量大致控制在10M级别单张224分辨率在边缘GPU上推理一遍约10到20毫秒具体取决于设备。如果你需要更高精度把stage3的block数量从4加到8stage4从2加到4就是seaformer_small级别的规模如果你要做移动端实时分类把stem的第一个卷积改成stride4的patch embed后两个stage各减一个block精度会掉一点但速度能上来一截。3.3 训练循环与参数设置优化器、标签平滑和EMA训练部分我给出目前最顺手的配置AdamW做优化器余弦退火管学习率混合精度降显存EMA给最后模型做加持。import torch import torch.nn as nn from torch.cuda.amp import autocast, GradScaler model SeaFormerTiny(num_classeslen(train_set.classes)).cuda() criterion nn.CrossEntropyLoss(label_smoothing0.1) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) scaler GradScaler() ema_model SeaFormerTiny(num_classeslen(train_set.classes)).cuda() ema_model.load_state_dict(model.state_dict()) ema_decay 0.999 for epoch in range(100): model.train() for images, labels in train_loader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) with autocast(): logits model(images) loss criterion(logits, labels) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() # EMA 更新直接在线更新权重 with torch.no_grad(): for ema_param, param in zip(ema_model.parameters(), model.parameters()): ema_param.data.mul_(ema_decay).add_((1 - ema_decay) * param.data) scheduler.step() val_acc evaluate(model, val_loader) ema_acc evaluate(ema_model, val_loader) print(fepoch {epoch1:03d}, acc{val_acc:.4f}, ema_acc{ema_acc:.4f})label_smoothing0.1是这个配置里对最终精度帮助最大的一项。它把one-hot标签往均匀分布推了一截模型不会为了把训练集置信度顶到99.9%而过分放大最后一层权重验证集精度通常能稳定涨0.3到1个点。EMA更新放在每个step里做ema_decay在0.999到0.9995之间取训练步数少的任务可以设到0.997否则EMA权重追不上模型变化反而拖累精度。autocast包裹的是前向和loss计算反向传播不需要单独处理梯度缩放交给scaler。每个step都做EMA更新代价只是多一次权重拷贝对显存几乎无影响。验证时用ema_model平滑后的权重往往落在一个更平缓的损失区域内比直接训练出来的模型泛化更好。关于Batch size和学习率我的经验法则是Batch size从64涨到256时学习率从1e-3同步放大到2e-3比较稳不要直接翻到4e-3。SeaFormer里LayerNorm和BN混用学习率过大时先崩的总是BN统计量表现是前几个epoch正常然后loss突然跳高。3.4 验证函数与训练监控evaluate函数在训练循环里被调用了我把实现放在这里。验证阶段要关闭梯度、切到eval模式。DropPath的self.training控制无需额外操作但BatchNorm在eval模式下使用running统计量务必切model.eval()。torch.no_grad() def evaluate(model, loader): model.eval() correct 0 total 0 for images, labels in loader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) with autocast(): logits model(images) pred logits.argmax(dim1) correct (pred labels).sum().item() total labels.size(0) return correct / totalpred logits.argmax(dim1)在fp16下和fp32下结果可能有一次位宽的微小差异在验证集上通常影响不到0.1个点。需要精确复现指标时验证阶段可以关掉autocast或者先logits.float()再argmax。从效率角度我一般留着autocast验证集的吞吐量往往决定调参效率。监控指标上除了整体acc还要顺带看每个类的召回率。用sklearn.metrics.classification_report打印一次重点看是不是有个别类永远分错。如果某个类召回率不到60骨子里不是模型问题而是该类训练样本太少需要回到3.1的数据增强上去给那个类单独加旋转或尺度扰动而不是在全数据集上加更强的增强。4. SeaFormer实战避坑与排查5个困扰我一周的问题方向这一章记的都是我自己在图像分类项目里真实翻过车的地方。每一条按现象、原因、解决三段式写方便对号入座。4.1 损失一直降不下来卡在0.8上下一动不动现象训练了十几个epoch训练集损失在0.8到0.9之间震荡分类精度也一直很低。前几个epoch下降很快后面完全停滞。原因最常见是初始学习率设置偏高AdamW的weight_decay又偏大。SeaFormer这种混合了BN和LayerNorm的结构对优化器的超参比纯CNN更敏感。另一个可能出现在2.2的代码上轴向注意力里qkv之后的reshape写错注意力作用到了错误轴上模型等于永远在跟自己打架。解决先把学习率降到3e-4跑20个epoch排除优化器问题。然后打印单个batch的前向特征图如果方差没有发散说明结构没有写穿。最后试一个trick把weight_decay临时设成0看loss是否松动。如果能松动说明L2正则把注意力权重压得太死通常降到0.02到0.05之间即可。4.2 训练精度98验证精度连60都不到现象训练集损失一路下到0.1附近训练精度刷到98但验证集只有50到60还随不同run波动很大。原因数据划分不均匀。图像分类数据集下载回来经常按类别目录放但很多人直接从一个总目录里按比例随机切train/val没有按类别做分层抽样。某个类别在验证集里只分到一两张正常样本其他全是遮挡图精度自然上不去。训练时图片增强过猛验证集又没有任何增强也容易让指标看起来差距过大。解决换成分层抽样。用scikit-learn按label做切分from sklearn.model_selection import StratifiedShuffleSplit from torchvision.datasets import ImageFolder dataset ImageFolder(./data_all) labels [s[1] for s in dataset.samples] split StratifiedShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(split.split(dataset.samples, labels))分完之后检查验证集每个类最少样本数低于5张的类别建议合并或补拍。数据切分永远是这类问题的第一排查点别一上来就调网络结构。4.3 混合精度开启后loss偶尔跳为NaN现象训练到中途某个batch的loss突然变成nan然后后面所有step都继续nan只能重跑。关掉autocast就正常。原因fp16可表示的数值范围有限。logits的绝对值超出65504时softmax计算会达到inf对fp16的loss来说就是溢出。多发生在训练起步、lr最大时或者分类头初始化权重过大时。输入里出现坏图、像素值全为0的图也会在极端情况下把它引爆。解决给head的Linear做更小的初始化或者把loss放在fp32里算with autocast(): logits model(images).float() # 先把 logits 转回 fp32 loss criterion(logits, labels)这个改动几乎不影响训练速度又能避开fp16溢出的边界。再把scaler.set_growth_interval(2000)放宽一点让GradScaler对loss增长的判断更保守也能减少中间爆nan的概率。4.4 加预训练权重后精度反而比随机初始化低现象加载一个在通用大图上预训练好的SeaFormer权重做微调训练loss能降但验证精度比只用随机初始化从零训练低了3到5个点。原因预训练模型的分类头输出维度是1000类自定义数据集只有十几个类直接截断会丢信息。如果不做分层学习率lr太大时头部权重被破坏主干也被冲得厉害。另一个常见问题预训练权重里包含num_batches_tracked这类BN专属keystrict加载后报key不匹配很多人图省事直接把权重删掉当随机初始化用。解决先兼容key差异再给主干和头部分配不同lr。backbone_params [p for n, p in model.named_parameters() if head not in n] head_params [p for n, p in model.named_parameters() if head in n] optimizer torch.optim.AdamW([ {params: backbone_params, lr: 3e-4}, {params: head_params, lr: 1e-3}, ], weight_decay0.05)不管vit、cnn还是SeaFormer自定义数据集微调的第一原则都是head用大lr快速适配主干用小lr维持已学到的表示。50个epoch内就能看到效果比单一口径lr省事得多。4.5 部署时精度和训练时对不上低2个点以上现象训练、验证都在GPU上的PyTorch里跑精度92。导出成ONNX后放到边缘设备推理精度掉到89甚至更低。原因不是ONNX算子问题而是训练和部署的输入管线不一致。训练用了RandomResizedCrop验证用了Resize加CenterCrop但部署侧常常直接把摄像头原始画面缩成正方形或没有归一化除以255。另外BatchNorm在部署转换时被折叠进卷积前提是统计量冻结而某些导出工具在动态batch时会对norm层处理得不干净。解决先做三重对齐。第一用PIL把单张图完整走一遍val_tf后再喂模型第二把喂进去的tensor保存下来部署代码里也用同一份预处理第三确认归一化的mean/std一致很多背景分类数据集的像素分布和ImageNet差异很大直接用ImageNet的mean/std会有偏差。这三项对齐做完部署精度和训练精度通常能回到0.1个点以内。5. 把SeaFormer精度再顶一截蒸馏、EMA和ONNX验证训练收敛之后如果还想把精度往上提我第一个会做的不是改网络结构而是知识蒸馏。用一个已经训好的大模型当老师SeaFormer当学生蒸馏损失加在logits上。我常用的损失写法是def distil_loss(student_logits, teacher_logits, labels, T3.0, alpha0.7): ce nn.CrossEntropyLoss()(student_logits, labels) kl nn.KLDivLoss(reductionbatchmean)( F.log_softmax(student_logits / T, dim1), F.softmax(teacher_logits / T, dim1) ) return alpha * kl * (T * T) (1 - alpha) * ceT * T这个系数是让KL损失的梯度和cross entropy在同一个量级上。alpha0.7表示70%的权重在蒸馏损失上30%还留在真实标签上避免学成老师的错误。这个操作在分类任务里稳定提升0.5到1个点代价只是多一次前向。EMA模型和当前模型的精度对比也值得盯住。如果ema_acc始终比当前模型低把ema_decay从0.999改到0.995让更新更快地跟随当前权重。如果ema_acc稳定高0.3以上说明模型已经进入过拟合区可以提前停掉训练以ema模型为最终交付。最后的验证动作是ONNX导出。导出时把动态batch打开同时固定分辨率避免部署引擎对动态尺寸做无意义的优化。dummy torch.randn(1, 3, 224, 224).cpu() torch.onnx.export( ema_model.cpu(), dummy, seaformer.onnx, input_names[input], output_names[logits], dynamic_axes{input: {0: batch_size}, logits: {0: batch_size}}, do_constant_foldingTrue, )导出后不要急着部署先用onnxruntime和PyTorch各跑一遍相同输入对比logits差的绝对值上界。差异超过1e-3就要回头检查预处理而不是去怀疑推理引擎。我习惯把这个对比写成脚本每次模型更新都自动跑一遍。在图像分类这种任务上模型结构决定精度上限训练配置决定能不能摸到上限部署对齐决定最终落地上限。说实话SeaFormer并不是榜单上最亮眼的模型但在我们这种算力受限的边缘场景里它是少有的从训练到部署都让人省心的选择。如果你正在边缘算力上做图像分类部署这个方向值得认真投入。希望帮到你。本文还有配套的精品资源点击获取
返回列表