news 2026/9/8 8:35:30

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

作者头像

张小明

前端开发工程师

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

水体区域分割数据集 海陆分割数据集 水与陆地分割检测 遥感水体分割数据集 原图(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 pipinstalltransformers

2. 数据准备

假设你的数据集目录结构如下:

Satellite_Images_of_Water_Bodies/ ├── images/ │ ├── 0001.jpg │ ├── 0002.jpg │ └── ... ├── masks/ │ ├── 0001.png │ ├── 0002.png │ └── ...

每个图像文件和对应的标签文件都以相同的文件名命名,例如0001.jpg0001.png

3. 创建数据加载器

创建一个数据加载器来读取图像和标签。我们使用PyTorch的DatasetDataLoader类。

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()
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/5 8:20:38

腾讯音乐秋招系统测试岗笔试全解析:考点、用例设计与时间分配

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/8 8:34:55

华为AI岗面试复盘:OD机试、Transformer与项目深挖全记录

2026年6月12号,我结束了华为AI岗的最后一轮面试。走出那栋楼的时候,我下意识把手机里存的机试草稿又翻了一遍,脑子里全是二叉树的递归、Transformer的KV Cache,以及简历里那个差点被面试官问穿的项目。这篇文章不是那种一句话概括…

作者头像 李华
网站建设 2026/9/7 16:10:06

掌阅前端笔试复盘:事件循环、Vue响应式与性能优化全解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/4 12:49:06

Excel高级筛选与表格结合:无需公式实现复杂多条件数据筛选

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/5 12:15:18

隐含波动率实战指南:用IV Rank与期限结构判断期权贵贱

隐含波动率是期权交易里绕不开的概念。接触期权一段时间后,你一定会发现:同样一张看涨期权,在标的价格涨跌幅度接近的时候,期权价格可能差得非常多。这个差异背后的主要变量,往往就是隐含波动率。 这篇文章要解决三个…

作者头像 李华
网站建设 2026/9/6 7:44:18

PSP游戏资源下载安全指南:警惕恶意压缩包与钓鱼陷阱

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华