news 2026/9/16 15:37:59

钢材表面缺陷检测:基于PyTorch的语义分割实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
钢材表面缺陷检测:基于PyTorch的语义分割实战

简介:面向计算机相关专业毕业设计、期末大作业及项目实战学习者,这一基于Python的钢材表面缺陷检测与分割竞赛解决方案覆盖了从数据增强、网络构建、损失函数设计到训练评估的完整流程。压缩包共7个文件,包含5个Python脚本——分别实现在线数据增强、模型定义、Lovász Softmax损失、训练及带增强评估等核心模块——外加1个可直接加载的预训练权重文件和1份说明文档,整体仅3.5MB,结构紧凑、易于复现。项目源自导师指导并认可的高分设计,评审98分,源码均经过本地编译与严格调试,可稳定运行,既能帮助初学者快速理解语义分割的工程代码组织方式,也为进阶者提供了良好的二次开发基础。目前已有83人学习下载,适用于课程设计、竞赛备赛练习,也可作为钢铁表面质量检测方向的基础参考工程。

1. 为什么钢材表面缺陷检测要先做分割,而不是只做分类

钢材表面缺陷检测与自然图像分类的差异不在网络结构,而在输出维度。裂纹、麻点、划痕这类缺陷往往细长、低对比度,灰度分布和氧化铁皮背景高度重叠,全图分类只能告诉你“有没有缺陷”,却给不出缺陷的位置、面积和形态,而这些恰恰是竞赛评分和产线复检最看重的信息。竞赛评测通常提交像素级 mask、用 mIoU 计分,所以完整解决方案必然落到语义分割或实例分割。这套基于 Python 的源码包把整条链路串好了:train.py 管理训练流程,model.py 定义分割网络,lovasz_softmax.py 提供面向 IoU 的损失,OnlineAugment.py 做在线增强,eval_with_aug.py 负责带增强的评估。正在做毕业设计、期末大作业,或第一次接触分割任务的 Python 开发者,可以直接对照源码跑通,再替换成自己的数据。

2. 数据解构:标注格式、类别分布与 Dataset 输入侧设计

2.1 钢材表面缺陷的类型与分布

拿到源码包先打开 README.md,确认数据目录和标注格式。公开钢材表面缺陷数据的标注粒度有两种:框和像素 mask。竞赛里为了算 mIoU,几乎都是像素级标注,具体编码又分两种——单通道 PNG 和 RLE 字符串。图像本身多数是灰度图,缺陷类型基本落在以下六类:

缺陷类型典型外观标注难点
裂纹细线状、低对比易与划痕混淆
夹杂颗粒状、亮度突变边界判定困难
斑块大面积灰度异常边缘过渡带长
麻点密集小圆坑单个体积小、数量大
氧化铁皮压入块状、纹理不均匀和背景融为一体
划痕长条状、方向性强可能横跨整幅图像

注意一个容易被忽略的分布事实:无缺陷的负样本占比经常超过三分之一。如果直接拿全图分类的思路去做分割,模型很容易学会把所有像素预测成背景,loss 看起来在下降,mIoU 却一直是 0。所以第一个环节要做的不是选模型,而是先统计 mask 的面积分布和类别占比,再决定损失函数和类别权重怎么配。

2.2 带预处理的 Dataset 类

train.py 里的数据入口通常是一个自定义 Dataset,读取图像和 mask,缩放、转 tensor、归一化。最常见的实现是 OpenCV 读图,训练时缩放到统一尺寸,参考写法如下:

import cv2 import numpy as np import torch from torch.utils.data import Dataset class SteelDefectDataset(Dataset): def __init__(self, image_paths, mask_paths, size=(256, 512), augment=None): self.image_paths = image_paths self.mask_paths = mask_paths self.size = size self.augment = augment def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img = cv2.imread(self.image_paths[idx], cv2.IMREAD_GRAYSCALE) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, self.size, interpolation=cv2.INTER_LINEAR) mask = cv2.resize(mask, self.size, interpolation=cv2.INTER_NEAREST) # mask 必须用最近邻重采样,线性插值会在缺陷边缘生成 0.x 的中间值 # 这些中间值会被当成额外类别,污染交叉熵损失 if self.augment is not None: augmented = self.augment(image=img, mask=mask) img = augmented['image'] mask = augmented['mask'] img_t = torch.from_numpy(img).float().unsqueeze(0) / 255.0 mask_t = torch.from_numpy(mask).long() return img_t, mask_t

这段代码里有两个点值得展开。第一,unsqueeze(0)把灰度图从(H, W)变成(1, H, W),因为 PyTorch 的 Conv2d 要求输入是四维(B, C, H, W),这里 batch 维度还没加,先补上 channel 维度。第二,归一化直接除以 255 而不是逐通道做 mean/std,速度快,而且对钢材这种灰度分布相对固定的场景,和完整归一化没有明显差别。真正影响精度的不是归一化方式,而是 mask 的 resample 方式,所以 mask 统一用cv2.INTER_NEAREST

