SegFormer 语义分割训练实战:自定义 Mask 数据集与预测可视化
2026/7/21 23:10:12 网站建设 项目流程

SegFormer 语义分割训练实战:自定义 Mask 数据集与预测可视化


这篇教程根据我复现 SegFormer 自定义分割训练流程时整理,重点演示环境安装、Mask 数据集加载、模型微调、测试评估和预测结果可视化。

本文整理自我的学习和项目复现过程,尽量按实操顺序保留 notebook 的关键步骤,同时把数据集获取方式调整为适合中文教程发布的写法。

本文会重点跑通以下流程:

  • 安装 PyTorch Lightning 和 Transformers 依赖
  • 从数据集后台获取分割数据
  • 封装语义分割数据集
  • 微调 SegFormer 模型
  • 可视化预测 Mask 和原图叠加效果

如果你正在系统学习目标检测、实例分割、OCR、多目标跟踪或视觉大模型,建议收藏本文;配套 notebook、示例图片和运行环境说明后续会继续整理。如果环境配置卡住,可以在评论区说明具体报错。

📚 文章目录

  • SegFormer 语义分割训练实战:自定义 Mask 数据集与预测可视化
    • ⚙️ 环境准备
    • 📦 从数据集后台获取分割数据
    • 🧱 封装语义分割数据集
    • 🧠 定义 SegFormer 微调模块
    • 🔧 构建数据加载器
    • 🏋️ 开始训练
    • 📏 测试模型
    • 🖼️ 预测结果可视化
    • 📌 小结
    • 📚 同系列教程汇总

⚙️ 环境准备

先安装训练所需依赖,并导入后续会用到的深度学习、评估和可视化模块。

!pip install-q pytorch-lightning==2.4.0
!pip install-q transformers==4.46.2datasets==2.21.0
importpytorch_lightningasplfrompytorch_lightning.callbacks.early_stoppingimportEarlyStoppingfrompytorch_lightning.callbacks.model_checkpointimportModelCheckpointfrompytorch_lightning.loggersimportCSVLoggerfromtransformersimportSegformerFeatureExtractor,SegformerForSemanticSegmentationfromdatasetsimportload_metricimporttorchfromtorchimportnnfromtorch.utils.dataimportDataset,DataLoaderimportosfromPILimportImageimportnumpyasnpimportrandom

📦 从数据集后台获取分割数据

从数据集后台导出分割数据后,修改路径变量即可接入自己的数据。

fromtypesimportSimpleNamespace# 从数据集后台下载 语义分割 格式数据集后,修改 DATASET_DIR 指向解压目录。DATASET_DIR="/content/dataset"# 修改为数据集后台导出的数据集目录dataset=SimpleNamespace(location=DATASET_DIR,version="1",name="custom-dataset")

🧱 封装语义分割数据集

这里把图像与 Mask 封装成 PyTorch Dataset,方便后续训练流程统一读取。

classSemanticSegmentationDataset(Dataset):"""图像语义分割数据集。"""def__init__(self,root_dir,feature_extractor):""" Args: root_dir (string): Root directory of the dataset containing the images + annotations. feature_extractor (SegFormerFeatureExtractor): feature extractor to prepare images + segmentation maps. train (bool): Whether to load "training" or "validation" images + annotations. """self.root_dir=root_dir self.feature_extractor=feature_extractor self.classes_csv_file=os.path.join(self.root_dir,"_classes.csv")withopen(self.classes_csv_file,'r')asfid:data=[l.split(',')fori,linenumerate(fid)ifi!=0]self.id2label={x[0]:x[1]forxindata}image_file_names=[fforfinos.listdir(self.root_dir)if'.jpg'inf]mask_file_names=[fforfinos.listdir(self.root_dir)if'.png'inf]self.images=sorted(image_file_names)self.masks=sorted(mask_file_names)def__len__(self):returnlen(self.images)def__getitem__(self,idx):image=Image.open(os.path.join(self.root_dir,self.images[idx]))segmentation_map=Image.open(os.path.join(self.root_dir,self.masks[idx]))# randomly crop + pad both image and segmentation map to same sizeencoded_inputs=self.feature_extractor(image,segmentation_map,return_tensors="pt")fork,vinencoded_inputs.items():encoded_inputs[k].squeeze_()# remove batch dimensionreturnencoded_inputs

🧠 定义 SegFormer 微调模块

LightningModule 中集中处理训练、验证、测试和指标计算逻辑。

