GhostNet系列瓶子分类实战:从PyTorch复现到模型部署 简介基于PyTorch框架实现的GhostNet轻量级神经网络三个版本瓶子垃圾图像分类方案主要面向希望快速上手图像分类、迁移学习与轻量化模型训练的开发者及在校学生。包内包含可直接运行的训练和验证脚本以及清晰的说明文档用户按说明将自备数据整理至对应目录后即可训练无需额外修改代码。资源包共2000个文件核心为1991张瓶子垃圾图片另有6个Python脚本、数据标注JSON、类别文本与说明文档压缩后大小约39.53MB方便下载和迁移。训练脚本支持在GhostNet三个版本之间切换也可选择SGD或Adam优化器训练过程中会保存最优与最终权重并生成训练集和验证集的损失、准确率曲线以及日志验证脚本可基于测试集计算混淆矩阵、召回率、精确率和F1分数便于量化比较不同模型的分类效果。目前已有149人学习使用适合用于课程设计、毕业设计或轻量化垃圾分类实验的参考起点。1. 从瓶子回收线说起为什么选 GhostNetV1/V2/V3做瓶子回收分拣的视觉项目你很快会碰到一个矛盾生产线要求毫秒级响应但瓶子之间差异极大——透明 PET 瓶在强光下近乎隐形深色玻璃瓶几乎不反光标签磨损让颜色特征完全失真。GhostNet 系列恰好是这类场景的常用解法V1 用廉价线性变换压缩冗余特征V2 补上空间长距离注意力V3 从频率域引入高频卷积三者在 PyTorch 下都能直接复现并迁移预训练权重。这篇文章按结构拆解、数据集准备、训练调参、部署优化的顺序展开目标是把一个可运行的瓶子分类流程落到你自己的工程里。2. GhostNetV1/V2/V3 结构演进与 PyTorch 复现2.1 GhostNetV1用 Ghost 模块替换普通卷积GhostNetV1 的核心观察是普通卷积输出的特征图中存在大量近似冗余的通道这些通道不必用完整卷积逐一生成。Ghost 模块先用一个普通卷积生成约一半的通道再对这个结果施加 depthwise 卷积得到另一半幻影通道最后在通道维度拼接。class GhostModuleV1(nn.Module): def __init__(self, inp, oup, dw_size3, ratio2, stride1): super().__init__() init_c oup // ratio new_c oup - init_c self.primary nn.Sequential( nn.Conv2d(inp, init_c, 1, stride, 0, biasFalse), nn.BatchNorm2d(init_c), nn.ReLU(inplaceTrue)) self.cheap nn.Sequential( nn.Conv2d(init_c, new_c, dw_size, 1, dw_size // 2, groupsinit_c, biasFalse), nn.BatchNorm2d(new_c), nn.ReLU(inplaceTrue)) def forward(self, x): x1 self.primary(x) x2 self.cheap(x1) return torch.cat([x1, x2], dim1)ratio2 表示输出通道中一半来自主卷积、一半来自派生变换。cheap 分支用 groupsinit_c 的 depthwise 卷积参数开销约等于一个 3x3 卷积核对单通道做变换远低于完整的 1x1 卷积投影。stride 只作用于主卷积cheap 分支的空间分辨率跟随主卷积拼接时无需额外对齐。Ghost Bottleneck 的堆叠方式是1x1 升维 → depthwise 空间变换 → Ghost 模块压回原通道stride1 时加残差。合成 Bottleneck 时注意 width_multiplier 会线性缩放每个 stage 的通道数1.0x 版本最终分类头输入是 1280 维。2.2 GhostNetV2DFC 注意力解决长距离依赖V2 要解决的问题是Ghost 模块只处理了通道冗余3x3 卷积在浅层感受野不到 20 像素而瓶子标签区域的宽度可能占特征图四分之一。V2 引入的解耦全连接注意力DFC把全局空间注意力分解为水平卷积和垂直卷积的组合。class DFC(nn.Module): def __init__(self, channels, k_size5): super().__init__() self.h nn.Conv2d(channels, channels, (1, k_size), padding(0, k_size // 2), groupschannels) self.v nn.Conv2d(channels, channels, (k_size, 1), padding(k_size // 2, 0), groupschannels) def forward(self, x): attn torch.sigmoid(self.v(self.h(x))) return x * attn两个一维卷积的复杂度之和是 O(HW(HK))远小于全连接注意力的 O(HW·HW)。k_size 一般取 5 或 7在瓶子分类实验中两者差异很小。注意 DFC 的卷积都做了 groupschannels逐通道建模空间关系不做跨通道融合——跨通道信息交给后面的 Ghost 模块处理。插入位置选在 Ghost 模块之后、残差相加之前。2.3 GhostNetV3HFConvolution 从频率域补充信息GhostNetV3 的改进视角是频率域。普通卷积天然倾向提取低频分量平滑区域对高频边缘、纹理响应偏弱而这些恰好是区分瓶口螺纹、标签边界的关键线索。HFConvolution 在主卷积旁并联一个固定拉普拉斯核的高通分支把提取到的高频响应按系数叠加回主输出。class HFConv(nn.Module): def __init__(self, channels): super().__init__() self.dw nn.Conv2d(channels, channels, 3, 1, 1, groupschannels, biasFalse) lap torch.tensor([[0, -1, 0], [-1, 4, -1], [0, -1, 0]], dtypetorch.float32).reshape(1, 1, 3, 3) self.register_buffer(hf, lap) def forward(self, x): return self.dw(x) 0.1 * F.conv2d(x, self.hf, padding1)实际工程里我一般把融合系数设为 0.1~0.15。系数太大高频噪声会被放大太小又起不到作用。拉普拉斯核保持固定不训练避免了随机初始化导致训练早期梯度剧烈波动。下表汇总三代模型的核心差异版本核心模块解决的问题典型用例V1Ghost Module通道冗余、FLOPs 过高资源受限的实时分类V2DFC Attention空间长距离依赖大目标、高分辨率输入V3HFConvolution高频细节丢失细粒度、边缘敏感任务3. 瓶子垃圾数据集的构建与 DataLoader 配置3.1 数据来源与目录组织做一个可用的瓶子分类数据集常见做法是公开数据集打底加自采数据微调。公开的垃圾分类数据集如 TrashNet包含玻璃瓶、塑料瓶类别可以作为预训练或对比实验的基准。自采数据时注意覆盖品牌、颜色、光照和磨损程度四个维度。每种瓶子建议至少拍 300 张用手机或工业相机都可以但分辨率统一缩放到 256x256 以上再进网络。目录按 ImageFolder 约定组织data/bottles/ ├── train/ │ ├── pet_bottle/ │ ├── glass_bottle/ │ └── aluminum_can/ ├── val/ └── test/3.2 Dataset 与 DataLoader 参数设置from torchvision import datasets, transforms from torch.utils.data import DataLoader train_loader DataLoader( datasets.ImageFolder(data/bottles/train, transformtrain_transform), batch_size64, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader( datasets.ImageFolder(data/bottles/val, transformeval_transform), batch_size64, shuffleFalse, num_workers4, pin_memoryTrue)num_workers4 在大多数四核 CPU 上刚好开太多反而会因为进程调度消耗拉高延迟。pin_memory 只在 GPU 训练时有意义它的作用是锁定页内存在 H2D 拷贝时减少一次 CPU 到 GPU 的中转。Windows 下多进程 DataLoader 务必把创建逻辑放在if __name__ __main__保护块里否则会无限递归产生子进程。3.2.1 不平衡类别与采样器实际回收线上 PET 瓶占比可能超过 60%此时直接用原始分布训练会让模型偏向高频类别。用 WeightedRandomSampler 做样本级过采样from torch.utils.data import WeightedRandomSampler labels [s[1] for s in dataset.samples] class_count torch.bincount(torch.tensor(labels)) sample_weights 1.0 / class_count[labels].double() sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue)注意 sample_weights 是逐样本的权重不是逐类别的。做法是把每个样本的标签映射到 1/class_freq。replacementTrue 时每个 epoch 会采样到若干重复的少类别样本。这会改变实际迭代步数如果用 StepLR 等基于 epoch 的调度器不用改但基于 iteration 的调度器需要重新计算。3.3 数据增强针对瓶子特征做组合瓶子的核心干扰因素是反光、标签遮挡和颜色失真。增强组合里保留随机裁切和翻转之外我一般加 RandomRotation(15) 模拟传送带抖动加 ColorJitter 模拟不同光源。还有一种值得加的是 RandomErasing用来模拟标签破损。train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.2), transforms.RandomErasing(p0.25, scale(0.02, 0.15)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ])验证集只做 Resize CenterCrop不做任何随机变换。Resize 到 256 再裁剪到 224 与 ImageNet 预训练一致避免预训练权重对尺寸的期望被破坏。RandomErasing 的 scale 上限不要超过 0.2面积太大容易把整瓶盖住训练损失会异常波动。ColorJitter 在瓶子分类里是必要项透明瓶在不同角度下的反光区域颜色变化明显模型不应依赖颜色做唯一判断依据。4. 训练与调参PyTorch 中的 GhostNet 迁移学习4.1 加载预训练权重并替换分类头GhostNet 在 ImageNet 的预训练权重可以直接加载分类头改成自己的类别数。关键是加载时过滤掉最后一层的键。from ghostnet_pytorch import GhostNetV1 model GhostNetV1(num_classes3, width_multiplier1.0) ckpt torch.load(ghostnetv1_1.0x.pth.tar, map_locationcpu) state_dict ckpt.get(state_dict, ckpt) filtered {k: v for k, v in state_dict.items() if fc not in k and classifier not in k} model.load_state_dict(filtered, strictFalse)strictFalse 允许缺失最后一层的权重。如果不加过滤条件直接加载会因 shape 不匹配抛出 RuntimeError。若数据集不到每类 100 张冻结前 4 个 stage 只训练最后两个 stage 和分类头会明显降低过拟合风险。4.2 优化器、标签平滑和 weight decay交叉熵在瓶类模糊样本上会输出过高置信度标签平滑可以缓解这个问题。同时对 BatchNorm 参数不做 weight decay这是训练稳定性上容易忽略的点。criterion nn.CrossEntropyLoss(label_smoothing0.1) decay [p for n, p in model.named_parameters() if bn not in n] no_decay [p for n, p in model.named_parameters() if bn in n] optimizer torch.optim.AdamW([ {params: decay, weight_decay: 1e-4}, {params: no_decay, weight_decay: 0.0}, ], lr1e-3)label_smoothing0.1 的含义是真实标签的概率不是 1.0而是 0.9剩余 0.1 分配给其他类别。对半透明瓶这种容易混淆的样本平滑后模型预测分布会更保守验证集准确率往往更高。BN 层的 gamma 和 beta 本身是归一化参数加 weight decay 会让它们偏离 1 和 0常见训练中 BN 收敛变慢。4.3 余弦退火与早停用 CosineAnnealingLR 配合早停是轻量分类任务的稳定搭配。T_max 设成规划的总 epoch 数学习率从 1e-3 平滑降到 1e-6。from torch.optim.lr_scheduler import CosineAnnealingLR scheduler CosineAnnealingLR(optimizer, T_max50, eta_min1e-6) best_acc, patience_counter 0.0, 0 for epoch in range(50): model.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() loss criterion(model(images), labels) loss.backward() optimizer.step() scheduler.step() val_acc evaluate(model, val_loader) if val_acc best_acc: best_acc val_acc patience_counter 0 torch.save(model.state_dict(), ghostnet_best.pth) else: patience_counter 1 if patience_counter 8: break提示loss 在 epoch 0 之后不下降先检查预处理是否和预训练一致再检查 BN 是否处于 eval 模式。早停的 patience 设 8 比较合适。瓶子分类的验证集如果只有几百张准确率本身有 ±2% 的随机波动patience 太小会在一个正常的波动低谷里提前停止。保存 best 模型而不是最后一个 epoch 的模型这个习惯在训练后期尤其重要。5. 量化感知训练与 ONNX 导出的部署验证5.1 QAT 微调被忽视的细节GhostNet 家族的体积优势意味着即使不量化也能跑但部署到 RK3588、Jetson 这类带 INT8 加速单元的板卡时量化感知训练是准确率不掉坑的前提。QAT 的常规流程是prepare_qat → 低学习率微调 5~10 个 epoch → convert 导出。微调学习率建议设正常训练的 20%~30%迭代次数不需要多重点是让 BN 统计量重新适配量化噪声。model.qconfig torch.ao.quantization.get_default_qat_qconfig(fbgemm) qat_model torch.ao.quantization.prepare_qat(model.train(), inplaceFalse) # 低学习率微调 qat_model torch.ao.quantization.convert(qat_model.eval(), inplaceFalse)量化掉点超过两个百分点时优先检查第一层卷积和最后的全连接层这两个层对量化误差最敏感。一般会将这两层的 qconfig 手动保持为 FP32。5.2 ONNX 导出后的数值一致性检查ONNX 导出是跨端部署的中转环节导出完必须做一致性验证。值得留意的是量化模型导出时要用 QDQ 格式保留量化节点导出后对比 PyTorch 和 onnxruntime 的 softmax 输出误差在 1e-4 以内算合格。import onnxruntime as ort sess ort.InferenceSession(ghostnet_bottles.onnx, providers[CPUExecutionProvider]) ort_out sess.run(None, {input: dummy.numpy()})[0] torch_out torch.softmax(model(dummy), dim1).detach().numpy() max_err torch.abs(torch.tensor(ort_out) - torch.tensor(torch_out)).max().item() print(fmax abs error: {max_err:.6f})除了数值对比留一张训练集外的真实瓶子照片做端到端回归确认前处理参数——尤其是 Normalize 的均值和标准差——在部署代码和训练代码中完全一致。这类问题在生产环境里出现频率远高于模型本身的精度问题往往一次 dirty 测试就能暴露。本文还有配套的精品资源点击获取