提示:如果 mask 的值域出现 0 和 255,说明原数据集把标签存成了可视化 PNG,训练前必须把 255 改成 1,否则交叉熵会把它当成独立的第二个类别。

2.3 RLE 解码与缺陷感知划分

如果竞赛给的是 RLE 字符串,要先解码成二维 mask。解码逻辑本身不复杂,但有边界条件要处理:空字符串表示无缺陷,应当返回全零矩阵,不能直接 split 报错。

def rle_decode(rle_string, shape): if rle_string is None or rle_string == '': return np.zeros(shape[0] * shape[1], dtype=np.uint8) s = rle_string.split() starts = np.asarray(s[0::2], dtype=int) - 1 lengths = np.asarray(s[1::2], dtype=int) flat = np.zeros(shape[0] * shape[1], dtype=np.uint8) for start, length in zip(starts, lengths): flat[start:start + length] = 1 return flat.reshape(shape)

RLE 的坐标区间经常是 1-based,所以starts要减 1;lengths描述的是开区间长度,直接用切片flat[start:start + length]赋值即可。钢材缺陷的 RLE 通常是按行扫描的,所以 decode 完 reshape 成(H, W)时顺序要和标注约定一致,否则 mask 会整体错位。

数据划分也不能直接 random split,要按缺陷类别做分层。同一张图里可能同时出现多个缺陷类别,一个简单的做法是先统计每张图包含的类别 id,再按类别频率做 StratifiedGroupKFold,保证 train/val 中的类别比例接近。训练尺寸和最终推理尺寸不一致时,记得在验证阶段把预测 mask 还原到原图尺寸再做 mIoU 统计,否则分数会虚高或虚低 1-2 个点。

3. model.py 与 lovasz_softmax.py:分割网络结构和损失函数选型

3.1 为什么选 encoder-decoder 而不是直接接全连接

钢材缺陷的尺度跨度很大,划痕可能覆盖整幅图像,麻点却只有几个像素。如果主干一路下采样到 1/32,一个 8×8 像素的麻点在特征图上只剩 1×1,几乎不可能被恢复。所以 model.py 里常见的分割网络是 DeepLabV3 或 UNet 这类 encoder-decoder 结构,encoder 提供语义,decoder 恢复空间分辨率。

工业上最稳的做法是拿一个预训练好的 ResNet 做 backbone,然后在最后一层接 ASPP(空洞空间金字塔池化),用不同膨胀率的空洞卷积并行捕捉多尺度上下文,最后上采样回原分辨率。参考实现:

import torch import torch.nn as nn import torchvision class ASPP(nn.Module): def __init__(self, in_channels, out_channels=256, rates=(6, 12, 18)): super().__init__() self.branches = nn.ModuleList() for rate in rates: self.branches.append( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=rate, dilation=rate, bias=False) ) self.project = nn.Sequential( nn.Conv2d(len(rates) * out_channels, out_channels, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), nn.Conv2d(out_channels, 1, kernel_size=1) ) def forward(self, x): outs = [branch(x) for branch in self.branches] return self.project(torch.cat(outs, dim=1))

这里 rates 是膨胀率。rate=6 表示卷积核在 3×3 的采样点上隔 6 个像素取一次,感受野瞬间变大,但参数不增加。把三个不同膨胀率的特征拼到一起,模型就能同时看到细划痕和大的氧化铁皮压入区域。output 通道最后接 1,是因为很多竞赛只有一个“缺陷”类别,输出一张(B, 1, H, W)的 logits 图就够了;多类别场景把out_channels改成类别数即可。

3.2 Lovász softmax 为什么比 cross entropy 适合 mIoU 竞赛

竞赛评分是 mIoU,交叉熵却是逐像素独立优化,两者在语义上是错位的。IoU 是集合层面的重叠度量,缺陷像素少的时候,交叉熵会把大量梯度分配给背景,前景区域学得很慢。Lovász softmax 的思路是把离散的 Jaccard 指数扩展成连续可导的损失,让梯度直接优化 IoU 的代理值。

那个 lovasz_softmax.py 文件的核心是排序梯度的过程,简化为 lovasz_grad:

def lovasz_grad(gt_sorted): p = len(gt_sorted) gts = gt_sorted.sum() intersection = gts - gt_sorted.float().cumsum(0) union = gts + (1 - gt_sorted).float().cumsum(0) jaccard = 1.0 - intersection / union if p > 1: jaccard[1:p] = jaccard[1:p] - jaccard[0:-1] return jaccard

