ARTICLE DETAIL

资讯详情

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

RepVgg图像分类实战:从训练到部署的结构重参数化全解析

RepVgg图像分类实战:从训练到部署的结构重参数化全解析 简介RepVgg图像分类实战资源面向深度学习入门到进阶的读者提供一套可直接运行的模型训练与推理项目。内容紧扣“VGG式”设计无分支结构、仅用3x3卷积和ReLU激活帮助理解plain架构在图像分类任务中的具体实现与调优方式。压缩包共2000个文件主体为2435张png图片构成整理好的分类数据集配套12个Python脚本贯穿数据加载、网络搭建、训练验证与推理预测另有2个json文件保存类别映射与配置参数、2个pth预训练权重可直接加载使用1个txt说明文档辅助快速上手整体大小约986.61MB能够支撑完整的本地复现。目前已有990人学习下载适合用于论文基准实验、课程设计或模型对比。通过这套资源可省去收集数据、整理格式和漫长预训练的时间直接基于数据集与权重开展训练和验证同时从源码层面拆解VGG式结构的设计细节帮助读者在图像分类项目中快速落地RepVgg。1. RepVgg实战图像分类项目从训练到推理的完整闭环这个项目的文件少得不像一个完整的图像分类工程一个class.json一个result.json几张测试图片。但真正用过RepVgg做图像分类的人会懂难点从来不在模型本身而在训练态和推理态怎么切换得干净。RepVgg训练时是三分支结构推理时却能重组回普通VGG那种无分支、只堆3x3卷积的plain网络速度不掉、精度不降。这套实战资源正好把流程串成一个闭环从class.json读取类别用图片生成数据集训练RepVgg分类模型最后把预测结果写回result.json并生成带标签的可视化图。想跑通最新图像分类模型RepVgg又不想从头趟一遍结构重参数化雷区的工程师拿它做底板最合适。项目更适合单标签分类场景比如森林图像分类、常见物体分类这类需求明确的任务拿到手就能改成自己的类别清单。2. 从VGG到RepVgg先搞懂三分支为什么能融合成一层卷积2.1 VGG式结构的三个硬约束摘要里提到“VGG式”的定义原文其实给了三个硬约束没有任何分支结构也就是plain或feed-forward架构只使用3x3卷积只使用ReLU作为激活函数。这三条约束放在今天看有点复古但每个字都有部署层面的考量。没有分支内存访问顺序就是线性的推理引擎不需要维护多条数据通路只用3x3卷积意味着整个模型只有一个算子类型Winograd、im2col等优化可以集中在同一种卷积上ReLU不参与参数和统计量的重参数化融合时不用做额外处理。所以VGG式结构在CPU和GPU上的运行效率都很高尤其适合对延迟敏感的图像分类服务端。RepVgg继承了这三个约束但它的高明之处是只在推理阶段满足这三条训练阶段还是用多分支网络提升表达力。很多人第一次看到RepVgg名字会误以为它是一个新模型结构实际上它更像是一套“训练完再剪枝”的方法论。计算图在训练结束后会被重写最终落地的模型已经不是一个带Branch的网络而是等价但更快的VGG风格卷积栈。2.2 训练态为什么是三分支而不是两条分支ResNet证明残差连接能有效改善梯度传播但残差连接在推理时也保留这会让每个Block都多一次加法显存和带宽都增加。RepVgg的出发点是能不能训练时享受残差的红利推理时又完全不保留分支结构。答案是用结构重参数化。训练态每个Block有三个分支3x3卷积、1x1卷积、Identity连接。3x3是主分支负责感受野内的特征提取1x1分支从通道维度做重排和3x3输出的语义信息互补Identity分支把原始特征直接传过去给梯度留一条直达通道。三个分支的输出逐元素相加再过一个ReLU。选择三条分支而不是更多是从计算量和有效性的平衡点考虑。再多加如5x5或7x7分支虽然也能融合但会把训练态FLOPs推高收益又不明显实际训练时间不太划算。这里还有一个容易误解的地方1x1和Identity为什么能融合进3x3因为卷积是线性算子。1x1卷积本质上是一个中心点有权重、其余位置为0的3x3卷积Identity则可以被视为一个单位矩阵形式的3x3卷积中心权重为1其余为0。三种卷积在同一个卷积核空间里是“可加”的所以理论上无论多少个分支只要补零对齐都能合并成一个卷积核。这个性质是结构重参数化能够成立的基础。2.3 ConvBN融合的数学操作要让三分支真正合并第一步是把每个分支里的BatchNorm参数卷进卷积层。BatchNorm在训练时用的是batch统计量推理时用的是running_mean和running_var所以融合必须发生在推理阶段不能拿训练阶段的状态做。融合公式我一般这样拆解假设卷积输出为WxBN的缩放参数是gamma偏置是beta则推理输出为gamma * (Wx - mean) / sqrt(var eps) beta。把它整理成新的卷积权重W W * gamma / sqrt(var eps)和新偏置b beta - mean * gamma / sqrt(var eps)整个表达式就是Wx b等价于一个带偏置卷积。下面这段是把一个Conv2d和BatchNorm2d合并的辅助函数import torch def fuse_bn_to_conv(conv, bn): 将Conv2d与BatchNorm2d合并为一个带偏置的Conv2d。 conv: 已有卷积层bias为False bn: 训练好的BatchNorm2d 返回合并后的weight和bias gamma bn.weight.data beta bn.bias.data mean bn.running_mean var bn.running_var eps bn.eps scale gamma / torch.sqrt(var eps) new_weight conv.weight.data * scale.reshape(-1, 1, 1, 1) new_bias beta - mean * scale return new_weight, new_bias参数scale是逐输出通道的所以要reshape成[out_ch, 1, 1, 1]才能和卷积权重相乘。如果卷积本身已经有bias也要先把原bias折进公式里公式会再复杂一点。项目里通常设置biasFalse因为后面接BN本身就会引入偏置没必要重复。注意融合得到的new_bias是逐通道的构造新Conv2d时必须设置biasTrue并赋值否则卷积输出会整体偏置准确率会崩掉。2.4 训练态RepVggBlock的PyTorch实现在工程代码里我一般把RepVggBlock定义成下面这种结构保证训练态能正常反向传播同时给后续融合留出清晰的接口import torch import torch.nn as nn import torch.nn.functional as F class RepVggBlock(nn.Module): def __init__(self, in_ch, out_ch, stride1): super().__init__() self.in_ch in_ch self.out_ch out_ch self.stride stride self.conv3x3 nn.Conv2d(in_ch, out_ch, kernel_size3, stridestride, padding1, biasFalse) self.bn3x3 nn.BatchNorm2d(out_ch) self.conv1x1 nn.Conv2d(in_ch, out_ch, kernel_size1, stridestride, padding0, biasFalse) self.bn1x1 nn.BatchNorm2d(out_ch) self.use_identity (in_ch out_ch and stride 1) if self.use_identity: self.bn_identity nn.BatchNorm2d(in_ch) def forward(self, x): y self.bn3x3(self.conv3x3(x)) y y self.bn1x1(self.conv1x1(x)) if self.use_identity: y y self.bn_identity(x) return F.relu(y)这里有几个参数需要注意。stride是2时特征图尺寸减半Identity分支没办法做对齐所以用use_identity来判断是否启用1x1分支的padding为0但在后续融合时它会被pad成一个3x3卷积核中心位置是原权重四周补零。第一层下采样通常会在stem上完成后续层保持stride1让分辨率缓慢下降。2.5 从Block到完整分类模型的结构设计有了Block之后怎么搭整个分类模型常见的做法是参考官方RepVgg-A0或B0的配置先用一个stride2的常规3x3卷积做stem然后用多个stage堆叠RepVggBlock每个stage的最后一块把stride设为2通道数翻倍最后接全局平均池化和全连接分类头。实际项目里不需要生搬硬套官方结构我会根据图像尺寸调整stage数量。比如224x224输入可以用4个stage而森林图像分类里如果原图是512x512可以考虑把下采样次数从5次降到4次保留更多纹理信息。这段代码里没有直接写完整模型类因为核心难点在Block上。完整模型类在资源代码里已经给出训练时只要传入类别数就能用。真正要留意的不是模型结构而是训练完切换到推理态的过程这一步做不好前面所有训练都白费。3. 数据准备把class.json、result.json和多张图片组织成训练集3.1 项目文件里class.json、result.json和散图各自干什么打开项目目录第一眼看到的是result.json、class.json和一堆PNG图片图片名类似77291b3ad.png、5a8b75712.png。class.json存的是类别映射result.json是模型预测的输出结果PNG图片是用于验证预测效果的真实样本。这套组合意味着你不需要自己造标签体系只需要把class.json读进来再定义好Dataset即可开始训练。不过散图本身没有目录级标签不能直接拿来当训练集它们更适合做一次前向推理的冒烟测试。我习惯先把它们复制到带类别名的目录里再让Dataset按目录扫描。这样既验证了推理也验证了数据加载链路后面换正式数据集时不需要改代码。3.2 从class.json里取出有序类别名class.json在不同项目里格式不一样常见的有两种一种是{cat: 0, dog: 1}另一种是{0: cat, 1: dog}。为了让标签顺序和模型输出序号一致我通常会先做一层兼容处理不直接依赖原始格式import json with open(class.json, r, encodingutf-8) as f: class_data json.load(f) if isinstance(class_data, dict): keys list(class_data.keys()) first_value class_data[keys[0]] if isinstance(first_value, int): # 键是类别名值是ID class_names [ k for k, v in sorted(class_data.items(), keylambda item: item[1]) ] else: # 键是ID值是类别名 class_names [ v for k, v in sorted(class_data.items(), keylambda item: int(item[0])) ] else: class_names list(class_data) print(类别列表, class_names)这段代码的关键在排序。如果不排序直接遍历dict的原始顺序遇到{cat: 0, dog: 1}没问题但如果JSON文件把dog写在cat前面模型输出就会和真实标签错位。宁可在数据准备阶段多花十秒钟排序也不愿在训练结束后才发现result.json里的ID对不上类别名。排序之后最好再检查一遍是否有重复类别名如果有应该立刻去查class.json的生成逻辑。3.3 自定义Dataset和DataLoader训练集准备通常按目录组织每个子目录是一个类别目录名就是class.json里的类别名。下面这个Dataset实现是我在分类项目里最常用的写法import os from PIL import Image from torch.utils.data import Dataset, DataLoader from torchvision import transforms class ImageFolderDataset(Dataset): def __init__(self, root_dir, class_names, transformNone): self.class_names class_names self.transform transform or transforms.ToTensor() self.samples [] self.labels [] for label, name in enumerate(class_names): folder os.path.join(root_dir, name) if not os.path.isdir(folder): continue for file in os.listdir(folder): if file.lower().endswith((.jpg, .jpeg, .png)): self.samples.append(os.path.join(folder, file)) self.labels.append(label) def __len__(self): return len(self.samples) def __getitem__(self, index): img Image.open(self.samples[index]).convert(RGB) if self.transform: img self.transform(img) label self.labels[index] return img, label transform_train transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) train_set ImageFolderDataset(data/train, class_names, transform_train) train_loader DataLoader(train_set, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue)这里有几个细节容易翻车图片统一转RGB避免灰度图通道数不一致RandomResizedCrop(224)会随机裁剪并缩放给模型提供尺度多样性归一化用的均值方差是ImageNet统计值如果换数据集且统计差异太大建议重新算。shuffle只在训练时开验证时应该改成False方便计算确定性指标。num_workers在Windows下建议改成0否则多进程初始化经常报错。提示如果发现加载慢优先检查是CPU解码瓶颈还是磁盘瓶颈不要在Dataset.__getitem__里做复杂的在线增强那会拖慢整个训练。3.4 数据增广如何选择和VGG式网络的关系RepVgg推理态是纯3x3卷积栈没有注意力机制也没有残差连接所以它对输入分布的敏感度比ResNet高一些。数据增广不只是为了涨点更是为了抹平不同数据来源的统计差异。我一般会在训练集上做随机裁剪、水平翻转和色彩抖动验证集只做resize和中心裁剪。色彩抖动对森林图像分类这种场景特别有用因为森林图片的光照变化大不同时间段拍摄的同一物种会有明显色差。增广强度不宜过大。VGG式网络容量有限过度增强反而会让模型学不到稳定的类别特征。一个通用的判断标准是先不加增广跑10个epoch再逐步加增广观察验证集准确率的浮动范围。如果加了增广后验证集准确率明显下降说明增强幅度已经超过了模型承受范围。3.5 建立val集并预写result.json模板如果项目里只有散图可以先用shell建一个临时目录结构类似于mkdir -p data/train/class_a data/train/class_b cp 77291b3ad.png 5a8b75712.png data/train/class_a/ cp 5d358beb9.png 8029e3396.png data/train/class_b/这样训练集只有四张图可能一两个epoch就过拟合但它能最快暴露代码层面的问题。正式训练时需要按类别把图片分进train和val比例通常是8比2。val集的划分要和class.json保持一致且不能让一张图同时出现在训练集和验证集里。result.json模板可以先写好等推理完成后把预测结果填进去字段一般包含图片路径、预测类别ID、类别名和置信度。这个模板在调试阶段可以帮助快速定位输出是不是已经错位。4. 训练与验证用一套固定配置跑通RepVgg分类4.1 超参数怎么选RepVgg继承VGG的朴素特性训练超参不需要很花哨。如果是森林图像分类这类中规模数据集我一般从下面的配置起步。参数推荐值理由优化器SGD momentum0.9比Adam收敛稳分类任务上泛化好初始学习率0.01大批量时配线性缩放学习率策略CosineAnnealing避免末尾震荡weight_decay1e-4对VGG式结构足够batch_size32或64视显存大小调整过大会影响BN统计epoch20~30小数据集20轮能看到收敛趋势SGD配合momentum在图像分类里依然是最稳妥的组合。学习率0.01只适合batch_size在256附近的情况如果你的batch_size减半建议把学习率也减半。weight_decay不要给太大因为VGG式全卷积结构本身没有那么多冗余参数过大正则反而让模型欠拟合。4.2 训练循环核心代码模型类在项目代码里已经写全我这里只抽训练循环方便看懂每步在做什么model RepVgg(num_classeslen(class_names), deployFalse) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD( model.parameters(), lr0.01, momentum0.9, weight_decay1e-4 ) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max20 ) best_acc 0.0 for epoch in range(20): model.train() running_loss 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) # 训练态走三分支 loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() scheduler.step() val_acc evaluate(model, val_loader) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best.pth) print(fepoch [{epoch 1}/20] loss {running_loss / len(train_loader):.4f} fval_acc {val_acc:.2f}%)这里deployFalse表示当前是训练态三分支全部参与前向。CrossEntropyLoss内部自带softmax不需要在模型输出上手动加。CosineAnnealing的T_max要跟epoch数一致否则余弦曲线会在周期内提前重置学习率变化就不平滑。保存best.pth的依据是val_acc不是loss因为分类任务最终看的是准确率避免训练loss过低但验证集已经过拟合。4.3 评估与模型保存的细节evaluate函数要避免梯度计算用torch.no_grad()包住并且切换到eval模式。RepVgg在eval模式下BatchNorm的统计量会切换到running_mean和running_var不会再用当前batch的统计值这个切换对结果影响很大。def evaluate(model, loader): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: outputs model(images) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return correct / total * 100.0如果total为0需要检查val目录是否为空或者在Dataset里加个断言。argmax选维度1因为outputs形状是[batch, num_classes]。保存模型状态时可以直接保存state_dict也可以保存整个model。我通常只存state_dict这样换框架或做部署时更方便加载。4.4 学习率调度和训练周期的把握学习率是训练里最容易踩坑的参数之一。使用CosineAnnealing时初期学习率下降得慢中期下降加速末段趋于平缓这正好配合RepVgg的特性前期快速收敛后期细调权重。训练20个epoch时前5个epoch可以当作预热观察loss是否明显下降。如果前5个epoch loss几乎不动先确认模型是不是训练态再确认优化器参数是否传对不要急着改学习率。训练周期不是越长越好尤其在小数据集上。RepVgg训练态分支多表达能力比普通VGG强因此过拟合更容易出现在20个epoch之后。判断过拟合的方法很简单看训练loss还在下降、val_acc却开始波动或下降就说明模型进入了记忆阶段这时候应该提前停止而不是继续跑。4.5 从训练日志定位过拟合和欠拟合我习惯在每个epoch末尾把所有指标追加到同一个txt文件里方便后面画曲线echo $(date %F-%T) epoch$epoch loss$loss val_acc$val_acc train.log观察日志时如果训练和验证准确率都低属于欠拟合优先增加模型容量或训练时长如果训练准确率高、验证准确率低属于过拟合加重weight_decay或增强增广如果两者都高但val_acc在训练中后期反复跳动说明学习率末段过大可以把T_max调大让学习率衰减得更慢一些。这套排查逻辑对RepVgg和其他CNN模型通用。5. 部署避坑RepVgg推理阶段最常见的5个问题5.1 训练结束直接推理准确率几乎为0现象用训练好的best.pth加载进去评估测试集准确率从99%掉到1%。原因模型类仍处于deployFalse状态推理走三分支而BatchNorm的running_mean和running_var只会被eval模式锁定。如果直接把模型放在训练模式里去推理BN会使用当前batch的小样本统计量分布会漂移。解决加载权重后先调用model.eval()让BN固定住running统计量。任何推理脚本的入口处都应该在模型加载完就显式设置eval不能留着训练状态。如果还是不对用torch.load的map_location检查权重是否完整并确认加载的是同一份类别清单。5.2 融合后的模型结果和训练态对不上现象手动把三分支融合成一个卷积后同一张图两次推理结果不一致差异在1e-2量级部分图片的分类结果翻转。原因融合过程中1x1和Identity分支的权重分别转换成3x3卷积再加总。常见错误是先把多个分支权重相加再统一折算BN实际上必须先各自和对应BN融合再相加。因为每个分支的BN统计量不同先加后融合会拿混合统计量去归一化。解决严格按分支独立融合再对齐加到同一个卷积核上。正确顺序是3x3分支做“ConvBN”融合1x1分支先pad成3x3再做“ConvBN”融合Identity分支构造单位卷积再做“ConvBN”融合最后三个核相加。每一步结束后打印一次权重和手动检查形状。5.3 class.json的ID顺序和模型输出错位现象result.json里输出的ID和class.json里对不上比如class.json里“猫”是0“狗”是1但result.json的0号类别概率对应的图像明显像狗。原因训练时类别列表经过了排序但推理脚本直接用class_data.items()的原始顺序导致模型输出序号和类别名错位。解决推理脚本里也要用同一份sorted class_names不能用两套处理逻辑。最好把class_names序列化为一个固定txt文件训练和推理都从文件读。在资源项目里class.json已经存在但建议额外生成一个class_names.txt避免JSON格式不同造成的二次解析错误。5.4 使用DataParallel后加载权重报错现象训练时用了nn.DataParallel(model)保存state_dict后推理时直接load报“Missing key(s)”和“Unexpected key(s)”。原因state_dict里的参数key全部多了module.前缀裸模型没有这个前缀所以键对不上。解决加载时做一层清洗new_state_dict {k.replace(module., ): v for k, v in state_dict.items()}。如果你没有多卡训练需求建议训练时就别套DataParallel单卡逻辑更省心能少踩这个常见坑。5.5 部署态模型FPS没有明显提升现象简单地把模型model.eval()后去推理测速度和普通VGG差不多甚至更慢。原因model.eval()只影响BN和Dropout不会自动把三分支融合成一个Conv2d。RepVgg的加速必须显式执行结构重参数化重新生成一个单分支模型。如果只是切eval模型内部依旧有1x1和Identity路径的额外加法FPS自然上不去。解决显式遍历RepVggBlock调用融合函数把整个网络替换成纯3x3卷积栈。下一节会给出一个轻量实现直接把模型导出为TorchScript用到服务端。6. 进阶技巧把融合后的RepVgg导出为TorchScript6.1 手动触发融合并替换BlockRepVgg官方源码里通常会提供一个switch_to_deploy方法但如果你只是想在现有工程里快速验证可以自己写循环遍历模型所有子模块遇到RepVggBlock就融合。我一般这样处理def fuse_model(model): for name, module in list(model.named_children()): if isinstance(module, RepVggBlock): fused_conv module.fuse_to_conv() setattr(model, name, fused_conv) elif len(list(module.children())) 0: fuse_model(module)module.fuse_to_conv是从训练态Block里导出的一个独立卷积层输入输出一致。融合完成后整个模型结构变成一个普通VGG式卷积栈参数数量比训练态少30%左右推理时只走一个算子。6.2 导出TorchScript并验证融合后的模型可以直接导出TorchScript方便脱离Python环境推理model.eval() example_input torch.randn(1, 3, 224, 224) traced_model torch.jit.trace(model, example_input) traced_model.save(repvgg_deploy.pt) exported torch.jit.load(repvgg_deploy.pt) with torch.no_grad(): out1 model(example_input) out2 exported(example_input) print(最大误差, (out1 - out2).abs().max().item())torch.jit.trace要求输入尺寸固定如果你后续想支持动态尺寸记得在服务端用torch.jit.RecursiveScriptModule做二次包装。最后一步验证最大误差非常关键正常情况下误差应小于1e-5如果超过这个量级多半是融合代码里的BN折入顺序有误。从那以后我每次训练完都强制走一遍流程eval模式验证、融合、导出、对比误差。这套动作已经成了肌肉记忆替我挡掉了不少远程部署才暴露的翻车问题。希望这篇笔记帮到你。本文还有配套的精品资源点击获取
返回列表