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 --name unet_env python=3.9
conda activate unet_env
4. 安装依赖项

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

pip install torch torchvision torchaudio
pip install opencv-python
pip install matplotlib
pip install scikit-image
pip install albumentations
pip install tqdm
pip install timm
pip install segmentation-models-pytorch

数据集准备

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

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

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

train_images: ./steel_defect_dataset/images/train
train_masks: ./steel_defect_dataset/masks/train
val_images: ./steel_defect_dataset/images/val
val_masks: ./steel_defect_dataset/masks/val
test_images: ./steel_defect_dataset/images/test
test_masks: ./steel_defect_dataset/masks/test

nc: 3
names: ['In', 'Pa', 'Sc']

使用U-Net训练模型

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

训练代码

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

import os
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import numpy as np
from torchvision import transforms

class SteelDefectDataset(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):
        return len(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,3

        if self.transform is not None:
            augmentations = self.transform(image=image, mask=mask)
            image = augmentations["image"]
            mask = augmentations["mask"]

        return image, mask

# 数据增强
import albumentations as A
from albumentations.pytorch import ToTensorV2

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)

接下来是训练部分:

import torch
import torch.nn as nn
from segmentation_models_pytorch import Unet
from tqdm import tqdm

# 初始化模型
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" if torch.cuda.is_available() else "cpu")
model.to(device)

# 训练循环
def train_model(model, train_loader, val_loader, loss_fn, optimizer, num_epochs=10):
    for epoch in range(num_epochs):
        model.train()
        loop = tqdm(train_loader)
        for batch_idx, (data, targets) in enumerate(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()
        with torch.no_grad():
            num_correct = 0
            num_pixels = 0
            dice_score = 0
            for data, targets in val_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)

推理代码

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

import cv2
from torchvision import transforms

# 加载训练好的模型
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)

with torch.no_grad():
    output = model(img_tensor)
    prediction = torch.argmax(output.squeeze(), dim=0).cpu().numpy()

# 显示结果
def label_to_color_image(label):
    colormap = np.array([[0, 0, 0], [255, 0, 0], [0, 255, 0], [0, 0, 255]])
    return colormap[label]

color_prediction = label_to_color_image(prediction)
cv2.imshow('Prediction', color_prediction)
cv2.waitKey(0)
cv2.destroyAllWindows()
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