使用UNet模型训练舌苔语义分割数据集,步骤:安装依赖、准备数据集、配置UNet模型、训练和评估模型、以及构建GUI应用程序来展示分割结果。
文章目录
- 使用UNet模型训练舌苔语义分割数据集,步骤:安装依赖、准备数据集、配置UNet模型、训练和评估模型、以及构建GUI应用程序来展示分割结果。
- 1. 安装依赖
- 2. 数据集准备
- 3. 配置UNet模型
- 4. 训练模型
- 5. 构建GUI应用程序
以下文字及代码仅供参考。
舌苔语义分割数据集
2460张数据(jpg)和mask掩码(png),
6类 红色舌苔厚腻 白色舌苔厚腻 黑色舌苔 白霉舌苔 紫色舌苔 红色舌苔黄腻厚
1
1
1
使用UNet模型训练舌苔语义分割数据集,步骤:安装依赖、准备数据集、配置UNet模型、训练和评估模型、以及构建GUI应用程序来展示分割结果。
1. 安装依赖
首先确保你的环境中已经安装了必要的库:
pipinstalltorch torchvision torchaudio opencv-python albumentations matplotlib tqdm2. 数据集准备
假设你已经有了一个标注好的数据集,包括图像及其对应的掩码文件(.png)。组织你的数据集如下:
tongue_dataset/ ├── images/ │ ├── img1.jpg │ ├── img2.jpg │ └── ... └── masks/ ├── mask1.png ├── mask2.png └── ...每个类别在mask中用不同的灰度值表示(如0代表背景,1代表红色舌苔厚腻等)。
创建一个自定义的PyTorch Dataset类来加载数据:
importosfromtorch.utils.dataimportDatasetimportcv2importnumpyasnpfromtorchvisionimporttransformsclassTongueDataset(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=cv2.imread(img_path)image=cv2.cvtColor(image,cv2.COLOR_BGR2RGB)mask=cv2.imread(mask_path,cv2.IMREAD_GRAYSCALE)ifself.transform:augmentations=self.transform(image=image,mask=mask)image=augmentations["image"]mask=augmentations["mask"]# 将mask转换为long类型,并保证是6类+背景(共7类)mask=mask.long()mask[mask>6]=0# 如果有超出范围的像素,设为背景returnimage,mask3. 配置UNet模型
这里我们采用一个基本的UNet架构:
importtorch.nnasnnimporttorchclassUNet(nn.Module):def__init__(self,n_channels,n_classes):super(UNet,self).__init__()# 这里省略了UNet的具体实现,你可以使用现成的UNet实现或者自己定义passdefforward(self,x):# 假设此处为UNet的前向传播过程pass# 初始化模型model=UNet(n_channels=3,n_classes=7)# 3通道输入,7类输出(6种舌苔+背景)4. 训练模型
定义损失函数和优化器,并开始训练:
fromtorch.utils.dataimportDataLoaderimporttorch.optimasoptim# 准备数据集和数据加载器transform=transforms.Compose([# 添加你需要的数据增强操作])dataset=TongueDataset(img_dir='path/to/images/',mask_dir='path/to/masks/',transform=transform)dataloader=DataLoader(dataset,batch_size=4,shuffle=True)# 损失函数和优化器criterion=nn.CrossEntropyLoss()optimizer=optim.Adam(model.parameters(),lr=0.001)# 训练循环forepochinrange(epochs):model.train()running_loss=0.0forimages,masksindataloader:optimizer.zero_grad()outputs=model(images)loss=criterion(outputs,masks)loss.backward()optimizer.step()running_loss+=loss.item()*images.size(0)print(f'Epoch{epoch+1}, Loss:{running_loss/len(dataloader.dataset)}')5. 构建GUI应用程序
接下来,我们将构建一个简单的PyQt5 GUI应用程序来展示UNet的分割结果。
importsysfromPyQt5.QtWidgetsimportQApplication,QLabel,QVBoxLayout,QWidget,QPushButton,QFileDialogfromPyQt5.QtGuiimportQPixmap,QImageimportcv2importnumpyasnpfromPILimportImageclassAppDemo(QWidget):def__init__(self):super().__init__()self.setWindowTitle('Tongue Segmentation')self.setGeometry(100,100,800,600)self.image_label=QLabel(self)self.button=QPushButton("Load Image",self)self.button.clicked.connect(self.load_image)vbox=QVBoxLayout()vbox.addWidget(self.image_label)vbox.addWidget(self.button)self.setLayout(vbox)# 加载已训练的UNet模型self.model=torch.load('path/to/your/trained_model.pth')self.model.eval()defload_image(self):fname,_=QFileDialog.getOpenFileName(self,'Open file','',"Image files (*.jpg *.png)")iffname:self.show_image(fname)defshow_image(self,image_path):image=cv2.imread(image_path)image_tensor=transforms.ToTensor()(image).unsqueeze(0)withtorch.no_grad():prediction=self.model(image_tensor)_,preds=torch.max(prediction,dim=1)pred_mask=preds.squeeze().cpu().numpy()pred_mask=Image.fromarray(pred_mask.astype(np.uint8),mode='P')pred_mask.putpalette([0,0,0,255,0,0,0,255,0,0,0,255,255,255,0,128,0,128])# 根据需要调整颜色板pred_mask=pred_mask.convert('RGB')height,width,channel=pred_mask.shape bytes_per_line=3*width q_img=QImage(pred_mask.data,width,height,bytes_per_line,QImage.Format_RGB888)pixmap=QPixmap.fromImage(q_img)self.image_label.setPixmap(pixmap)if__name__=='__main__':app=QApplication(sys.argv)demo=AppDemo()demo.show()sys.exit(app.exec_())请根据实际情况调整上述代码中的路径、参数和逻辑。