classSegformerFinetuner(pl.LightningModule):def__init__(self,id2label,train_dataloader=None,val_dataloader=None,test_dataloader=None,metrics_interval=100):super(SegformerFinetuner,self).__init__()self.id2label=id2label self.metrics_interval=metrics_interval self.train_dl=train_dataloader self.val_dl=val_dataloader self.test_dl=test_dataloader self.num_classes=len(id2label.keys())self.label2id={v:kfork,vinself.id2label.items()}self.model=SegformerForSemanticSegmentation.from_pretrained("nvidia/segformer-b0-finetuned-ade-512-512",return_dict=False,num_labels=self.num_classes,id2label=self.id2label,label2id=self.label2id,ignore_mismatched_sizes=True,)self.train_mean_iou=load_metric("mean_iou")self.val_mean_iou=load_metric("mean_iou")self.test_mean_iou=load_metric("mean_iou")self.validation_step_outputs=[]defforward(self,images,masks):outputs=self.model(pixel_values=images,labels=masks)returnoutputsdeftraining_step(self,batch,batch_nb):images,masks=batch['pixel_values'],batch['labels']outputs=self(images,masks)loss,logits=outputs[0],outputs[1]upsampled_logits=nn.functional.interpolate(logits,size=masks.shape[-2:],mode="bilinear",align_corners=False)predicted=upsampled_logits.argmax(dim=1)self.train_mean_iou.add_batch(predictions=predicted.detach().cpu().numpy(),references=masks.detach().cpu().numpy())ifbatch_nb%self.metrics_interval==0:metrics=self.train_mean_iou.compute(num_labels=self.num_classes,ignore_index=255,reduce_labels=False,)metrics={'loss':loss,"mean_iou":metrics["mean_iou"],"mean_accuracy":metrics["mean_accuracy"]}fork,vinmetrics.items():self.log(k,v)return(metrics)else:return({'loss':loss})defvalidation_step(self,batch,batch_nb):images,masks=batch['pixel_values'],batch['labels']outputs=self(images,masks)loss,logits=outputs[0],outputs[1]upsampled_logits=nn.functional.interpolate(logits,size=masks.shape[-2:],mode="bilinear",align_corners=False)predicted=upsampled_logits.argmax(dim=1)self.val_mean_iou.add_batch(predictions=predicted.detach().cpu().numpy(),references=masks.detach().cpu().numpy())self.validation_step_outputs.append({'val_loss':loss})return({'val_loss':loss})defon_validation_epoch_end(self):metrics=self.val_mean_iou.compute(num_labels=self.num_classes,ignore_index=255,reduce_labels=False,)avg_val_loss=torch.stack([x["val_loss"]forxinself.validation_step_outputs]).mean()val_mean_iou=metrics["mean_iou"]val_mean_accuracy=metrics["mean_accuracy"]metrics={"val_loss":avg_val_loss,"val_mean_iou":val_mean_iou,"val_mean_accuracy":val_mean_accuracy}fork,vinmetrics.items():self.log(k,v)self.validation_step_outputs.clear()returnmetricsdeftest_step(self,batch,batch_nb):images,masks=batch['pixel_values'],batch['labels']outputs=self(images,masks)loss,logits=outputs[0],outputs[1]upsampled_logits=nn.functional.interpolate(logits,size=masks.shape[-2:],mode="bilinear",align_corners=False)predicted=upsampled_logits.argmax(dim=1)self.test_mean_iou.add_batch(predictions=predicted.detach().cpu().numpy(),references=masks.detach().cpu().numpy())return({'test_loss':loss})deftest_epoch_end(self,outputs):metrics=self.test_mean_iou.compute(num_labels=self.num_classes,ignore_index=255,reduce_labels=False,)avg_test_loss=torch.stack([x["test_loss"]forxinoutputs]).mean()test_mean_iou=metrics["mean_iou"]test_mean_accuracy=metrics["mean_accuracy"]metrics={"test_loss":avg_test_loss,"test_mean_iou":test_mean_iou,"test_mean_accuracy":test_mean_accuracy}fork,vinmetrics.items():self.log(k,v)returnmetricsdefconfigure_optimizers(self):returntorch.optim.Adam([pforpinself.parameters()ifp.requires_grad],lr=2e-05,eps=1e-08)deftrain_dataloader(self):returnself.train_dldefval_dataloader(self):returnself.val_dldeftest_dataloader(self):returnself.test_dl

🔧 构建数据加载器

设置特征提取器、标签映射和 train/valid/test 三组数据加载器。

