去年下半年我接了一个挺典型的现场需求:客户产线某个部件的外观检测,正常样本只有几十张,缺陷倒是五花八门——划痕、压伤、脏污,还有叫不上名字的异物。要求也非常明确:漏检尽量低,误报率控制在产线能接受的范围,而且必须部署在工控机上,不能用大服务器慢慢跑。
这个需求几乎是工业异常检测的标准模板。现在缺陷样本稀缺、缺陷形态开放、边缘端算力有限,这三个约束凑在一起,恰好把 Transformer 这类擅长全局建模的模型推到了台前,又逼着我把“能跑”变成“跑得稳”。我后来把这套从数据准备、模型训练到边缘部署的完整流程整理成了一个项目,代号叫 Transformer2edge。这篇文章就把整个链路摊开讲一遍,适合正在做工业视觉落地、或者想把 Transformer 架构搬到边缘设备上的朋友参考。
1. 先想清楚:异常检测到底在解决什么问题
1.1 工业异常检测和普通分类任务完全是两码事
分类任务的前提是有明确的类别体系:猫、狗、车、行人,样本天然是平衡的,模型要做的只是学一个判别边界。工业异常检测完全不同,它本质上是“开放集”问题——你永远不知道下一块不良品上会长出什么样的缺陷。
以我当时接的案例来说,客户给的数据里只有 68 张正常品图片,缺陷图总共 40 多张,而且划痕、压伤、脏污、异物这几种形态之间差异极大。如果用传统分类思路去做,要么缺陷类别不够用,要么过拟合到那 40 多张图上,产线一换光照条件就崩。
异常检测的底层逻辑是另一套:把“正常”的分布学好,凡是偏离这个分布的都是异常。这个概念说起来简单,做起来要命。因为“正常”本身也有波动——光照变一点、角度偏一点、机台抖动一下,都算正常。模型得太懂什么是合理的波动范围,才能把真正的不正常挑出来。
1.2 为什么选 Transformer,而不是继续用 CNN
传统 CNN 做异常检测也不是不行,PatchCore 这类基于特征存储库的方法在 MVTec 数据集上效果一直不错。但 CNN 的感受野受限,对“大范围上下文”的敏感度不够。工业缺陷有一个特点:很多缺陷本身很小,但它的“异常感”恰恰来自与周围环境的对比。一个细微的划痕,单独裁剪出来看可能毫不起眼,但放在整个部件的纹理背景下,它就是不正常的。
Transformer 的自注意力机制天然擅长捕捉这种长程依赖。图像被打成 patch 序列之后,每个 patch 都能直接和其他任意位置的 patch 做交互,相当于模型看每个局部特征时,都会自动参考全局背景。这个能力在纹理复杂的工件表面特别有用。
不过我得说句公道话,别神话 Transformer。小数据集上它非常容易过拟合,训练不稳定也是家常便饭。后面会详细讲怎么用蒸馏、轻量化设计和数据增强把这些坑填上。
1.3 三种主流方案,我为什么这么选
工业异常检测目前有三条主流技术路线,对比起来看会更清楚:
| 方案 | 核心思路 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|---|
| 基于重构 | 用 AE/VAE/ViT 学习重建正常样本,重建误差大的区域即为异常 | 原理直观,训练简单 | 对纹理细节不敏感,容易把正常纹理变化也重建得很好 | 纹理简单、缺陷明显的场景 |
| 基于嵌入 | 用预训练网络提取特征,建立正常特征存储库,推理时计算最近邻距离 | 精度高,泛化性较好 | 特征库占内存,推理耗时长,边缘端有压力 | 缺陷形态多样、对精度要求高 |
| 基于合成异常 | 人为生成异常样本,训练模型区分正常和合成异常 | 可控性强,可针对具体缺陷类型 | 合成分布与真实分布存在偏差,可能过拟合到合成特征 | 已知缺陷类型较明确 |
我做 Transformer2edge 时,最终选了“嵌入 + 合成异常”结合的路线,但在特征提取部分用轻量化 Transformer 替代了传统的 ResNet 主干。原因有两个:一是嵌入式的精度上限更高,二是 Transformer 的全局建模能力能缓解特征存储库方法在复杂背景下的误判。代价是模型结构需要重新设计,这也是 Transformer2edge 名字的由来——把 Transformer 的能力真正用到边缘设备上。
2. 数据准备与合成异常:没有缺陷图也能训练
2.1 数据采集的四个硬性要求
工业项目的第一关永远是数据。很多团队上来就想要更多的缺陷样本,但实际上正常样本才是最关键的。我在这个项目里总结了四个硬性要求:
第一,正常样本必须覆盖真实产线的全部变化范围。不能只在固定光照、固定角度下采集,要分别抽不同班次、不同光照、不同机台状态下的正常品,否则模型会把光照变化当成异常。
第二,正常样本数量不少于 50 张,200 张左右效果就比较稳定了。少于 50 张,特征空间根本建立不起来。我这里最终用了 68 张,说实话偏少,后面靠合成策略和更强的正则化硬扛过来了。
第三,拍摄时工件在视野内的位置、角度尽量一致。如果位置飘忽很大,模型会花大量容量去学“位置不变性”,而这些能力对缺陷检测本身没有帮助。
第四,如果有缺陷样本,哪怕只有十几张,也要保留下来。它们不参与训练,但可以作为验证集,用来观察模型能不能真正把已知缺陷检出来。
2.2 合成异常的正确打开方式
缺陷样本不够,最常用的补救手段就是合成异常。业界两个经典做法是 CutPaste 和 DRAEM。CutPaste 的思路是把正常图像中的一块区域随机裁剪后粘贴回图像的其他位置,模拟局部纹理突变;DRAEM 更进一步,用 Perlin 噪声生成不规则的缺陷掩码,再在掩码区域叠加随机纹理。
我实际跑下来的感受是:CutPaste 简单有效,但对“结构性异常”的模拟能力弱;DRAEM 的效果更接近真实缺陷,尤其是压伤、磕碰这类不规则形态。Transformer2edge 里我两种都在用,但有个重要的原则问题——合成异常不能太“认真”。
什么意思?如果你合成异常时纹理选择太固定,颜色太统一,模型学到的是“这种特定纹理不正常”,而不是“偏离正常纹理分布不正常”。产线上的真实缺陷千变万化,一旦和合成纹理对不上,漏检率就飙升。所以我的建议是:合成时故意做得粗糙一点、随机一点,宁可是“看起来奇怪”而不是“看起来像某种缺陷”。
还有一个特别坑的细节:合成异常和真实正常样本的比例要控制好。比例太高,模型会退化成一个普通二分类器,对未知类型的异常几乎没有泛化能力;比例太低,模型又学不到异常信号。我试过 1:1 到 1:5,最终稳定在 1:3 左右,正常样本占多数,异常作为少数信号去引导特征空间的结构。
2.3 数据增强里的三个教训
除了合成异常,常规数据增强也有一堆坑。我踩过最狠的三个:
一是随机裁剪要慎用。工业部件位置相对固定,过强的随机裁剪会让模型误以为“位置偏移是正常的”,推理时部件位置稍微偏一点,模型反而就不报异常了。我最后只用了中心裁剪加轻微平移,幅度控制在 5% 以内。
二是颜色增强要保守。HSV 抖动、随机亮度这些对自然图像很有效,但在工业场景下,光照变化本来就是异常的重要信号。如果增强做得太狠,模型对真实的光照异常会非常迟钝。我把饱和度抖动和亮度抖动的幅度都降到了默认值的一半以下。
三是 MixUp 和 CutMix 这类混合增强,在异常检测里效果不稳定。我的经验是它们在分类任务里能提升鲁棒性,但在异常检测里会让正常特征空间变模糊,反而增加误报率。Transformer2edge 的最终方案里没有用这两种增强,这是反复对比实验之后得出的结论。
3. 轻量化 Transformer 模型设计与训练关键细节
3.1 模型结构怎么改,才适合边缘设备
边缘设备上跑 Transformer,最大的敌人是参数量、计算量和内存占用。ViT-Base 有 8600 万参数,边缘端跑 224x224 输入可能只有十几 FPS,根本没戏。我在设计 Transformer2edge 的骨干网络时,做了三处关键改动:
第一,patch embedding 不走标准卷积,而是用步长为 2 的 3x3 卷积堆叠三次,逐步把分辨率从 256 降下来。这样做的好处是早期就能保留更多局部纹理特征,同时减少后续自注意力层的序列长度,计算量能省下不少。
第二,自注意力层只保留 4 层,每层维度降到 192,head 数量设 4。这个规模对工业缺陷检测来说是够用的,因为我们要的不是 ImageNet 级别的语义理解能力,而是对正常纹理分布的敏感度。dim 太大在数据量少的时候反而容易过拟合。我试过 dim=384 的版本,训练集 Loss 降得很快,但验证集上的 AUC 反而低了两个点。
第三,前馈网络 FFN 部分不直接上 GELU 激活函数。GELU 在部分 NPU 上效率很差,量化时也容易出精度问题。我用的是经过重参数化设计的 ReLU 结构,推理阶段可以融合成单个矩阵乘法,对 TensorRT 和 RKNN 这类推理框架都很友好。实测下来,这个改动在边缘端能带来约 15% 的延迟收益。
下面是我最终落地的一组配置,可以照着直接用:
| 参数 | 数值 | 说明 |
|---|---|---|
| 输入分辨率 | 256x256 | 平衡细节和计算量 |
| Patch size | 16x16 | 每个 patch 包含足够纹理 |
| Embedding dim | 192 | 降低过拟合风险 |
| Transformer 层数 | 4 | 够用即可 |
| Head 数 | 4 | 与 dim 匹配 |
| FFN 隐藏层 | 384 | 2 倍 dim |
| Dropout | 0.1 | 稳定性关键 |
3.2 训练策略:对比学习损失是核心
有了模型结构,训练策略也要跟得上。Transformer2edge 的核心训练目标不是分类损失,而是对比学习损失。具体来说,我把正常样本做了两次不同的随机增强,得到两个视角的表示,然后让模型学习拉近“同一张图的不同视角”,同时推开“不同正常图之间的表示距离”。
同样一张图的不同增强版本在特征空间里应该挨得很近,不同正常图之间允许有一定距离,但也不能分布得太开。我用的是 InfoNCE 损失的变体,温度参数设了 0.07。训练到后期,正常样本在特征空间里会形成一个紧凑的簇,推理时新来一张图,如果它的特征表示距离这个簇的中心太远,就判定为异常。
优化器选了 AdamW,学习率设置为 1e-3,warmup 10 个 epoch 后按 cosine schedule 衰减到 1e-5。EMA 权重衰减设 0.999,这个对稳定训练很有帮助,尤其是数据量不到 100 张的时候,EMA 简直是我最后的救命稻草,它能显著降低训练后期的 Loss 震荡。
还有一个细节值得说:训练时的 batch size 不需要大。我最后用的是 16,每个 batch 里恰好能放下 8 组正负样本对比。batch size 太大反而会让对比学习任务变简单,模型会走捷径去学“区分不同图片本身”,而不是学“区分正常纹理的变化范围”。
3.3 蒸馏:大模型带小模型,边缘端照样吃香
轻量化 Transformer 在边缘端跑是没问题了,但从零开始训练小模型的上限往往有限。一个更聪明的做法是:先用一个更强的教师模型(可以是 ViT-Base,也可以是大规模预训练的 CNN)在同样的数据上把特征空间建好,然后让学生模型去对齐教师模型的输出特征。
我这里具体怎么做呢?教师模型用 CLIP 的 ViT-B/16 和 ResNet-50 双塔结构分别提取特征,然后做特征拼接,得到教师表征。学生模型就是前面说的轻量化 Transformer,训练时的 Loss 由两部分组成:一是学生和教师特征的余弦相似度,二是异常检测任务本身的对比损失。两项之间权重调到 0.6 和 0.4,蒸馏为主,任务为辅。
使用蒸馏之后,学生模型在验证集上的 AUC 从 0.93 提到了 0.97,而推理速度几乎没变。这算是我在整个项目里投入产出比最高的一步。
4. 边缘部署实操:模型转换、量化与推理优化
4.1 边缘硬件怎么选
模型训练好了,下一步是落到边缘设备上。Transformer2edge 这个项目的硬件选型上,我前后试过四类设备,各有利弊:
| 硬件平台 | 算力 | 显存/内存 | 功耗 | 适合场景 |
|---|---|---|---|---|
| 工控机 + Intel/AMD 核显 | 中等 | 共享内存 | 高 | 已有产线改造,兼容性好 |
| Jetson Orin Nano/NX | 高 | 8-16GB 共享 | 中 | 独立视觉检测单元 |
| 瑞芯微 RK3588 | 中上 | NPU 6 TOPS | 低 | 对成本敏感的嵌入式场景 |
| 算能 BM1684X | 高 | 32 TOPS | 中 | 多路并发检测 |
我这个项目最终选了 Jetson Orin NX,原因是客户现场同时要跑 4 路相机,每路画面都要实时推理,Orin NX 的 16GB 统一内存和 100 TOPS 算力比较充足,而且 TensorRT 对 Transformer 的支持比较成熟。如果是单路检测或者预算紧张的场景,RK3588 也是不错的选择,但要注意它对某些算子的支持不如 NVIDIA 平台顺滑。
在选型时还有一个容易被忽视的点:看产线的实际安装环境。有些工控机放在电柜里,散热差,高功耗设备容易降频导致推理速度不稳定。Orin NX 功耗控制在 15-25W,装上被动散热片在电柜里也不会热降频,这是我选它的一个现实原因。
4.2 模型导出和量化:ONNX 转 TensorRT 的完整流程
PyTorch 模型要部署到 Jetson 上,标准路径是 PyTorch -> ONNX -> TensorRT。整个过程踩坑不断,我重点说几个关键环节。
首先是 ONNX 导出。导出前模型要切换到 eval 模式,关闭 dropout 和 EMA。输入输出建议固定分辨率,边缘端尽量不要用动态输入尺寸,因为动态维度在 TensorRT 里会触发多组 kernel 优化,不仅增加转换时间,推理时还会有额外的调度开销。我这边固定为 256x256,输出一个 1x256 维的特征向量和异常得分,导出时用 opset_version=17 比较稳妥。
导出后先用 onnxsim 做一遍图优化,把多余的 Shape 和 Reshape 节点清掉。之前我有一版模型导出后有 200 多个无效节点,转 TensorRT 时直接卡死,用 onnxsim 优化后降到 60 多个,问题瞬间解决。
然后是 TensorRT 转换,这里 FP16 和 INT8 要分开说。
FP16 基本是白捡的收益,精度损失可以忽略不计,推理速度能翻倍。我在 Transformer2edge 里默认就用 FP16。
INT8 更快,但坑也多。TensorRT 的 INT8 需要提供校准数据集,校准数据的选择直接影响量化精度。我的经验是:校准集必须用真实产线的正常样本,数量 100-200 张,覆盖不同光照、不同角度。千万不能用训练集里的图直接当校准集,否则模型在校准过的分布上表现很好,一到现场分布稍微偏移就全线崩溃。
| 精度模式 | 推理延迟(4路并发) | 精度损失 | 推荐场景 |
|---|---|---|---|
| FP32 | 约 22ms/张 | 基准 | 调试阶段 |
| FP16 | 约 11ms/张 | 几乎无损 | 大多数场景 |
| INT8 | 约 6ms/张 | 需验证 | 高并发、低延迟 |
INT8 校准不一定要一上来就全量化。我建议先用 TensorRT 的逐层敏感度分析找出对量化最敏感的几个层,把这些层回退到 FP16,其余层用 INT8。这个流程我现在已经做成脚本了,跑一遍大概十几分钟,能省下现场调试的好几个小时。
4.3 推理端的关键代码与内存优化
TensorRT 引擎加载后,Python 侧核心代码如下:
import tensorrt as trt import numpy as np class TrtInference: def __init__(self, engine_path): logger = trt.Logger(trt.Logger.WARNING) runtime = trt.Runtime(logger) with open(engine_path, "rb") as f: engine = runtime.deserialize_cuda_engine(f.read()) self.context = engine.create_execution_context() self.inputs = [] self.outputs = [] self.allocations = [] for i in range(engine.num_io_tensors): name = engine.get_tensor_name(i) mode = engine.get_tensor_mode(name) shape = engine.get_tensor_shape(name) dtype = trt.nptype(engine.get_tensor_dtype(name)) if mode == trt.TensorIOMode.INPUT: self.context.set_input_shape(name, shape) self.inputs.append(name) else: self.outputs.append(name) buf = np.empty(shape, dtype=dtype) self.allocations.append(buf) def infer(self, img): # 假设 img 已经是 (1,3,256,256) 且归一化 self.allocations[0][:] = img self.context.execute_v2(self.allocations) return self.allocations[1].copy()这里有几个细节要强调。execute_v2 里的 allocations 必须用 numpy 数组的底层指针,不能传 Python list。另外,单帧推理模式下输入输出分开分配内存没有问题,但如果做多路相机的并发推理,显存会吃紧,这时候有两个优化手段:一是所有输入共享同一块 GPU 显存,轮流写入;二是把图像预处理放到 GPU 上做,不要 CPU 转 BGR 再转 GPU,直接 device 端用 cv2.cuda 完成 resize 和 normalize。
我这边四路并发最终用了双缓冲结构:GPU 上一块缓冲用于预处理和推理,另一块用于下一帧图像的拷贝。这样 GPU 在推理的同时不需要等待 CPU 完成图像读取,延迟从 22ms 降到了 16ms,吞吐量提升非常明显。
4.4 多路相机并发与 CPU 流水线调度
多路并发还有一个常见误区:不要为每一路相机单独创建一个 TensorRT context。我的经验是,在 Jetson Orin NX 上最多开 2 个 context,再往上会触发 GPU 显存碎片化和上下文切换开销,推理总吞吐量反而下降。正确的做法是:所有相机推图到一个共享的任务队列,一个个按顺序进 GPU 推理,4 路 1080p 画面在 20ms 内的延迟完全够用,不需要真并行。
CPU 侧的流水线调度同样重要。我的做法是三个线程并行:线程 A 负责从相机拉流和 ROI 提取;线程 B 负责预处理和推理;线程 C 负责把结果推给 PLC 或上位机。三线程之间用队列解耦,队列长度控制在 5 帧以内,慢了就丢帧,避免内存持续堆积。
5. 从仿真到产线:精度回退、温度漂移与稳定性问题排查
5.1 常见问题速查表
部署上线只是开始,真正的麻烦都在现场。我整理了一张高频问题排查表,基本涵盖了边缘端部署异常检测算法时遇到的主要坑:
| 现象 | 可能原因 | 解决思路 |
|---|---|---|
| FP16/INT8 后误报率明显升高 | 量化敏感层未保留 FP16 | 逐层敏感度分析,关键层回退 FP16 |
| 特征分布与训练时不一致 | 校准集与现场分布偏移 | 用现场真实数据重新校准 |
| 白天正常、傍晚误报飙升 | 光照变化超出建模范围 | 增加多光照样本,或增加白平衡预处理 |
| 机台温度升高后推理变慢 | GPU 热降频 | 功耗模式限制,或加强散热 |
| 长时间运行内存不断上涨 | 队列堆积或显存泄漏 | 检查队列长度,显存显式释放 |
| 检测到缺陷但位置对不上 | ROI 坐标换算错误 | 在输送给 PLC 前加坐标系映射验证 |
5.2 精度回退的排查流程
模型从 PyTorch 转到 TensorRT 后,如果发现精度明显下降,我一般会按下面的顺序排查:
第一步,确认 FP32 引擎的精度和 PyTorch 原模型一致。这一步能排除算子实现差异,比如 LayerNorm 在 TensorRT 里的 epsilon 实现可能和 PyTorch 不完全一样。
第二步,对比 FP16 和 FP32 的逐层输出差异。TensorRT 可以开启层级别的调试输出,找到输出差异最大的几层。一般来说,注意力层里的 Softmax 和最后的 Embedding 层是量化重灾区,把这几层手动设为 FP16 或 FP32,其他保持 INT8,能恢复大部分精度损失。
第三步,检查输入预处理是否完全对齐。PyTorch 训练时用的归一化参数(mean、std)必须一字不差地搬到部署端。这个听起来简单,但我在现场遇到过两次因为 BGR/RGB 通道顺序不一致导致的精度崩盘,排查了整整一个下午。
5.3 温度漂移:INT8 量化的隐形杀手
温度对 INT8 推理精度的影响,行业内讨论得不多,但确实存在。卷积和矩阵乘法在 INT8 下对数值范围非常敏感,当设备温度升高,某些硬件单元的执行精度可能会有细微波动,导致同一张图片在冷机时正常、热机时报异常的现象。
我的解决办法有两层:第一层是硬件层面,给 Jetson 板子加了工业级散热片和风扇,设置功耗模式为 20W,宁可牺牲一点峰值性能,也要保证持续推理的稳定性。第二层是软件层面,在推理服务里加了一个温度监控模块,当核心温度超过 75 度时自动切换更保守的阈值参数,异常判定阈值向上浮动 5%,避免温度引起的特征漂移直接触发误报。
5.4 长尾缺陷漏检与规则后处理兜底
深度学习模型再强,也不可能覆盖所有类型的缺陷。Transformer2edge 上线后,客户反馈有一类“浅划痕”经常漏检——它在特征空间里偏离正常分布不够远,阈值稍微提高一点就藏进去了。
我的做法是加一道传统视觉规则后处理做兜底。针对浅划痕的特点,用方向性滤波增强,再配合梯度阈值和连通域分析做一个轻量级的“浅划痕探测器”。两个检测器是“或”的关系:深度学习检出来或者规则探到都算异常。这不算什么高深技术,但在产线落地时特别实用,客户只看最终漏检率和误报率,不会关心你用的是深度学习还是传统视觉。
规则后处理还有另一个好处:可以快速响应新增缺陷类型。当客户反馈某类缺陷漏检时,传统规则可以在当天就加一道检测逻辑,而做模型重训需要收集数据、标注、训练、验证,至少三五天。在项目初期,先快后准的节奏非常重要。
5.5 持续迭代:数据回流与周期性重训
边缘部署不是终点,而是另一个起点。Transformer2edge 的数据闭环长这样:前端推理服务把每张检测图的特征向量和异常得分异步记录到本地数据库;每周导出一份“高置信正常”和“高置信异常”的样本集合,人工复核后回流到训练集;每两周用积累的数据增量微调一次模型,产线切换产品型号时重新做一次快速验证。
增量训练时要特别小心灾难性遗忘的问题。我用的策略是:把过去三周内的典型正常样本固定作为锚点,和新增样本混在一起训练,同时用学习率 5e-5 这样的低学习率做微调,最大程度保留原有特征空间的结构。目前这套机制在产线上稳定运行了几个月,误报率从初期的每天 8 次降到了每天 1-2 次,漏检率也符合客户预期。
6. 写在最后的工程经验
Transformer2edge 这个项目做下来,我最大的体会是:模型结构是重要,但它只是整套系统的一小部分。数据策略、训练策略、部署优化、现场调试,每一环抠出来的时间都比想象中多。边缘端的 AI 落地,考验的是把算法、硬件、生产环境串起来的系统工程能力。
如果你也打算在边缘设备上做 Transformer 类模型,有几点建议可以提前排雷:第一,别一上来就追求最复杂的模型,先打通端到端流程,用最简单的 ViT-Tiny 跑通,再去迭代精度;第二,把量化敏感度分析和校准脚本提前固化,现场调试时能省一半时间;第三,部署端一定要加温度监控和日志回传,没有数据闭环,后续优化就是盲人摸象。
最后分享一个小技巧:在产线试运行阶段,把模型的异常得分做实时可视化,保存成热力图叠加在原图上,这比任何指标报表都更能说服现场工程师。他们看到热力图准确圈出缺陷位置的那一刻,才算真正信任这套算法。这个信任建立的瞬间,往往比模型 AUC 提升零点几个点更有价值。