
SegFormer 语义分割训练实战自定义 Mask 数据集与预测可视化这篇教程根据我复现 SegFormer 自定义分割训练流程时整理重点演示环境安装、Mask 数据集加载、模型微调、测试评估和预测结果可视化。本文整理自我的学习和项目复现过程尽量按实操顺序保留 notebook 的关键步骤同时把数据集获取方式调整为适合中文教程发布的写法。本文会重点跑通以下流程安装 PyTorch Lightning 和 Transformers 依赖从数据集后台获取分割数据封装语义分割数据集微调 SegFormer 模型可视化预测 Mask 和原图叠加效果如果你正在系统学习目标检测、实例分割、OCR、多目标跟踪或视觉大模型建议收藏本文配套 notebook、示例图片和运行环境说明后续会继续整理。如果环境配置卡住可以在评论区说明具体报错。 文章目录SegFormer 语义分割训练实战自定义 Mask 数据集与预测可视化⚙️ 环境准备 从数据集后台获取分割数据 封装语义分割数据集 定义 SegFormer 微调模块 构建数据加载器️ 开始训练 测试模型️ 预测结果可视化 小结 同系列教程汇总⚙️ 环境准备先安装训练所需依赖并导入后续会用到的深度学习、评估和可视化模块。!pip install-q pytorch-lightning2.4.0!pip install-q transformers4.46.2datasets2.21.0importpytorch_lightningasplfrompytorch_lightning.callbacks.early_stoppingimportEarlyStoppingfrompytorch_lightning.callbacks.model_checkpointimportModelCheckpointfrompytorch_lightning.loggersimportCSVLoggerfromtransformersimportSegformerFeatureExtractor,SegformerForSemanticSegmentationfromdatasetsimportload_metricimporttorchfromtorchimportnnfromtorch.utils.dataimportDataset,DataLoaderimportosfromPILimportImageimportnumpyasnpimportrandom 从数据集后台获取分割数据从数据集后台导出分割数据后修改路径变量即可接入自己的数据。fromtypesimportSimpleNamespace# 从数据集后台下载 语义分割 格式数据集后修改 DATASET_DIR 指向解压目录。DATASET_DIR/content/dataset# 修改为数据集后台导出的数据集目录datasetSimpleNamespace(locationDATASET_DIR,version1,namecustom-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_dirroot_dir self.feature_extractorfeature_extractor self.classes_csv_fileos.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.jpginf]mask_file_names[fforfinos.listdir(self.root_dir)if.pnginf]self.imagessorted(image_file_names)self.maskssorted(mask_file_names)def__len__(self):returnlen(self.images)def__getitem__(self,idx):imageImage.open(os.path.join(self.root_dir,self.images[idx]))segmentation_mapImage.open(os.path.join(self.root_dir,self.masks[idx]))# randomly crop pad both image and segmentation map to same sizeencoded_inputsself.feature_extractor(image,segmentation_map,return_tensorspt)fork,vinencoded_inputs.items():encoded_inputs[k].squeeze_()# remove batch dimensionreturnencoded_inputs 定义 SegFormer 微调模块LightningModule 中集中处理训练、验证、测试和指标计算逻辑。classSegformerFinetuner(pl.LightningModule):def__init__(self,id2label,train_dataloaderNone,val_dataloaderNone,test_dataloaderNone,metrics_interval100):super(SegformerFinetuner,self).__init__()self.id2labelid2label self.metrics_intervalmetrics_interval self.train_dltrain_dataloader self.val_dlval_dataloader self.test_dltest_dataloader self.num_classeslen(id2label.keys())self.label2id{v:kfork,vinself.id2label.items()}self.modelSegformerForSemanticSegmentation.from_pretrained(nvidia/segformer-b0-finetuned-ade-512-512,return_dictFalse,num_labelsself.num_classes,id2labelself.id2label,label2idself.label2id,ignore_mismatched_sizesTrue,)self.train_mean_iouload_metric(mean_iou)self.val_mean_iouload_metric(mean_iou)self.test_mean_iouload_metric(mean_iou)self.validation_step_outputs[]defforward(self,images,masks):outputsself.model(pixel_valuesimages,labelsmasks)returnoutputsdeftraining_step(self,batch,batch_nb):images,masksbatch[pixel_values],batch[labels]outputsself(images,masks)loss,logitsoutputs[0],outputs[1]upsampled_logitsnn.functional.interpolate(logits,sizemasks.shape[-2:],modebilinear,align_cornersFalse)predictedupsampled_logits.argmax(dim1)self.train_mean_iou.add_batch(predictionspredicted.detach().cpu().numpy(),referencesmasks.detach().cpu().numpy())ifbatch_nb%self.metrics_interval0:metricsself.train_mean_iou.compute(num_labelsself.num_classes,ignore_index255,reduce_labelsFalse,)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,masksbatch[pixel_values],batch[labels]outputsself(images,masks)loss,logitsoutputs[0],outputs[1]upsampled_logitsnn.functional.interpolate(logits,sizemasks.shape[-2:],modebilinear,align_cornersFalse)predictedupsampled_logits.argmax(dim1)self.val_mean_iou.add_batch(predictionspredicted.detach().cpu().numpy(),referencesmasks.detach().cpu().numpy())self.validation_step_outputs.append({val_loss:loss})return({val_loss:loss})defon_validation_epoch_end(self):metricsself.val_mean_iou.compute(num_labelsself.num_classes,ignore_index255,reduce_labelsFalse,)avg_val_losstorch.stack([x[val_loss]forxinself.validation_step_outputs]).mean()val_mean_ioumetrics[mean_iou]val_mean_accuracymetrics[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,masksbatch[pixel_values],batch[labels]outputsself(images,masks)loss,logitsoutputs[0],outputs[1]upsampled_logitsnn.functional.interpolate(logits,sizemasks.shape[-2:],modebilinear,align_cornersFalse)predictedupsampled_logits.argmax(dim1)self.test_mean_iou.add_batch(predictionspredicted.detach().cpu().numpy(),referencesmasks.detach().cpu().numpy())return({test_loss:loss})deftest_epoch_end(self,outputs):metricsself.test_mean_iou.compute(num_labelsself.num_classes,ignore_index255,reduce_labelsFalse,)avg_test_losstorch.stack([x[test_loss]forxinoutputs]).mean()test_mean_ioumetrics[mean_iou]test_mean_accuracymetrics[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],lr2e-05,eps1e-08)deftrain_dataloader(self):returnself.train_dldefval_dataloader(self):returnself.val_dldeftest_dataloader(self):returnself.test_dl 构建数据加载器设置特征提取器、标签映射和 train/valid/test 三组数据加载器。feature_extractorSegformerFeatureExtractor.from_pretrained(nvidia/segformer-b0-finetuned-ade-512-512)feature_extractor.do_reduce_labelsFalsefeature_extractor.size128train_datasetSemanticSegmentationDataset(f{dataset.location}/train/,feature_extractor)val_datasetSemanticSegmentationDataset(f{dataset.location}/valid/,feature_extractor)test_datasetSemanticSegmentationDataset(f{dataset.location}/test/,feature_extractor)batch_size8num_workers2train_dataloaderDataLoader(train_dataset,batch_sizebatch_size,shuffleTrue,num_workersnum_workers)val_dataloaderDataLoader(val_dataset,batch_sizebatch_size,num_workersnum_workers)test_dataloaderDataLoader(test_dataset,batch_sizebatch_size,num_workersnum_workers)segformer_finetunerSegformerFinetuner(train_dataset.id2label,train_dataloadertrain_dataloader,val_dataloaderval_dataloader,test_dataloadertest_dataloader,metrics_interval10,)️ 开始训练配置早停、模型保存和 Trainer 后启动训练。early_stop_callbackEarlyStopping(monitorval_loss,min_delta0.00,patience10,verboseFalse,modemin,)checkpoint_callbackModelCheckpoint(save_top_k1,monitorval_loss)trainerpl.Trainer(callbacks[early_stop_callback,checkpoint_callback],max_epochs500,val_check_intervallen(train_dataloader),)trainer.fit(segformer_finetuner)%load_ext tensorboard%tensorboard--logdir lightning_logs/ 测试模型加载最佳权重在测试集上查看最终表现。restrainer.test(ckpt_pathbest)️ 预测结果可视化将预测 Mask 转换为彩色图并与原图叠加直观看模型效果。color_map{0:(0,0,0),1:(255,0,0),}defprediction_to_vis(prediction):vis_shapeprediction.shape(3,)visnp.zeros(vis_shape)fori,cincolor_map.items():vis[predictioni]color_map[i]returnImage.fromarray(vis.astype(np.uint8))forbatchintest_dataloader:images,masksbatch[pixel_values],batch[labels]outputssegformer_finetuner.model(images,masks)loss,logitsoutputs[0],outputs[1]upsampled_logitsnn.functional.interpolate(logits,sizemasks.shape[-2:],modebilinear,align_cornersFalse)predicted_maskupsampled_logits.argmax(dim1).cpu().numpy()masksmasks.cpu().numpy()n_plots4frommatplotlibimportpyplotasplt f,axarrplt.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_idx0input_image_fileos.path.join(test_dataset.root_dir,test_dataset.images[test_idx])input_imageImage.open(input_image_file)test_batchtest_dataset[test_idx]images,maskstest_batch[pixel_values],test_batch[labels]imagestorch.unsqueeze(images,0)maskstorch.unsqueeze(masks,0)outputssegformer_finetuner.model(images,masks)loss,logitsoutputs[0],outputs[1]upsampled_logitsnn.functional.interpolate(logits,sizemasks.shape[-2:],modebilinear,align_cornersFalse)predicted_maskupsampled_logits.argmax(dim1).cpu().numpy()maskprediction_to_vis(np.squeeze(masks))maskmask.resize(input_image.size)maskmask.convert(RGBA)input_imageinput_image.convert(RGBA)overlay_imgImage.blend(input_image,mask,0.5)overlay_img 小结这篇教程完整整理了SegFormer 语义分割训练的核心复现流程。实际操作时建议先确认 GPU、依赖版本、数据集路径和模型权重路径再逐段运行 notebook。后续我会继续按源项目顺序整理同系列中的目标检测、实例分割、OCR、多目标跟踪和视觉大模型教程。 同系列教程汇总Google Gemini 3.5 Flash 零样本目标检测教程从提示词到可视化结果GLM-OCR 文档识别实战教程从验证码、公式到车牌 OCRRF-DETR ByteTrack 多目标跟踪实战教程从命令行到 Python 视频轨迹可视化SAM 3 图像分割实战教程文本、框和点提示的多种分割方式SegFormer 语义分割训练实战自定义 Mask 数据集与预测可视化-本文