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-smi2. 安装Anaconda
访问 Anaconda官网 下载并安装适合您操作系统的版本。
3. 创建Python虚拟环境
打开终端或Anaconda Prompt,然后输入以下命令来创建并激活新的Python环境:
conda create--nameunet_envpython=3.9conda activate unet_env4. 安装依赖项
在激活的环境中运行以下命令以安装必要的库:
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.yamldata.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()