gt_sorted是按预测概率排序后的标签,0 表示背景,1 表示缺陷。cumsum 是前缀和,intersection通过“总正样本数减去已扫描的正样本累计数”不断更新,union则是把背景的累计数加到总正样本上。最终jaccard计算每个排序位置的 IoU 增量,再取差分得到梯度权重。简单理解就是:预测越离谱的 hard negative,Lovász 给它的梯度权重越大。

3.3 损失搭配和类别权重

实际训练很少单独用 Lovász,因为它对冷启动阶段不友好。我一般会把它和交叉熵按loss = 0.7 * lovasz + 0.3 * ce混合。前几个 epoch,交叉熵提供稳定的逐像素梯度,等 mask 预测有一点形状之后,Lovász 再主导优化方向。这个比例不固定,数据类别特别不平衡时,可以把 Lovász 权重提到 0.8。

损失函数优化目标适用场景注意点
CrossEntropy逐像素分类通用基线对类别不平衡敏感
Dice Loss前景区域重叠二分类 mask训练初期容易震荡
Lovász SoftmaxJaccard 指数mIoU 竞赛 / 掩膜评价需要对 logits 先过 softmax

还有一个细节:由于 mask 的前背景比例极端,可以在交叉熵里把背景的类别权重调低,例如weight=torch.tensor([0.5, 1.0])。这不会直接提升 mIoU,但能防止 CNN 前几轮就把所有像素判定为背景,导致 loss 下降到一定数值后 mIoU 卡死在 0。

4. OnlineAugment.py 与 train.py:在线增强、训练循环与 checkpoint 管理

4.1 在线增强为什么要放在数据读取端

钢材图像往往来自同一卷钢带的连续拍摄,直接连续采样会碰到大量高度相似的帧,模型很快就会记住“背景统计量”,泛化性极差。OnlineAugment 的好处是每个 epoch 都用不同的随机变换组合,相当于每个 epoch 得到一份新数据,而离线增强会把磁盘占用放大 N 倍。放在 Dataset 里还避免了在 GPU 上做增强的同步开销。

4.2 增强组合的常见做法

这是一个稳健、不会破坏缺陷物理形态的增强管线,使用 albumentations:

import albumentations as A train_aug = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.OneOf([ A.GaussianBlur(blur_limit=(3, 5), p=0.5), A.MotionBlur(blur_limit=3, p=0.5), ], p=0.3), A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5), ])

注意,这里几乎没有使用 Scale/Translate/Rotate 这类任意角度形变。钢材表面的氧化铁皮压入和麻点是刚性的,不存在物体被扭曲的物理可能,过度形变反而会让模型学到错误的形状先验。RandomRotate90 本质是重排像素,不会产生新的插值伪影,在钢材灰度图上非常安全。如果担心图像上下方向在生产线上有语义,可以去掉 VerticalFlip。

4.3 train.py 里的训练循环

训练循环看起来和普通分割任务差别不大,但有几个容易被忽略的步骤。一个标准 epoch 的骨架如下:

for epoch in range(start_epoch, total_epochs): model.train() for images, masks in train_loader: images = images.cuda(non_blocking=True) masks = masks.cuda(non_blocking=True) logits = model(images) probas = torch.softmax(logits, dim=1) lovasz_loss = lovasz_softmax(probas, masks) ce_loss = nn.CrossEntropyLoss()(logits, masks) loss = 0.7 * lovasz_loss + 0.3 * ce_loss optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() val_miou = validate(model, val_loader)

clip_grad_norm_是为了处理偶发的大梯度。钢材图像里的极端像素值(比如强反光)会产生很大的 loss 尖刺,把梯度范数限制在 1.0 能避免一个 step 就把权重冲坏。validate函数每次返回验证集 mIoU,用它来判断是否保存新的 best checkpoint,而不是看训练 loss。

4.4 checkpoint 与断点续训

竞赛工程里 checkpoint 不能只存模型权重,至少要包含四样东西:

torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_miou': best_miou, }, 'checkpoint.pth')

恢复训练时把 epoch 和 optimizer 状态一并载入,才能继续原本的学习率调度。常见的错误是只 loadmodel_state_dict,导致恢复后 learning rate scheduler 从头开始,学习率瞬间跳高,模型几步之内就发散。

超参数上以这块竞赛数据为参考,常用范围如下:

参数参考值说明
输入尺寸256×256 / 256×512显存不足时优先缩短短边
batch size8-32与学习率同步调整
优化器AdamWlr 从 1e-4 开始
调度器CosineAnnealingLRT_max 设为总 epoch 数
训练轮数50-100验证 mIoU 连续 3 轮不升则早停
梯度裁剪max_norm=1.0防止异常像素值导致发散