feature_extractor=SegformerFeatureExtractor.from_pretrained("nvidia/segformer-b0-finetuned-ade-512-512")feature_extractor.do_reduce_labels=Falsefeature_extractor.size=128train_dataset=SemanticSegmentationDataset(f"{dataset.location}/train/",feature_extractor)val_dataset=SemanticSegmentationDataset(f"{dataset.location}/valid/",feature_extractor)test_dataset=SemanticSegmentationDataset(f"{dataset.location}/test/",feature_extractor)batch_size=8num_workers=2train_dataloader=DataLoader(train_dataset,batch_size=batch_size,shuffle=True,num_workers=num_workers)val_dataloader=DataLoader(val_dataset,batch_size=batch_size,num_workers=num_workers)test_dataloader=DataLoader(test_dataset,batch_size=batch_size,num_workers=num_workers)segformer_finetuner=SegformerFinetuner(train_dataset.id2label,train_dataloader=train_dataloader,val_dataloader=val_dataloader,test_dataloader=test_dataloader,metrics_interval=10,)

🏋️ 开始训练

配置早停、模型保存和 Trainer 后启动训练。

early_stop_callback=EarlyStopping(monitor="val_loss",min_delta=0.00,patience=10,verbose=False,mode="min",)checkpoint_callback=ModelCheckpoint(save_top_k=1,monitor="val_loss")trainer=pl.Trainer(callbacks=[early_stop_callback,checkpoint_callback],max_epochs=500,val_check_interval=len(train_dataloader),)trainer.fit(segformer_finetuner)
%load_ext tensorboard%tensorboard--logdir lightning_logs/

📏 测试模型

加载最佳权重,在测试集上查看最终表现。

res=trainer.test(ckpt_path="best")

🖼️ 预测结果可视化

将预测 Mask 转换为彩色图,并与原图叠加,直观看模型效果。

color_map={0:(0,0,0),1:(255,0,0),}defprediction_to_vis(prediction):vis_shape=prediction.shape+(3,)vis=np.zeros(vis_shape)fori,cincolor_map.items():vis[prediction==i]=color_map[i]returnImage.fromarray(vis.astype(np.uint8))forbatchintest_dataloader:images,masks=batch['pixel_values'],batch['labels']outputs=segformer_finetuner.model(images,masks)loss,logits=outputs[0],outputs[1]upsampled_logits=nn.functional.interpolate(logits,size=masks.shape[-2:],mode="bilinear",align_corners=False)predicted_mask=upsampled_logits.argmax(dim=1).cpu().numpy()masks=masks.cpu().numpy()n_plots=4frommatplotlibimportpyplotasplt f,axarr=plt.subplots(n_plots,2)f.set_figheight(15)f.set_figwidth(15)foriinrange(n_plots):axarr[i,0].imshow(prediction_to_vis(predicted_mask[i,:,:]))axarr[i,1].imshow(prediction_to_vis(masks[i,:,:]))

#Predict on a test image and overlay the mask on the original imagetest_idx=0input_image_file=os.path.join(test_dataset.root_dir,test_dataset.images[test_idx])input_image=Image.open(input_image_file)test_batch=test_dataset[test_idx]images,masks=test_batch['pixel_values'],test_batch['labels']images=torch.unsqueeze(images,0)masks=torch.unsqueeze(masks,0)outputs=segformer_finetuner.model(images,masks)loss,logits=outputs[0],outputs[1]upsampled_logits=nn.functional.interpolate(logits,size=masks.shape[-2:],mode="bilinear",align_corners=False)predicted_mask=upsampled_logits.argmax(dim=1).cpu().numpy()mask=prediction_to_vis(np.squeeze(masks))mask=mask.resize(input_image.size)mask=mask.convert("RGBA")input_image=input_image.convert("RGBA")overlay_img=Image.blend(input_image,mask,0.5)
overlay_img

📌 小结

这篇教程完整整理了SegFormer 语义分割训练的核心复现流程。实际操作时,建议先确认 GPU、依赖版本、数据集路径和模型权重路径,再逐段运行 notebook。

后续我会继续按源项目顺序整理同系列中的目标检测、实例分割、OCR、多目标跟踪和视觉大模型教程。

📚 同系列教程汇总

Google Gemini 3.5 Flash 零样本目标检测教程:从提示词到可视化结果

  • GLM-OCR 文档识别实战教程:从验证码、公式到车牌 OCR

  • RF-DETR + ByteTrack 多目标跟踪实战教程:从命令行到 Python 视频轨迹可视化

  • SAM 3 图像分割实战教程:文本、框和点提示的多种分割方式

  • SegFormer 语义分割训练实战:自定义 Mask 数据集与预测可视化-本文

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询