水体区域分割数据集 海陆分割数据集 水与陆地分割检测 遥感水体分割数据集 原图(3841张)和对应的分割mask(3841张) 陆地上的水体区域进行图像分割
解决卫星遥感水体图像分割任务:
unet++,transunet实现
包含遥感水体分割数据集
(Satellite_Images_of_Water_Bodies)
用于对陆地上的水体区域进行图像分割。
包含原图(3841张)和对应的分割mask(3841张)
附深度学习网络或改进网络实现分割。
解决遥感图像的水体区域分割任务
卫星遥感水体图像分割任务,如何准备数据、训练模型、评估模型和可视化结果。使用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,transform=None):self.image_dir=image_dir self.mask_dir=mask_dir self.transform=transform self.images=os.listdir(image_dir)def__len__(self):returnlen(self.images)def__getitem__(self,index):img_path=os.path.join(self.image_dir,self.images[index])mask_path=os.path.join(self.mask_dir,self.images[index].replace('.jpg','.png'))image=np.array(Image.open(img_path).convert("RGB"))mask=np.array(Image.open(mask_dir).convert("L"),dtype=np.float32)mask[mask==255.0]=1.0ifself.transformisnotNone:augmentations=self.transform(image=image,mask=mask)image=augmentations['image']mask=augmentations['mask']returnimage,mask# 数据增强transform=A.Compose([A.Resize(height=256,width=256),A.Rotate(limit=35,p=1.0),A.HorizontalFlip(p=0.5),A.VerticalFlip(p=0.1),A.Normalize(mean=[0.485,0.456,0.406],std=[0.229,0.224,0.225],max_pixel_value=255.0,),ToTensorV2(),])# 创建数据加载器train_dataset=WaterBodyDataset(image_dir="Satellite_Images_of_Water_Bodies/images",mask_dir="Satellite_Images_of_Water_Bodies/masks",transform=transform,)train_loader=DataLoader(train_dataset,batch_size=16,shuffle=True,num_workers=2)4. 定义UNet++模型
UNet++是一种改进的UNet模型,通过引入更多的跳跃连接来提高性能。
importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassDoubleConv(nn.Module):def__init__(self,in_channels,out_channels):super(DoubleConv,self).__init__()self.conv=nn.Sequential(nn.Conv2d(in_channels,out_channels,3,1,1,bias=False),nn.BatchNorm2d(out_channels),nn.ReLU(inplace=True),nn.Conv2d(out_channels,out_channels,3,1,1,bias=False),nn.BatchNorm2d(out_channels),nn.ReLU(inplace=True),)defforward(self,x):returnself.conv(x)classUNetPlusPlus(nn.Module):def__init__(self,in_channels=3,out_channels=1,features=[32,64,128,256]):super(UNetPlusPlus,self).__init__()self.features=features self.encoder1=DoubleConv(in_channels,features[0])self.encoder2=DoubleConv(features[0],features[1])self.encoder3=DoubleConv(features[1],features[2])self.encoder4=DoubleConv(features[2],features[3])self.upconv3=nn.ConvTranspose2d(features[3],features[2],kernel_size=2,stride=2)self.upconv2=nn.ConvTranspose2d(features[2],features[1],kernel_size=2,stride=2)self.upconv1=nn.ConvTranspose2d(features[1],features[0],kernel_size=2,stride=2)self.decoder3=DoubleConv(features[3]+features[2],features[2])self.decoder2=DoubleConv(features[2]+features[1],features[1])self.decoder1=DoubleConv(features[1]+features[0],features[0])self.final_conv=nn.Conv2d(features[0],out_channels,kernel_size=1)defforward(self,x):enc1=self.encoder1(x)enc2=self.encoder2(F.max_pool2d(enc1,2))enc3=self.encoder3(F.max_pool2d(enc2,2))enc4=self.encoder4(F.max_pool2d(enc3,2))dec3=self.upconv3(enc4)dec3=torch.cat((dec3,enc3),dim=1)dec3=self.decoder3(dec3)dec2=self.upconv2(dec3)dec2=torch.cat((dec2,enc2),dim=1)dec2=self.decoder2(dec2)dec1=self.upconv1(dec2)dec1=torch.cat((dec1,enc1),dim=1)dec1=self.decoder1(dec1)returnself.final_conv(dec1)5. 定义TransUNet模型
TransUNet结合了Transformer和UNet的优点,适用于高分辨率图像的分割任务。
importtorchimporttorch.nnasnnfromtransformersimportViTModelfromeinopsimportrearrangeclassTransUNet(nn.Module):def__init__(self,in_channels=3,out_channels=1,vit_name='google/vit-base-patch16-224-in21k'):super(TransUNet,self).__init__()self.vit=ViTModel.from_pretrained(vit_name)self.upconv1=nn.ConvTranspose2d(768,256,kernel_size=2,stride=2)self.upconv2=nn.ConvTranspose2d(256,128,kernel_size=2,stride=2)self.upconv3=nn.ConvTranspose2d(128,64,kernel_size=2,stride=2)self.decoder1=DoubleConv(768+256,256)self.decoder2=DoubleConv(256+128,128)self.decoder3=DoubleConv(128+64,64)self.final_conv=nn.Conv2d(64,out_channels,kernel_size=1)defforward(self,x):x=self.vit(pixel_values=x)['last_hidden_state']x=rearrange(x,'b (h w) c -> b c h w',h=14,w=14)dec1=self.upconv1(x)dec1=self.decoder1(dec1)dec2=self.upconv2(dec1)dec2=self.decoder2(dec2)dec3=self.upconv3(dec2)dec3=self.decoder3(dec3)returnself.final_conv(dec3)6. 训练模型
定义训练和验证函数。
importtorch.optimasoptimfromtqdmimporttqdm device=torch.device("cuda"iftorch.cuda.is_available()else"cpu")deftrain_fn(loader,model,optimizer,loss_fn,scaler):loop=tqdm(loader)forbatch_idx,(data,targets)inenumerate(loop):data=data.to(device)targets=targets.unsqueeze(1).to(device)# Forwardwithtorch.cuda.amp.autocast():predictions=model(data)loss=loss_fn(predictions,targets)# Backwardoptimizer.zero_grad()scaler.scale(loss).backward()scaler.step(optimizer)scaler.update()# Update tqdm looploop.set_postfix(loss=loss.item())defcheck_accuracy(loader,model,device="cuda"):num_correct=0num_pixels=0dice_score=0model.eval()withtorch.no_grad():forx,yinloader:x=x.to(device)y=y.to(device).unsqueeze(1)preds=torch.sigmoid(model(x))preds=(preds>0.5).float()num_correct+=(preds==y).sum()num_pixels+=torch.numel(preds)dice_score+=(2*(preds*y).sum())/((preds+y).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(loader)}")model.train()defmain():model=UNetPlusPlus(in_channels=3,out_channels=1).to(device)# 或者使用 TransUNet# model = TransUNet(in_channels=3, out_channels=1).to(device)loss_fn=nn.BCEWithLogitsLoss()optimizer=optim.Adam(model.parameters(),lr=1e-4)scaler=torch.cuda.amp.GradScaler()forepochinrange(100):# Number of epochstrain_fn(train_loader,model,optimizer,loss_fn,scaler)check_accuracy(train_loader,model,device=device)# Save modelcheckpoint={"state_dict":model.state_dict(),"optimizer":optimizer.state_dict(),}torch.save(checkpoint,f"water_body_segmentation_checkpoint.pth.tar")if__name__=="__main__":main()