ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

水体区域分割数据集 海陆分割数据集 水与陆地分割检测 遥感水体分割数据集 原图(3841张)和对应的分割mask(3841张) 陆地上的水体区域进行图像分割

水体区域分割数据集 海陆分割数据集 水与陆地分割检测 遥感水体分割数据集 原图(3841张)和对应的分割mask(3841张) 陆地上的水体区域进行图像分割 水体区域分割数据集 海陆分割数据集 水与陆地分割检测 遥感水体分割数据集 原图3841张和对应的分割mask3841张 陆地上的水体区域进行图像分割解决卫星遥感水体图像分割任务:unettransunet实现包含遥感水体分割数据集Satellite_Images_of_Water_Bodies用于对陆地上的水体区域进行图像分割。包含原图3841张和对应的分割mask3841张附深度学习网络或改进网络实现分割。解决遥感图像的水体区域分割任务卫星遥感水体图像分割任务如何准备数据、训练模型、评估模型和可视化结果。使用UNet和TransUNet两种模型进行水体区域分割任务并提供完整的代码示例。1. 环境准备首先确保你已经安装了必要的库和工具。你可以使用以下命令安装所需的库pipinstalltorch torchvision pipinstallnumpy pipinstallpandas pipinstallmatplotlib pipinstallscikit-image pipinstallalbumentations pipinstalltqdm pipinstalleinops pipinstalltransformers2. 数据准备假设你的数据集目录结构如下Satellite_Images_of_Water_Bodies/ ├── images/ │ ├── 0001.jpg │ ├── 0002.jpg │ └── ... ├── masks/ │ ├── 0001.png │ ├── 0002.png │ └── ...每个图像文件和对应的标签文件都以相同的文件名命名例如0001.jpg和0001.png。3. 创建数据加载器创建一个数据加载器来读取图像和标签。我们使用PyTorch的Dataset和DataLoader类。importosimporttorchfromtorch.utils.dataimportDataset,DataLoaderfromPILimportImageimportnumpyasnpimportalbumentationsasAfromalbumentations.pytorchimportToTensorV2classWaterBodyDataset(Dataset):def__init__(self,image_dir,mask_dir,transformNone):self.image_dirimage_dir self.mask_dirmask_dir self.transformtransform self.imagesos.listdir(image_dir)def__len__(self):returnlen(self.images)def__getitem__(self,index):img_pathos.path.join(self.image_dir,self.images[index])mask_pathos.path.join(self.mask_dir,self.images[index].replace(.jpg,.png))imagenp.array(Image.open(img_path).convert(RGB))masknp.array(Image.open(mask_dir).convert(L),dtypenp.float32)mask[mask255.0]1.0ifself.transformisnotNone:augmentationsself.transform(imageimage,maskmask)imageaugmentations[image]maskaugmentations[mask]returnimage,mask# 数据增强transformA.Compose([A.Resize(height256,width256),A.Rotate(limit35,p1.0),A.HorizontalFlip(p0.5),A.VerticalFlip(p0.1),A.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225],max_pixel_value255.0,),ToTensorV2(),])# 创建数据加载器train_datasetWaterBodyDataset(image_dirSatellite_Images_of_Water_Bodies/images,mask_dirSatellite_Images_of_Water_Bodies/masks,transformtransform,)train_loaderDataLoader(train_dataset,batch_size16,shuffleTrue,num_workers2)4. 定义UNet模型UNet是一种改进的UNet模型通过引入更多的跳跃连接来提高性能。importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassDoubleConv(nn.Module):def__init__(self,in_channels,out_channels):super(DoubleConv,self).__init__()self.convnn.Sequential(nn.Conv2d(in_channels,out_channels,3,1,1,biasFalse),nn.BatchNorm2d(out_channels),nn.ReLU(inplaceTrue),nn.Conv2d(out_channels,out_channels,3,1,1,biasFalse),nn.BatchNorm2d(out_channels),nn.ReLU(inplaceTrue),)defforward(self,x):returnself.conv(x)classUNetPlusPlus(nn.Module):def__init__(self,in_channels3,out_channels1,features[32,64,128,256]):super(UNetPlusPlus,self).__init__()self.featuresfeatures self.encoder1DoubleConv(in_channels,features[0])self.encoder2DoubleConv(features[0],features[1])self.encoder3DoubleConv(features[1],features[2])self.encoder4DoubleConv(features[2],features[3])self.upconv3nn.ConvTranspose2d(features[3],features[2],kernel_size2,stride2)self.upconv2nn.ConvTranspose2d(features[2],features[1],kernel_size2,stride2)self.upconv1nn.ConvTranspose2d(features[1],features[0],kernel_size2,stride2)self.decoder3DoubleConv(features[3]features[2],features[2])self.decoder2DoubleConv(features[2]features[1],features[1])self.decoder1DoubleConv(features[1]features[0],features[0])self.final_convnn.Conv2d(features[0],out_channels,kernel_size1)defforward(self,x):enc1self.encoder1(x)enc2self.encoder2(F.max_pool2d(enc1,2))enc3self.encoder3(F.max_pool2d(enc2,2))enc4self.encoder4(F.max_pool2d(enc3,2))dec3self.upconv3(enc4)dec3torch.cat((dec3,enc3),dim1)dec3self.decoder3(dec3)dec2self.upconv2(dec3)dec2torch.cat((dec2,enc2),dim1)dec2self.decoder2(dec2)dec1self.upconv1(dec2)dec1torch.cat((dec1,enc1),dim1)dec1self.decoder1(dec1)returnself.final_conv(dec1)5. 定义TransUNet模型TransUNet结合了Transformer和UNet的优点适用于高分辨率图像的分割任务。importtorchimporttorch.nnasnnfromtransformersimportViTModelfromeinopsimportrearrangeclassTransUNet(nn.Module):def__init__(self,in_channels3,out_channels1,vit_namegoogle/vit-base-patch16-224-in21k):super(TransUNet,self).__init__()self.vitViTModel.from_pretrained(vit_name)self.upconv1nn.ConvTranspose2d(768,256,kernel_size2,stride2)self.upconv2nn.ConvTranspose2d(256,128,kernel_size2,stride2)self.upconv3nn.ConvTranspose2d(128,64,kernel_size2,stride2)self.decoder1DoubleConv(768256,256)self.decoder2DoubleConv(256128,128)self.decoder3DoubleConv(12864,64)self.final_convnn.Conv2d(64,out_channels,kernel_size1)defforward(self,x):xself.vit(pixel_valuesx)[last_hidden_state]xrearrange(x,b (h w) c - b c h w,h14,w14)dec1self.upconv1(x)dec1self.decoder1(dec1)dec2self.upconv2(dec1)dec2self.decoder2(dec2)dec3self.upconv3(dec2)dec3self.decoder3(dec3)returnself.final_conv(dec3)6. 训练模型定义训练和验证函数。importtorch.optimasoptimfromtqdmimporttqdm devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)deftrain_fn(loader,model,optimizer,loss_fn,scaler):looptqdm(loader)forbatch_idx,(data,targets)inenumerate(loop):datadata.to(device)targetstargets.unsqueeze(1).to(device)# Forwardwithtorch.cuda.amp.autocast():predictionsmodel(data)lossloss_fn(predictions,targets)# Backwardoptimizer.zero_grad()scaler.scale(loss).backward()scaler.step(optimizer)scaler.update()# Update tqdm looploop.set_postfix(lossloss.item())defcheck_accuracy(loader,model,devicecuda):num_correct0num_pixels0dice_score0model.eval()withtorch.no_grad():forx,yinloader:xx.to(device)yy.to(device).unsqueeze(1)predstorch.sigmoid(model(x))preds(preds0.5).float()num_correct(predsy).sum()num_pixelstorch.numel(preds)dice_score(2*(preds*y).sum())/((predsy).sum()1e-8)print(fGot{num_correct}/{num_pixels}with acc{num_correct/num_pixels*100:.2f})print(fDice score:{dice_score/len(loader)})model.train()defmain():modelUNetPlusPlus(in_channels3,out_channels1).to(device)# 或者使用 TransUNet# model TransUNet(in_channels3, out_channels1).to(device)loss_fnnn.BCEWithLogitsLoss()optimizeroptim.Adam(model.parameters(),lr1e-4)scalertorch.cuda.amp.GradScaler()forepochinrange(100):# Number of epochstrain_fn(train_loader,model,optimizer,loss_fn,scaler)check_accuracy(train_loader,model,devicedevice)# Save modelcheckpoint{state_dict:model.state_dict(),optimizer:optimizer.state_dict(),}torch.save(checkpoint,fwater_body_segmentation_checkpoint.pth.tar)if__name____main__:main()
返回列表