5. eval_with_aug.py:TTA 推理、mIoU 统计与断点续训的细节

5.1 从 checkpoint 恢复模型状态

eval_with_aug.py 先要解决的是“怎么从训练状态切到推理状态”。关键是把权重载进来,再切换成评估模式:

ckpt = torch.load('checkpoint.pth', map_location='cuda:0') model.load_state_dict(ckpt['model_state_dict']) start_epoch = ckpt['epoch'] + 1 best_miou = ckpt['best_miou'] model.eval()

如果你的 checkpoint 里有优化器状态,并且想继续训练,就把optimizer.load_state_dict(ckpt['optimizer_state_dict'])也加上,再进入 train 循环;如果只是做验证和推理,则不要加载 optimizer。model.eval()必须放在循环外,它只会改 BatchNorm 和 Dropout,不会影响 Conv 权重,但漏掉它会让 mIoU 因为 BatchNorm 统计量不一致而掉 1-3 个点。

5.2 多尺度 TTA 的叠加

TTA 的思路是让模型在推理时看到更多“视角”,然后把概率平均。一个常用组合是三种尺度加水平翻转:

import torch.nn.functional as F def predict_tta(model, img, scales=(0.8, 1.0, 1.2)): H, W = img.shape[-2:] probs = [] for scale in scales: scaled = F.interpolate(img, scale_factor=scale, mode='bilinear', align_corners=False) logits = model(scaled) logits = F.interpolate(logits, size=(H, W), mode='bilinear', align_corners=False) p = torch.softmax(logits, dim=1) logits_f = model(torch.flip(scaled, dims=[3])) logits_f = F.interpolate(logits_f, size=(H, W), mode='bilinear', align_corners=False) p = p + torch.flip(torch.softmax(logits_f, dim=1), dims=[3]) probs.append(p) return torch.stack(probs).mean(dim=0)

每个尺度内部先各自处理翻转,再把所有尺度求平均,避免不同尺度权重失衡。多尺度 TTA 通常能带来 1-2 个百分点的 mIoU 提升,代价是推理时间乘以 6,所以要权衡是在线评分还是离线提交。

5.3 评估口径要统一

最后的 mIoU 统计要保证训练、验证、推理三个阶段对 mask 的插值方式一致。训练时代码用INTER_NEAREST缩 mask,验证时就别改用INTER_LINEAR,否则边缘像素的对齐方式不同,mIoU 的波动会掩盖真实的增益。统计指标时按类别分别计算 IoU,再求平均得到 mIoU,不要把背景类算进去。

最后提一个实际踩过的坑:把model.pth导出给后端推理时,模型输出 logits 的形状是(B, C, H, W),但如果脚本里有一次.permute(0, 2, 3, 1),最后 mask 会以(B, H, W, C)输出。在 torchvision 或 Flask 接口里,这两者混用时 argmax 的轴对不上,得到的 mask 就是错的。建议在 eval 脚本末尾加一行断言:

assert pred_mask.shape == (batch_size, height, width), "mask 形状必须是 B,H,W"

这样就算之后有人改动了前向分支,提交前也能被挡在门口。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 15:37:29

企业三大核心痛点解析与解决方案

1. 行业痛点深度剖析最近在和几位不同领域的企业主交流时,发现三个反复被提及的共性问题。这些问题看似简单,实则直击行业发展的核心痛点。作为从业十余年的行业观察者,我想结合具体案例,拆解这些"重灾区"背后的深层原因…

作者头像 李华
网站建设 2026/9/16 15:36:18

SkyPilot 快速上手:用 PyTorch DDP 在云端启动 minGPT 分布式训练

SkyPilot 快速上手:用 PyTorch DDP 在云端启动 minGPT 分布式训练 【免费下载链接】skypilot The AI Compute Platform for frontier teams. SkyPilot turns fragmented AI compute into one AI supercomputer, so frontier AI teams build custom intelligence fas…

作者头像 李华
网站建设 2026/9/16 15:35:18

tmux终端复用器实战:从防断线到高效工作流

用过终端的人,应该都有过这种经历:远程连了台服务器,部署跑了一半,网络一抖,SSH一断,整个任务跟着终端一起没了。那个瞬间,悔恨、无奈、想砸电脑的情绪交织在一起,最后只能老老实实重…

作者头像 李华
网站建设 2026/9/16 15:33:01

Python命令行参数类型管理实战指南

1. 为什么需要参数类型管理在Python命令行工具开发中,参数解析是每个开发者都要面对的基础问题。记得我第一次写命令行工具时,处理用户输入的各种参数格式简直让人抓狂 - 数字被当成字符串、文件路径需要手动验证、布尔值判断写了一大堆if...else。直到深…

作者头像 李华