第26课:工业零部件外观缺陷检测系统:从学术Demo到产线工程的重构实战

文章目录

    • 一、先说结论:为什么你的分类模型上了产线就废了
    • 二、改造前后:一张表看清差距
    • 三、架构重构:五层分层设计
    • 四、数据层:把"脏"当成常态
      • [4.1 清单驱动,而不是靠目录结构猜标签](#4.1 清单驱动,而不是靠目录结构猜标签)
      • [4.2 批次分组划分:这 10 行代码能救你的验证指标](#4.2 批次分组划分:这 10 行代码能救你的验证指标)
      • [4.3 数据体检与训练期兜底](#4.3 数据体检与训练期兜底)
      • [4.4 增强:模拟的是产线扰动,不是花哨](#4.4 增强:模拟的是产线扰动,不是花哨)
      • [4.5 类别不平衡:双管齐下](#4.5 类别不平衡:双管齐下)
    • 五、模型层:小样本友好的骨干工厂
    • 六、训练引擎:稳定性优先于花哨
      • [6.1 优化配置](#6.1 优化配置)
      • [6.2 学习率调度:预热是必需的](#6.2 学习率调度:预热是必需的)
      • [6.3 早停:统一转成"越大越好"](#6.3 早停:统一转成"越大越好")
      • [6.4 产物契约:权重与元信息分离](#6.4 产物契约:权重与元信息分离)
    • 七、决策层:这才是整个系统的核心
      • [7.1 三态判定,而不是二分类](#7.1 三态判定,而不是二分类)
      • [7.2 判定规则:只有"足够像合格"才放行](#7.2 判定规则:只有"足够像合格"才放行)
      • [7.3 业务指标:这三个才是老板要看的](#7.3 业务指标:这三个才是老板要看的)
      • [7.4 阈值寻优:把风险权衡交给算法](#7.4 阈值寻优:把风险权衡交给算法)
      • [7.5 配置方式](#7.5 配置方式)
    • 八、评估与报告:给质量负责人看得懂的东西
      • [8.1 报告字段](#8.1 报告字段)
      • [8.2 报告产物的工程细节](#8.2 报告产物的工程细节)
    • [九、部署:ONNX + FastAPI](#九、部署:ONNX + FastAPI)
      • [9.1 ONNX 导出与一致性校验](#9.1 ONNX 导出与一致性校验)
      • [9.2 HTTP 接口](#9.2 HTTP 接口)
      • [9.3 判定结果示例](#9.3 判定结果示例)
      • [9.4 双后端一致性设计](#9.4 双后端一致性设计)
    • [十、实战:5 分钟跑通全流程](#十、实战:5 分钟跑通全流程)
    • 十一、踩坑记录:这些坑我都踩过
      • [坑 1:`torch.cuda.amp` 在新版本 PyTorch 上报错](#坑 1:torch.cuda.amp 在新版本 PyTorch 上报错)
      • [坑 2:`pretrained=True` 被移除](#坑 2:pretrained=True 被移除)
      • [坑 3:为了下载权重关掉 SSL 证书校验](#坑 3:为了下载权重关掉 SSL 证书校验)
      • [坑 4:`torch.load` 默认参数的安全隐患](#坑 4:torch.load 默认参数的安全隐患)
      • [坑 5:日志重复打印](#坑 5:日志重复打印)
      • [坑 6:TensorBoard 初始化失败导致训练中断](#坑 6:TensorBoard 初始化失败导致训练中断)
    • [十二、测试体系:46 个用例的分层设计](#十二、测试体系:46 个用例的分层设计)
    • 十三、已知边界与后续方向
    • 十四、总结
    • 项目源码下载

关键词:PyTorch 工业质检 缺陷检测 阈值寻优 ONNX 部署 FastAPI 类别不平衡 工程化

摘要:CIFAR-10 上跑通四分类,准确率 90%+,就能拿去产线检缺陷了吗?答案是否定的。本文从一个典型的"学术 Demo"项目出发,逐条拆解它上不了产线的 12 个致命问题,并给出完整的重构方案:清单驱动的数据层、批次分组划分防泄漏、业务阈值决策层、漏检/过杀/复核三指标体系、ONNX 一致性校验与在线服务。全文含可运行代码、真实运行数据与完整架构图。


一、先说结论:为什么你的分类模型上了产线就废了

先看一段非常典型的"学术 Demo"代码------它来自一个真实的缺陷分类项目:

python 复制代码
# 从 CIFAR-10 里抽前 4 类,"模拟"缺陷数据
target_classes: [0, 1, 2, 3]   # CIFAR-10 前4类模拟缺陷
class_names: ['scratch', 'dent', 'stain', 'none']

# 训练 5 轮,验证集直接从训练集随机切
train_val_split: 0.9

# 32×32 的图片硬拉到 224 送进 ResNet
transforms.RandomResizedCrop(224)

# 最后:argmax 一把梭
_, pred = outputs.max(1)

这段代码在 notebook 里能跑,准确率也能到 85%+。但把它放到产线上,会连续踩下面这些坑:

# 问题 产线后果
1 用 CIFAR-10 假装缺陷数据 32×32 的马车、青蛙,和产线上 2000×2000 的划痕毫无关系
2 训练集随机切验证集 同批次泄漏:同一卷料、同一班次、同一相机参数下拍的图同时出现在训练和验证集,验证指标虚高 10~20 个点
3 只报 Top-1 准确率 质量负责人真正要的是"漏了多少缺陷件",准确率根本没法决策
4 argmax 直接判定 缺陷件只要概率 0.51 > 0.49 就被判合格放行,漏检率不可控
5 无类别不平衡处理 ok 占 90%、划痕占 2%,模型学会"全猜 ok",准确率 90% 但缺陷一个没抓到
6 无数据脏治理 一张坏图直接让 DataLoader 崩在半夜
7 requirements 写了 fastapi/onnx 却没代码 训完不知道怎么部署,模型躺在 .pth 里吃灰
8 torch.cuda.amp / pretrained= 旧 API 新版 PyTorch 直接报错或警告
9 ssl._create_unverified_context 关掉证书校验,安全审计过不了
10 无任何测试 改一行代码不知道有没有弄坏别的
11 无配置快照 三个月后想复现实验,不知道当时用的什么参数
12 权重里塞满训练状态 部署时要拖着整个 optimizer 状态,还可能被 weights_only=True 拒绝加载

根本矛盾在于:学术分类追求"整体分对多少",工业质检追求"缺陷件绝不能放过"。这两件事的优化目标完全不同,甚至相互冲突。

本文的改造思路可以浓缩成一句话:

把"类别预测"升级为"业务判定",把"准确率"升级为"风险权衡"。


二、改造前后:一张表看清差距

维度 v1.0 学术 Demo v2.0 产线工程
数据源 CIFAR-10 前 4 类 目录 + CSV 清单,支持批次分组划分
划分方式 对训练集随机切 按 batch_id 分组划分,杜绝同批次泄漏
输入分辨率 32×32 强拉 224 可配置(默认 224),增强围绕产线扰动
类别不平衡 无 损失加权 + 加权采样,权重归一化并截断
核心指标 Top-1 准确率 每类 P/R/F1 + 漏检率 / 过杀率 / 复核率
决策方式 argmax 阈值保守判定 + 验证集自动寻优 + 人工复核带
部署 无 ONNX 导出 + 一致性校验 + FastAPI 服务
可复现 仅随机种子 种子 + 确定性算法 + 配置快照 + 结构化指标
工程化 单文件 main.py 6 个 CLI 子命令 + 分层包结构
测试 无 46 个单元测试(含端到端流水线)
代码注释 部分 逐行注释,覆盖率 100%

三、架构重构:五层分层设计

重构后的系统采用严格的五层分层,依赖方向自上而下,每层可独立替换与测试。

图 3-1 系统五层架构

复制代码
接入层   CLI(6 个子命令)/ HTTP(4 个接口)
   ↓
引擎层   trainer(训练循环·早停·阈值搜索)/ evaluator(概率推理·质检报告)
   ↓
模型层 + 决策层   models/factory(骨干工厂)/ utils/threshold(阈值判定·寻优)
   ↓
数据层   manifest(清单·分组划分)/ transforms(增强)/ dataset(兜底)/ dataloader(采样)
   ↓
基础层   config(校验·快照)/ seed(可复现)/ utils(日志·检查点·指标·可视化·设备)

这里有个刻意的设计 :决策层(utils/threshold.py)独立于模型层,且用纯 NumPy 实现。

为什么要这么设计?因为业务规则会频繁变(今天漏检率要求 1%,下个月客户要求 0.5%),而模型结构相对稳定。把决策逻辑抽离出来后:

  • 调阈值不需要重新训练,跑一遍验证集就行
  • 决策逻辑可以在没有 GPU、没有 PyTorch 的环境里测试
  • 离线回测和在线服务共用同一份代码,不会出现"离线看着准、上线不一样"

四、数据层:把"脏"当成常态

4.1 清单驱动,而不是靠目录结构猜标签

真实产线的标注来自 MES 系统或质检记录表,导出就是一张 CSV。所以清单才是第一等公民:

csv 复制代码
image_path,label,batch_id,split
scratch/IMG_20240501_0001.png,scratch,20240501,
ok/IMG_20240502_0007.png,ok,20240502,train

三个字段各有讲究:

  • image_path:相对 data.root 的路径,方便数据集整体迁移
  • batch_id:批次号,这是防泄漏的关键
  • split:可选。业务方指定了就用业务方的,没指定才自动划分

清单不存在时,自动扫描 root/<类别名>/*.jpg 目录结构生成,兼顾两种习惯。

4.2 批次分组划分:这 10 行代码能救你的验证指标

这是整个数据层最值钱的一段代码。产线数据的特点是:同一批次(同一卷料、同一班次、同一光照)内样本高度相似。如果随机划分,训练集和验证集里都有同一批次的图,模型只要"记住批次特征"就能拿高分------但上线遇到新批次就原形毕露。

python 复制代码
def group_split(records, val_ratio=0.15, test_ratio=0.15,
                seed=42, group_key="batch_id"):
    """按批次分组划分训练/验证/测试,避免同批次泄漏"""
    groups = defaultdict(list)                            # 批次 -> 样本列表
    for record in records:                                # 遍历全部记录
        groups[str(record.get(group_key, "common"))].append(dict(record))   # 按批次聚合

    group_items = sorted(groups.items(), key=lambda kv: kv[0])   # 按批次名排序,保证确定性
    rng = random.Random(seed)                             # 独立随机源,不污染全局
    group_items = list(group_items)                       # 转列表
    rng.shuffle(group_items)                              # 打乱"批次顺序"(注意:不是打乱样本)

    splits = {"train": [], "val": [], "test": []}         # 三个集合
    n_total = sum(len(v) for _, v in group_items)         # 样本总数
    n_test = int(n_total * test_ratio)                    # 测试集目标数量
    n_val = int(n_total * val_ratio)                      # 验证集目标数量

    for _, items in group_items:                          # 按打乱后的批次依次分配
        if len(splits["test"]) < n_test:                  # 测试集还没填满
            splits["test"].extend(items)                  # 整个批次进测试集
        elif len(splits["val"]) < n_val:                  # 验证集还没填满
            splits["val"].extend(items)                   # 整个批次进验证集
        else:                                             # 其余归训练
            splits["train"].extend(items)                 # 整个批次进训练集
    return splits                                         # 返回划分结果

核心就一句话:打乱的是"批次顺序",而不是"样本"。 同一个 batch_id 的所有样本必定落在同一个集合里。

实测提醒:如果你的数据没有 batch_id,系统会从文件名推断(比如 IMG_20240501_xxx.png 里的日期前缀)。但最好还是显式提供,别赌推断逻辑。

4.3 数据体检与训练期兜底

产线数据脏是常态:文件被移动过、标注写错、图片传输截断。系统在两个层面处理:

构造期体检(早失败,早发现):

python 复制代码
# 1. 文件不存在 → 过滤并记日志(常见于数据集迁移后)
# 2. 标签不在 class_names → 丢弃(避免训练时出现越界索引)
# 3. Pillow verify() 校验解码 → 剔除截断/损坏图片
# 4. 每类样本数 < min_samples_per_class → 直接报错提示补数

训练期兜底(不中断长任务):

python 复制代码
def __getitem__(self, idx):
    """带兜底重试的样本读取:坏图最多向后重试 3 次"""
    for attempt in range(self.retry + 1):                 # 最多重试 retry 次
        real_idx = (idx + attempt) % len(self.records)    # 向后顺延,取模防越界
        try:                                              # 尝试读取
            img = self._load_image(real_idx)              # 加载并解码图片
            if img is None:                               # 解码失败
                raise ValueError("图片解码失败")            # 抛错进入下一轮重试
            label = self.records[real_idx]["label_index"] # 取标签索引
            if self.transform is not None:                # 有变换
                img = self.transform(img)                 # 应用变换
            return img, label                             # 正常返回
        except Exception as exc:                          # 捕获所有异常
            self.skipped.append((real_idx, str(exc)))     # 记录被跳过的索引与原因
            continue                                      # 继续下一次尝试
    # 全部重试失败:返回全零样本,保证 DataLoader 不中断
    return torch.zeros(self.fallback_shape), 0            # 兜底返回

关键点:兜底返回全零样本 而不是抛异常。一个 8 小时的训练任务,不应该因为一张坏图在第 7 小时崩掉。被跳过的索引都存在 dataset.skipped 里,训练结束后统一排查。

4.4 增强:模拟的是产线扰动,不是花哨

增强不是为了"增加数据多样性",而是为了让模型对你无力改变的环境因素免疫。每一条算子都要对应一个真实物理原因:

增强算子 模拟的真实因素 你为什么改不了它
RandomResizedCrop 相机到产品距离 / 视野微变 夹具公差、振动
水平/垂直翻转 + 小角度旋转 夹具安装方向与角度偏差 人工上料方向不固定
ColorJitter(亮度/对比度/饱和度/色相) 车间光照漂移、灯管老化 车间照明不受你控制
GaussianBlur / RandomAdjustSharpness 失焦、镜头脏污 镜头会脏、会失焦
RandomErasing 局部遮挡、反光、脏点 现场就是有遮挡

强度分三档 light / medium / strong,由 augment_strength 配置。别一上来就 strong------过强的增强会让细小的划痕特征被破坏,反而学不到东西。

⚠️ 铁律 :评估、测试、推理三处必须用完全相同的确定性变换 (Resize → CenterCrop → ToTensor → Normalize)。训练增强只作用于训练集。我见过太多"验证准确率 95%、线上 70%"的事故,最后查出来是推理时忘了归一化,或者验证时误用了随机裁剪。

4.5 类别不平衡:双管齐下

产线上 ok 类样本天然占绝大多数(良率 95% 是好事,但对训练是灾难)。系统用了两个手段,且都必须做限幅:

手段一:损失加权

python 复制代码
# sqrt_inverse 比 inverse 温和,是默认选择
#   inverse       : w_c = N / (C * n_c)        ------ 极端不平衡时权重爆炸
#   sqrt_inverse  : w_c = sqrt(N / n_c)        ------ 推荐,抑制过头
# 权重归一化后按 max_weight(默认 20)截断,防止训练发散

手段二:加权采样

python 复制代码
# sampler=weighted 启用 WeightedRandomSampler,少数类被多次抽中
# ⚠️ 启用采样器后必须关闭 shuffle,否则语义冲突

为什么两个都要?损失加权改变的是梯度贡献 ,加权采样改变的是每个 epoch 看到的样本分布。小样本场景下,加权采样能让少数类在有限步数内被充分看到,收敛更快。


五、模型层:小样本友好的骨干工厂

统一结构是"预训练骨干 + 新分类头",分类头为 Dropout → Linear,并在 model.head 上保留引用,供分层学习率与冻结策略定位。

python 复制代码
def _replace_head(model, model_name, num_classes, dropout):
    """替换分类头为 Dropout + Linear,并在 model.head 上保留引用"""
    parent, in_features, attr = _head_info(model, model_name)   # 定位原分类头
    head = nn.Sequential(                        # 构造新分类头
        nn.Dropout(p=dropout),                   # Dropout:抑制小样本过拟合
        nn.Linear(in_features, num_classes),     # 新的分类线性层
    )
    if attr == "fc":                             # ResNet 系
        model.fc = head                          # 直接替换 fc
    else:                                        # MobileNetV3 / EfficientNet
        parent[-1] = head                        # 替换 classifier 序列最后一项
    model.head = head                            # ★ 保存引用,便于分层学习率定位
    return model

支持的骨干与选型建议:

骨干 可训练参数(4 类) 适用场景
resnet18 ≈ 11.2 M 默认,精度与速度均衡
resnet34 / resnet50 ≈ 21 M / 23.5 M 服务器端,精度优先
mobilenet_v3_small ≈ 0.93 M 产线边缘工控机
mobilenet_v3_large ≈ 4.2 M 边缘设备,精度优先
efficientnet_b0 ≈ 4.0 M 参数效率优先

几个容易忽略但很实用的兼容处理:

① 新旧 torchvision API 兼容

python 复制代码
# 优先用新 API:weights=Enum.DEFAULT
# 失败(旧版本 / 离线无权重)→ 自动降级随机初始化并打印告警
# 好处:离线环境、内网环境也能跑起来,不会直接崩

② 灰度图输入

有些产线用黑白相机。直接把单通道图送进 ResNet 会维度不匹配。系统重建首层卷积,权重按通道维取均值继承,保留预训练特征:

python 复制代码
# 原 conv1: (64, 3, 7, 7)  →  新 conv1: (64, 1, 7, 7)
# 新权重 = 原权重沿通道维取均值,形状 (64, 1, 7, 7)

③ 分阶段冻结(小样本救命)

yaml 复制代码
model:
  freeze_backbone: false    # true = 只训练分类头(数据极少时先用这个)
  freeze_stages: 0          # 冻结 stem + 前 N 个 stage(ResNet 有效)

数据量 < 1000 张时,强烈建议 freeze_backbone=true 先训 5~10 轮,再解冻全网络微调。


六、训练引擎:稳定性优先于花哨

6.1 优化配置

配置项 默认值 为什么这么定
optimizer adamw AdamW 解耦权重衰减,小样本更稳
lr 3e-4 微调预训练模型的常用起点
head_lr_scale 5.0 分类头是随机初始化的,需要更快收敛
weight_decay 0.05 抑制过拟合
label_smoothing 0.05 缓解人工标注噪声(质检标注本来就有争议)
grad_clip 1.0 防止异常样本导致梯度爆炸
amp auto CPU 自动关闭,不会因为没 GPU 就报错

head_lr_scale 是个容易被忽略但效果明显的技巧:骨干带着 ImageNet 预训练权重,已经"差不多对了",学习率要小;分类头是随机初始化的,学习率要放大 5 倍。用同一个学习率,要么头训不动,要么骨干被破坏。

6.2 学习率调度:预热是必需的

yaml 复制代码
train:
  scheduler: cosine
  scheduler_params:
    T_max: 30
    warmup_epochs: 2        # ★ 预热 2 轮
    min_lr: 1.0e-6          # 余弦退火下限

为什么必须预热? 预训练骨干的权重是"好"的,如果第一轮就用 3e-4 的学习率猛冲,会把预训练特征破坏掉。预热让学习率从 0 线性上升到目标值,给骨干一个缓冲。

实现上用 SequentialLR 把预热和主调度拼起来,旧版本 PyTorch 自动降级为主调度。

6.3 早停:统一转成"越大越好"

yaml 复制代码
train:
  early_stopping:
    enable: true
    monitor: val_macro_f1    # val_loss / val_acc / val_macro_f1 / val_miss_rate
    mode: max                # max=越大越好,min=越小越好
    patience: 8
    min_delta: 0.0005

实现小技巧:把监控值统一转成"越大越好"(损失和漏检率取负),这样比较逻辑只有一套,不容易写反。

6.4 产物契约:权重与元信息分离

这是个安全性 + 可用性双重考虑的设计:

复制代码
outputs/checkpoints/
├── best_model.pt           # 仅 state_dict → torch.load(weights_only=True) 可安全加载
├── best_model.meta.json    # 类别名、ok 索引、阈值、复核带、分辨率、策略、指标快照
├── last_model.pt           # 末轮权重
└── epoch_N.pt              # 完整训练状态(含 optimizer,可恢复训练)
json 复制代码
{
  "class_names": ["scratch", "dent", "stain", "ok"],
  "ok_index": 3,
  "threshold": 0.86,
  "review_band": 0.15,
  "image_size": 224,
  "strategy": "threshold",
  "metrics": {"val_macro_f1": 0.87, "val_miss_rate": 0.008}
}

部署侧只需要 model.onnx + model.meta.json 两个文件,不需要训练配置、不需要源码、不需要 PyTorch。这样部署同学拿到东西就能干活。


七、决策层:这才是整个系统的核心

前面所有工作都是铺垫。这一节是本文的重点。

7.1 三态判定,而不是二分类

模型输出的是概率,但产线要的是动作。系统输出三态:

判定 含义 下游动作
OK 判定合格 自动放行,进入下一工位
NG 判定存在缺陷 拦截、返修或报废
REVIEW 置信度不足 推送人工复检工位

图 7-1 在线推理与业务决策流程

7.2 判定规则:只有"足够像合格"才放行

python 复制代码
def decide_from_probs(probs, ok_index, threshold):
    """按阈值把概率矩阵转换为类别预测(向量化实现)"""
    probs = np.asarray(probs, dtype=np.float64)   # 转浮点数组
    if probs.ndim != 2:                           # 必须是二维
        raise ValueError("probs 必须是形状为 (N, C) 的二维数组")
    ok_probs = probs[:, ok_index]                 # 取出"合格类"的概率
    tmp = probs.copy()                            # 复制一份,避免修改入参
    tmp[:, ok_index] = -np.inf                    # ★ 屏蔽合格类,便于取"最可能的缺陷类"
    defect_pred = tmp.argmax(axis=1)              # 最可能的缺陷类别索引
    # ★ 核心:只有 p(ok) >= 阈值才放行,否则保守判为最可能的缺陷
    preds = np.where(ok_probs >= threshold, ok_index, defect_pred)
    return preds.astype(np.int64)                 # 返回整数预测

这段代码体现了保守原则:

  • p(ok) >= 阈值 → 放行(足够确信是合格品)
  • p(ok) < 阈值 → 不是判 ok,而是判"最可能的那个缺陷类"

注意第二点。很多人的实现是"不 ok 就判 NG",但这样丢失了缺陷类型信息------维修工需要知道是划痕还是凹坑。这里用 tmp[:, ok_index] = -inf 屏蔽掉 ok 类再 argmax,就拿到了最可能的缺陷类型。

复核带独立生效:

python 复制代码
def need_review_flags(probs, review_band):
    """标记最大置信度低于复核带的样本(需人工复核)"""
    max_probs = probs.max(axis=1) if probs.size else np.zeros(0)   # 每行最大概率
    return (max_probs < float(review_band)).astype(np.int64)       # 返回 0/1 标记

复核带独立于阈值 。max(p) < review_band 的样本,无论被判为什么,一律转人工。这是给"模型也没把握"的情况留的出口。

7.3 业务指标:这三个才是老板要看的

指标 定义 业务含义
漏检率 miss_rate 缺陷件被判为 ok 的比例 💀 最危险,直接决定客诉风险
过杀率 overkill_rate 合格件被判为缺陷的比例 💰 影响产能、返工与报废成本
复核率 review_rate 置信度低于复核带的比例 👷 决定人工复检人力成本

三者的关系:漏检和过杀此消彼长。阈值调高 → 更严格 → 漏检率下降、过杀率上升。

python 复制代码
miss_rate     = 漏检数 / 缺陷总数     # 越低越好,但不能为零(那意味着全判 NG)
overkill_rate = 过杀数 / 合格总数     # 越低越好,但不能为零(那意味着全放行)
review_rate   = 复核数 / 总样本数     # 越低越好,但过低说明复核带形同虚设

7.4 阈值寻优:把风险权衡交给算法

阈值不该拍脑袋定。系统的做法是在验证集上做约束搜索:

python 复制代码
def search_threshold(probs, labels, ok_index,
                     target_miss_rate=0.01, grid_start=0.5,
                     grid_end=0.99, num_points=50, review_band=0.0):
    """在验证集上搜索满足漏检率约束的最优阈值

    选择策略:优先满足 miss_rate <= target;
    在满足条件的候选中取"过杀率最小"的阈值;
    若没有任何阈值满足,则取漏检率最低者并标记 feasible=False,供人工介入。
    """
    grid = np.linspace(float(grid_start), float(grid_end), int(num_points))   # 阈值网格
    results = []                                                              # 每个阈值的结果
    for thr in grid:                                                          # 遍历候选阈值
        results.append(evaluate_threshold(probs, labels, ok_index, float(thr), review_band))

    # ① 约束筛选:只保留满足漏检率上限的候选
    feasible = [r for r in results if r["miss_rate"] <= float(target_miss_rate) + 1e-12]

    if feasible:                                                              # ② 存在可行解
        # 过杀最小优先;过杀相同时取更宽松(更小的)阈值,避免过度保守
        best = min(feasible, key=lambda r: (r["overkill_rate"], r["threshold"]))
        best = dict(best)
        best["feasible"] = True                                               # 标记可行
    else:                                                                     # ③ 无解兜底
        best = dict(min(results, key=lambda r: r["miss_rate"]))               # 取漏检最低者
        best["feasible"] = False                                              # ★ 标记不可行

    best["target_miss_rate"] = float(target_miss_rate)                        # 记录目标
    best["curve"] = results                                                   # ★ 完整曲线,便于画图/审计
    return best

这里有一个我认为最重要的设计决策:

python 复制代码
best["feasible"] = False    # 没有任何阈值满足漏检约束时

当验证集上找不到任何满足 target_miss_rate 的阈值时,系统不会自动放宽业务目标 ,而是返回 feasible=False 并把漏检最低的那个阈值给你,同时明确告诉你"这个目标达不到"。

这是刻意的。工业场景里,"系统悄悄接受了一个高风险阈值"比"系统报错说做不到"要危险得多。前者会让你在不知情的情况下持续漏检,后者至少逼你去补数据、换模型或调特征。

图 7-2 阈值---业务风险权衡曲线 (本项目 search_threshold 真实产出)

图中可以看到典型的权衡形态:阈值从 0.5 升到 0.99,漏检率(红)单调下降、过杀率(蓝)单调上升。绿色虚线是在"漏检 ≤ 1%"约束下搜索到的最优阈值------它选的是满足约束时过杀最小的那个点,而不是漏检最低的点(后者会让过杀率飙到不可接受)。

7.5 配置方式

yaml 复制代码
decision:
  ok_class: ok
  strategy: threshold          # argmax=直接取最大概率;threshold=按阈值保守判定
  target_miss_rate: 0.01       # ★ 线上可接受的漏检率上限
  review_band: 0.15            # 置信度低于该值 → 人工复核
  threshold_grid: 0.5          # 阈值搜索起点
  auto_threshold: true         # 训练结束自动在验证集上搜索最优阈值

调参建议:

  • target_miss_rate 由质量协议决定,不是由模型能力决定。客户要求 0.5% 就设 0.005
  • 设了达不到 → 看 feasible 字段,为 false 就该补数据,而不是改目标
  • review_band 由人力预算决定:复检工位能处理多少件,反过来定这个值

八、评估与报告:给质量负责人看得懂的东西

评估不止输出一个数字,而是生成一份可直接决策的质检报告。

8.1 报告字段

json 复制代码
{
  "accuracy": 0.6111,
  "macro_f1": 0.3694,
  "argmax_accuracy": 0.8333,          // ★ 不套阈值的原始模型能力
  "per_class": [
    {"class": "scratch", "precision": 0.0, "recall": 0.0, "f1": 0.0},
    {"class": "dent",    "precision": 0.0, "recall": 0.0, "f1": 0.0},
    {"class": "stain",   "precision": 1.0, "recall": 0.5, "f1": 0.667},
    {"class": "ok",      "precision": 0.68,"recall": 1.0, "f1": 0.811}
  ],
  "confusion_matrix": [[0,0,0,7],[0,0,7,0],[0,0,7,0],[0,0,0,15]],
  "business": {
    "miss_rate": 0.3333,              // 漏检率
    "overkill_rate": 0.0,             // 过杀率
    "defect_total": 21, "missed": 7,
    "ok_total": 15, "overkill": 0,
    "review_rate": 0.0, "review_count": 0
  },
  "recommended_threshold": {
    "threshold": 0.85,
    "feasible": true,
    "miss_rate": 0.0,
    "overkill_rate": 0.8667
  },
  "artifacts": {
    "confusion_matrix": "outputs/reports/test_confusion_matrix.png",
    "threshold_curve": "outputs/reports/test_threshold_curve.png"
  }
}

argmax_accuracy 这个字段非常有用 。它可能高于最终 accuracy(比如这里 0.83 vs 0.61),说明:模型本身判别能力还行,是阈值太严导致的整体准确率下降。

这能帮你快速定位问题:

  • argmax_accuracy 低 → 模型弱,去调模型/补数据
  • argmax_accuracy 高但最终 accuracy 低 → 阈值严,这是业务选择,不是 bug

图 8-1 混淆矩阵(行归一化)

📌 说明:这张图和上面的指标来自 300 张合成数据训练 8 轮的演示结果,用于展示报告能力,不代表真实产线指标。真实产线应使用足量数据与更多轮次。

图 8-2 训练过程曲线 (读取 metrics.jsonl 绘制)

小样本下验证指标有波动是正常的,这正是早停机制存在的意义------它会挑出最佳轮次(图中红色虚线标注)的权重,而不是最后一轮。

8.2 报告产物的工程细节

python 复制代码
# 用 matplotlib 的 Agg 无界面后端 → 服务器环境可直接出图
matplotlib.use("Agg")

# 中文字体缺失时静默降级 → 不会因字体问题中断报告生成
try:
    plt.rcParams["font.sans-serif"] = ["Noto Sans CJK SC"]
except Exception:
    pass    # 降级为默认字体,图照出,只是中文可能显示成方块

九、部署:ONNX + FastAPI

requirements.txt 里写了 fastapi 和 onnx 却没有任何代码,这是 Demo 项目的通病。补齐部署链路:

图 9-1 典型部署拓扑:相机 → 边缘工控机 → 云端训练 → 质量看板

9.1 ONNX 导出与一致性校验

这一步千万别省。 ONNX 转换过程中算子实现可能有细微差异,不校验就上线,等于埋雷。

bash 复制代码
# 导出 ONNX 并做数值一致性校验(随机输入对比 torch 与 onnxruntime 输出)
python -m src.cli.export --checkpoint ./outputs/checkpoints/best_model.pt \
    --onnx ./outputs/model.onnx --verify

# 实测输出
# 一致性校验:{'ok': True, 'max_diff': 7.629e-06, 'tolerance': 0.0001}

实测最大误差 7.6e-06,远低于 1e-4 容差。导出配置:opset 17、输入名 input、输出名 logits、支持动态 batch。

9.2 HTTP 接口

接口 方法 说明
/healthz GET 健康检查,返回状态与类别列表
/model_info GET 类别、ok 索引、阈值、复核带、分辨率、后端类型
/predict POST (multipart) 单图判定
/predict/batch POST (multipart) 多图批量判定

9.3 判定结果示例

json 复制代码
{
  "image": "IMG_20240501_0001.png",
  "pred_label": "scratch",
  "pred_index": 0,
  "confidence": 0.873,
  "decision": "NG",
  "need_review": false,
  "threshold": 0.86,
  "probs": {
    "scratch": 0.873, "dent": 0.060,
    "stain": 0.040,  "ok": 0.027
  }
}

注意 threshold: 0.86,但 p(ok) = 0.027 远低于它 → 判 NG。同时 confidence = 0.873 > review_band = 0.15 → 不需要复核,直接拦截。每个字段都能追溯到判定依据,出了客诉能复盘。

9.4 双后端一致性设计

python 复制代码
class DefectPredictor:
    """统一预测器:PyTorch 与 ONNX Runtime 暴露完全一致的接口"""
    # predict_pil()    ------ 单张 PIL 图
    # predict_file()   ------ 单个文件
    # predict_paths()  ------ 路径列表
    # predict_dir()    ------ 整个目录

离线回测用 PyTorch,线上服务用 ONNX,但判定逻辑是同一份代码。这避免了经典的"离线看着准、上线不一样"。


十、实战:5 分钟跑通全流程

bash 复制代码
# ① 安装依赖
pip install -r requirements.txt

# ② 生成合成演示数据(★ 不需要下载任何数据集,离线可跑)
python -m src.cli.make_demo_data --out-root ./data/raw --manifest ./data/manifest.csv \
    --classes scratch,dent,stain,ok --counts 40,40,40,80 --image-size 128

# ③ 训练(离线冒烟:2 轮、不加载预训练权重)
python -m src.cli.train --config configs/smoke.yaml --epochs 2 --set model.pretrained=false

# ④ 评估并生成质检报告(混淆矩阵 + 阈值曲线 + JSON 报告)
python -m src.cli.evaluate --config configs/smoke.yaml \
    --checkpoint ./outputs/checkpoints/best_model.pt --split test

# ⑤ 导出 ONNX 并做数值一致性校验
python -m src.cli.export --checkpoint ./outputs/checkpoints/best_model.pt \
    --onnx ./outputs/model.onnx --verify

# ⑥ 批量判定(torch 权重或 onnx 模型都支持)
python -m src.cli.predict --onnx ./outputs/model.onnx --input ./data/raw/ok \
    --output ./outputs/ok_predictions.csv

# ⑦ 启动在线服务:http://127.0.0.1:8000/docs
python -m src.cli.serve --onnx ./outputs/model.onnx --port 8000

常用配置覆盖

不用改 YAML,命令行直接覆盖:

bash 复制代码
# 改学习率、批次与轮数
python -m src.cli.train --config configs/default.yaml \
    --set train.lr=5e-4 data.batch_size=64 train.epochs=50

# 换骨干 / 冻结骨干 / 灰度输入
python -m src.cli.train --config configs/default.yaml \
    --set model.name=efficientnet_b0 model.freeze_backbone=true data.image_mode=L

# 指定数据与输出目录(多实验对比)
python -m src.cli.train --config configs/default.yaml --manifest ./data/manifest.csv \
    --data-root ./data/raw --output-dir ./outputs/exp001

项目结构

复制代码
industrial_defect_inspection/
├── configs/                  # default.yaml(生产)/ smoke.yaml(CI 冒烟)
├── src/
│   ├── config.py             # 配置加载、命令行覆盖、校验、快照
│   ├── seed.py               # 全局随机种子与确定性算法
│   ├── data/                 # manifest / transforms / dataset / dataloader
│   ├── models/               # factory:骨干工厂与分类头替换
│   ├── engines/              # trainer:训练循环;evaluator:质检报告
│   ├── inference/            # predictor / export_onnx / service
│   ├── utils/                # logger / checkpoint / metrics / threshold / visualization
│   └── cli/                  # 6 个子命令入口
├── tests/                    # 46 个单元测试(含端到端冒烟)
├── scripts/                  # run_demo.sh / run_tests.sh
├── README.md
└── ARCHITECTURE.md

十一、踩坑记录:这些坑我都踩过

坑 1:torch.cuda.amp 在新版本 PyTorch 上报错

python 复制代码
# ❌ 旧写法(PyTorch 2.x 已废弃)
from torch.cuda.amp import autocast, GradScaler

# ✅ 新写法 + 兼容封装
# 优先 torch.amp.autocast(device_type=...),失败降级 torch.cuda.amp
# 并且:CPU 环境自动关闭 AMP,不会因为没 GPU 就崩

坑 2:pretrained=True 被移除

python 复制代码
# ❌ 旧写法
models.resnet18(pretrained=True)

# ✅ 新写法
models.resnet18(weights=ResNet18_Weights.DEFAULT)

系统做了双重兜底:优先用新 API,失败则自动降级随机初始化并打印告警(不是静默失败)。这样内网/离线环境也能跑。

坑 3:为了下载权重关掉 SSL 证书校验

python 复制代码
# ❌ 危险!原项目就是这么干的
import ssl
ssl._create_default_https_context = ssl._create_unverified_context

这行代码会全局关闭证书校验 ,任何 HTTPS 请求都变得可被中间人攻击,安全审计必挂。正确做法是:权重加载失败就降级随机初始化 + 告警。不要为了可用性牺牲安全性。

坑 4:torch.load 默认参数的安全隐患

python 复制代码
# ❌ 权重里塞满 optimizer 状态,部署时 weights_only=True 会拒绝加载
torch.save({"model": model.state_dict(), "optimizer": opt.state_dict()}, path)

# ✅ 分离:部署权重只存 state_dict
torch.save(model.state_dict(), "best_model.pt")        # 可 weights_only 加载
torch.save(full_state, "epoch_10.pt")                  # 显式声明不可 weights_only
# 元信息(类别/阈值/分辨率)单独存 .meta.json

坑 5:日志重复打印

python 复制代码
# ❌ 每次实例化 Logger 都 addHandler,测试里实例化 5 次就打印 5 遍
self.logger.addHandler(console_handler)

# ✅ 先清理历史 handler
for h in list(self.logger.handlers):
    self.logger.removeHandler(h)

坑 6:TensorBoard 初始化失败导致训练中断

python 复制代码
# ✅ 降级处理:TB 挂了也要能训练
try:
    self.tb_writer = SummaryWriter(tb_dir)
except Exception:
    self.tb_writer = None    # 降级为纯文本日志

十二、测试体系:46 个用例的分层设计

层级 覆盖范围 依赖 PyTorch 用途
纯逻辑 清单解析与划分、指标、阈值搜索、配置校验 ❌ 否 最快的回归网,秒级反馈
数据管线 变换、数据集、坏图兜底、采样与加载器 ✅ 验证数据层健壮性
模型 各骨干前向、灰度输入、冻结策略 ✅ 防止换骨干踩坑
引擎 真实训练 2 轮、验证指标、元信息可部署 ✅ 端到端训练闭环
评估 概率推理、报告字段、图表产物 ✅ 报告契约不被破坏
推理 Torch/ONNX 一致性、批量预测、CSV/JSON 导出 ✅ 部署链路可信赖
服务 全部 4 个 HTTP 接口 ✅ 接口契约
端到端 造数 → 训练 → 评估 → 导出 → 预测 ✅ 全 CLI 串通
bash 复制代码
python -m pytest -v
# 46 passed in ~30s

为什么把"纯逻辑"单独分层? 阈值搜索、业务指标、清单划分这些是业务规则,改动频率远高于模型代码,而且完全不依赖 GPU。把它们做成秒级可测的纯 NumPy 函数,能极大提升迭代信心。

数据全部由 make_demo_data 合成,不依赖外网下载任何数据集或预训练权重,CI 可直接运行。


十三、已知边界与后续方向

诚实说明这套系统的能力边界:

边界 说明 后续方向
只分类,不定位 输出整图类别,不给出缺陷位置与尺寸 引入检测/分割模型,清单扩展为带框标注
阈值依赖数据分布 换相机、光照改造、新产品导入后需重跑寻优 接入阈值漂移监控看板
小样本风险 每类样本过少时指标置信度低,系统仅告警不阻断 主动学习 + 难例挖掘
ONNX 导出链路 默认用稳定的 TorchScript 路径(dynamo=False) 切 dynamo 时需复核容差与动态轴

关于阈值漂移的提醒(这是我见过最现实的问题):

阈值是在验证集上寻优得到的,它只在"验证集分布 ≈ 线上分布"时成立。一旦:

  • 更换相机或镜头
  • 车间照明改造
  • 导入新产品/新产线
  • 供应商更换原材料

就必须重新采集数据并重跑阈值寻优。建议把漏检率/过杀率做成监控看板,一旦漂移超阈值就告警。


十四、总结

从学术 Demo 到产线工程,改动最大的不是模型,而是围绕模型的工程决策。回头看,最有价值的五个改动是:

  1. 按批次分组划分 ------ 一行配置的差别,可能是验证指标虚高 20 个点的差别
  2. 阈值保守判定替代 argmax ------ 让漏检率从"不可控"变成"可配置、可验证"
  3. feasible=False 的诚实设计 ------ 达不到业务目标就明说,绝不悄悄放宽
  4. 权重与元信息分离 ------ 部署侧零依赖,安全性与可用性双赢
  5. 决策层用纯 NumPy 实现 ------ 业务规则可独立测试、独立迭代,不受模型改动影响

最后给一句我认为最重要的建议:

在工业质检场景,不要问"模型准确率多少",要问"漏检率能不能接受、过杀率要付出多少代价"。

前者是技术指标,后者才是业务决策。

项目源码下载

去下载 ⏬

相关推荐
wflynn1 小时前
语言判别增强多语言语音模型的语言学习能力
人工智能·ai
空心木偶☜1 小时前
A2A2A(多智能体协作)
开发语言·python·ai·ai编程
其实防守也摸鱼1 小时前
DeepSeek Harness 开源贡献手记:从 Issue 到 Merge 的完整旅程
android·数据库·学习·ai·oracle·自动化
ai小陈10 小时前
GPU服务器租用存储验收:检查点写入与磁盘吞吐实战
运维·服务器·人工智能·ai·ssh·gpu算力
bigdata-余建新11 小时前
week10
ai
JackSparrow41411 小时前
和AI一起将全部CSDN博文迁移到个人博客站
人工智能·程序人生·ai·github·cloudflare·astro·静态博客
搬砖的小码农_Sky11 小时前
AI Agent:如何处理Claude Code 最近版本(2026年更新)引入的模型上下文限制
人工智能·windows·ai·ai编程
kaixin_啊啊12 小时前
【零基础学AI】第 1 章课后练习与答案
人工智能·ai
李航198314 小时前
AI定制柜建模,需要详细的建模规范和标准流程
人工智能·python·计算机视觉·ai·ai编程