当前位置: 首页 > news >正文

U-Net模型进行训练钢材表面缺陷语义分割数据集 通过钢材缺陷分割数据集的权重模型,推理识别钢材分割

U-Net模型进行训练钢材表面缺陷语义分割数据集 通过钢材缺陷分割数据集的权重模型,推理识别钢材分割

文章目录

      • 环境搭建
        • 1. 安装CUDA驱动
        • 2. 安装Anaconda
        • 3. 创建Python虚拟环境
        • 4. 安装依赖项
      • 数据集准备
      • 使用U-Net训练模型
        • 训练代码
      • 推理代码

以下文字及代码仅供参考学习使用。

钢材表面缺陷语义分割数据集
4432张数据(jpg)和mask掩码(png),有颜色映射关系,另外有转换成coco格式(json)和yolo格式(txt)
三种缺陷类型(像素标签)
0为背景 1为夹杂物(In) 2为补丁(Pa) 3为划痕(Sc)

共4432张数据(jpg)和mask掩码(png),有颜色映射关系,另外有转换成coco格式(json)和yolo格式(txt)

含三种缺陷类型(像素标签)
0为背景 1为夹杂物(In) 2为补丁(Pa) 3为划痕(Sc)

mask标签颜色映射

为了使用U-Net模型对钢材表面缺陷进行语义分割,我们需要从环境搭建开始,到数据集准备、模型训练和推理。以下是详细的步骤指南。仅供参考学习使用

labelme查看

环境搭建

1. 安装CUDA驱动

确保您的系统已经安装了与GPU兼容的CUDA驱动版本。可以使用以下命令检查:

nvidia-smi
2. 安装Anaconda

访问 Anaconda官网 下载并安装适合您操作系统的版本。

3. 创建Python虚拟环境

打开终端或Anaconda Prompt,然后输入以下命令来创建并激活新的Python环境:

conda create--nameunet_envpython=3.9conda activate unet_env
4. 安装依赖项

在激活的环境中运行以下命令以安装必要的库:

pipinstalltorch torchvision torchaudio pipinstallopencv-python pipinstallmatplotlib pipinstallscikit-image pipinstallalbumentations pipinstalltqdm pipinstalltimm pipinstallsegmentation-models-pytorch

数据集准备

假设同学你的数据集按照如下结构组织:

steel_defect_dataset/ ├── images/ │ ├── train/ │ ├── val/ │ └── test/ ├── masks/ │ ├── train/ │ ├── val/ │ └── test/ └── data.yaml

data.yaml文件内容示例(请根据实际情况调整路径):

train_images:./steel_defect_dataset/images/traintrain_masks:./steel_defect_dataset/masks/trainval_images:./steel_defect_dataset/images/valval_masks:./steel_defect_dataset/masks/valtest_images:./steel_defect_dataset/images/testtest_masks:./steel_defect_dataset/masks/testnc:3names:['In','Pa','Sc']

使用U-Net训练模型

使用segmentation_models.pytorch库中的U-Net模型进行训练。的Python脚本示例,用于加载U-Net模型并使用提供的数据集进行训练。仅供参考学习使用。

训练代码

首先,编写一个数据加载器函数,用于加载图像和掩码,并应用必要的预处理。

importosfromtorch.utils.dataimportDataset,DataLoaderfromPILimportImageimportnumpyasnpfromtorchvisionimporttransformsclassSteelDefectDataset(Dataset):def__init__(self,img_dir,mask_dir,transform=None):self.img_dir=img_dir self.mask_dir=mask_dir self.transform=transform self.images=os.listdir(img_dir)def__len__(self):returnlen(self.images)def__getitem__(self,idx):img_path=os.path.join(self.img_dir,self.images[idx])mask_path=os.path.join(self.mask_dir,self.images[idx].replace('.jpg','.png'))image=np.array(Image.open(img_path).convert("RGB"))mask=np.array(Image.open(mask_path).convert("L"),dtype=np.float32)mask[mask==255.0]=1.0# 背景为0,其他类别为1,2,3ifself.transformisnotNone:augmentations=self.transform(image=image,mask=mask)image=augmentations["image"]mask=augmentations["mask"]returnimage,mask# 数据增强importalbumentationsasAfromalbumentations.pytorchimportToTensorV2 transform=A.Compose([A.Resize(height=256,width=256),A.Normalize(mean=[0.0,0.0,0.0],std=[1.0,1.0,1.0],max_pixel_value=255.0,),ToTensorV2(),],)train_ds=SteelDefectDataset(img_dir="path/to/train/images",mask_dir="path/to/train/masks",transform=transform,)val_ds=SteelDefectDataset(img_dir="path/to/val/images",mask_dir="path/to/val/masks",transform=transform,)train_loader=DataLoader(train_ds,batch_size=16,shuffle=True)val_loader=DataLoader(val_ds,batch_size=16,shuffle=False)

接下来是训练部分:

importtorchimporttorch.nnasnnfromsegmentation_models_pytorchimportUnetfromtqdmimporttqdm# 初始化模型model=Unet(encoder_name="resnet34",classes=3,activation=None)# 损失函数和优化器loss_fn=nn.CrossEntropyLoss()optimizer=torch.optim.Adam(model.parameters(),lr=1e-4)# 设备配置device=torch.device("cuda"iftorch.cuda.is_available()else"cpu")model.to(device)# 训练循环deftrain_model(model,train_loader,val_loader,loss_fn,optimizer,num_epochs=10):forepochinrange(num_epochs):model.train()loop=tqdm(train_loader)forbatch_idx,(data,targets)inenumerate(loop):data=data.to(device=device)targets=targets.long().to(device=device)# 前向传播predictions=model(data)loss=loss_fn(predictions,targets)# 反向传播和优化optimizer.zero_grad()loss.backward()optimizer.step()# 更新进度条loop.set_postfix(loss=loss.item())# 验证阶段model.eval()withtorch.no_grad():num_correct=0num_pixels=0dice_score=0fordata,targetsinval_loader:data=data.to(device)targets=targets.to(device).unsqueeze(1)predictions=torch.softmax(model(data),dim=1)preds=torch.argmax(predictions,dim=1).float()num_correct+=(preds==targets).sum()num_pixels+=torch.numel(preds)dice_score+=(2*(preds*targets).sum())/((preds+targets).sum()+1e-8)print(f"Got{num_correct}/{num_pixels}with acc{num_correct/num_pixels*100:.2f}")print(f"Dice score:{dice_score/len(val_loader)}")# 开始训练train_model(model,train_loader,val_loader,loss_fn,optimizer,num_epochs=10)

推理代码

训练完成后,您可以使用训练好的模型对新图片进行预测。以下是一个简单的例子:

importcv2fromtorchvisionimporttransforms# 加载训练好的模型model=Unet(encoder_name="resnet34",classes=3,activation=None)model.load_state_dict(torch.load('path/to/best_model.pth'))model.eval()# 图像预处理preprocess=transforms.Compose([transforms.ToPILImage(),transforms.Resize((256,256)),transforms.ToTensor(),transforms.Normalize(mean=[0.0,0.0,0.0],std=[1.0,1.0,1.0]),])# 对单张图片进行预测image_path='path/to/new/image.jpg'img=cv2.imread(image_path)img_tensor=preprocess(img).unsqueeze(0).to(device)withtorch.no_grad():output=model(img_tensor)prediction=torch.argmax(output.squeeze(),dim=0).cpu().numpy()# 显示结果deflabel_to_color_image(label):colormap=np.array([[0,0,0],[255,0,0],[0,255,0],[0,0,255]])returncolormap[label]color_prediction=label_to_color_image(prediction)cv2.imshow('Prediction',color_prediction)cv2.waitKey(0)cv2.destroyAllWindows()
http://www.cnnetsun.cn/news/1905420.html

相关文章:

  • 紫光同创PDS在线仿真避坑指南:手把手教你处理信号被优化的问题
  • 如何彻底卸载Microsoft Edge:EdgeRemover工具终极指南
  • 从CLIP到Qwen-VL-MoE:多模态域适应演进图谱(2018–2024关键论文脉络+工业落地成熟度雷达图)
  • 液态神经网络(LTCs)在连续时间控制中的可解释性设计与应用
  • AGI 的进化之路:从技术突破到伦理挑战
  • Cellpose-SAM细胞分割技术深度解析与实践指南
  • 别再死记硬背DDS概念了!用ROS2实战案例带你搞懂Topic、Service、Action的QoS调优
  • MM32 MCU烧录失败?5个常见硬件问题排查指南(附电路设计建议)
  • STM32F103RCT6三线SPI驱动ADS8866避坑指南(附完整代码)
  • 别再只盯着HA了!聊聊vSphere FT容错的真实应用场景与那些“不起眼”的限制
  • 网络安全术语解析:通用平台枚举CPE实战指南
  • 告别Termux折腾!在华为平板上用AidLux搭建Python开发环境,自带VSCode到底香不香?
  • 仿真系列专栏介绍
  • 傅里叶变换实战:如何用Python避免频谱分析中的泄露效应?
  • 野火串口调试助手PID协议详解:从数据包解析到弱函数重写,一步步打造你的专属上位机
  • 终极指南:如何用Coconut优雅解决Python 2/3跨版本兼容难题
  • [Python3高阶编程] - Waitress 源码剖析03: WSGI 服务器核心引擎 - server.py 解析
  • 如何在3分钟内搭建Sakura-13B-Galgame翻译API:免费离线日语游戏翻译终极指南
  • 3分钟搞定Windows UEFI启动画面:告别单调开机界面
  • 革命性AI工具gptcommit:让GPT-3为你自动编写完美的Git提交信息
  • YOLOv11、PyQt5、火灾烟雾检测 智慧火灾监测-YOLOv11火灾检测系统【YOLO火灾检测系统】智能预警,守护安全 火灾监测数据集的训练及应用
  • Vivado ILA实战:5分钟搞定FPGA信号抓取与波形分析(附常见问题排查)
  • 5大核心模块:重新定义英雄联盟游戏辅助体验
  • APK Installer:Windows平台安卓应用安装的完整解决方案
  • Kazumi番剧播放器:从零开始的完整使用指南
  • PPTist:基于Vue3+TypeScript的在线演示文稿创作平台终极指南
  • 华为路由器OSPF多区域配置详解:从零到实战一步到位
  • Video-subtitle-remover:AI驱动的视频硬字幕去除终极指南
  • Windows 10下Veins+SUMO+OMNeT++环境搭建全攻略(避坑指南)
  • 深度解析APK文件:Java开发者必备的apk-parser完全实战指南