Transformers 中 SegGPT 的上下文语义分割实战:从图像处理器到掩码后处理

发布时间:2026/9/8 21:46:02
Transformers 中 SegGPT 的上下文语义分割实战:从图像处理器到掩码后处理 Transformers 中 SegGPT 的上下文语义分割实战从图像处理器到掩码后处理【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformersSegGPT 是 Hugging Face Transformers 中实现上下文学习in-context learning图像分割的模型给定一张待分割图像、一张提示图像及其提示掩码模型即可通过单个 decoder-only Transformer 一次性生成分割掩码无需逐类训练。本文基于仓库中的官方模型文档 SegGPT 展开结合 模型实现、图像处理器 与 测试用例 的源码完整讲解SegGptConfig参数、提示掩码的两种输入格式、特征融合feature ensemble机制以及从原始pred_masks到最终语义分割图的后处理流程帮助读者可复现地跑通 one-shot 语义分割推理。SegGPT 的核心理念把分割当成按上下文填色SegGPT 出自论文《SegGPT: Segmenting Everything In Context》arXiv: 2304.03284Xinlong Wang 等人模型文档注明该模型于 2023-04-06 发布于 HF papers、2024-02-26 合入 Transformers。其论文摘要指出SegGPT 将各类分割任务统一为一个通用的上下文学习框架把不同形式的分割数据转换为相同格式的图像训练被形式化为一个上下文填色问题in-context coloring problem每个数据样本使用随机的颜色映射训练目标是根据上下文完成任务而不是依赖特定颜色。训练完成后它可以在图像或视频上执行任意图上下文推理任务如对象实例、stuff、部件、轮廓和文字分割并覆盖 few-shot 语义分割、视频目标分割、语义分割、全景分割等多种任务。官方文档中给出的代表结果为COCO-20 上 56.1 mIoUone-shot、FSS-1000 上 85.6 mIoU。从 Transformers 的实现看这套思想落在三段式结构上上下文拼接把提示图像与待分割图像沿高度方向拼接成一张高为两倍的伪图像输入同一个 ViT 编码器中间层特征采集在编码器的若干中间层默认第 5、11、17、23 层取出特征作为解码器输入RGB 空间解码解码器输出 3 通道的颜色图即直接在 RGB 像素空间预测掩码再通过调色板palette映射回类别索引。仓库内该模型的完整文件布局如下均为相对仓库根目录的路径文件作用configuration_seggpt.pySegGptConfig架构超参数定义modeling_seggpt.pySegGptModel编码器、SegGptForImageSegmentation含解码器与损失image_processing_seggpt.pySegGptImageProcessortorchvision 后端image_processing_pil_seggpt.pySegGptImageProcessorPil纯 PIL/numpy 后端convert_seggpt_to_hf.py将原始 Painter 仓库权重转换为 HF 格式test_modeling_seggpt.py、test_image_processing_seggpt.py模型与图像处理测试推荐加载的官方检查点为BAAI/seggpt-vit-large。SegGptConfig关键参数与默认值SegGptConfig继承自PreTrainedConfigmodel_type seggpt。以下参数与默认值均直接取自 configuration_seggpt.py参数默认值含义hidden_size1024Transformer 隐藏维度num_hidden_layers24编码器层数num_attention_heads16注意力头数hidden_actgelu激活函数image_size(896, 448)输入图像尺寸高度为提示与图像拼接后的总高patch_size16patch 大小mlp_dimNone回退为hidden_size * 4MLP 维度在__post_init__中若为None则置为hidden_size * 4pretrain_image_size224绝对位置编码的预训练尺寸用于双三次插值use_relative_position_embeddingsTrue注意力中是否使用分解式相对位置编码merge_index2提示特征与输入特征合并取平均的编码器层索引intermediate_hidden_state_indices(5, 11, 17, 23)供解码器使用的中间层索引decoder_hidden_size64解码器内部特征维度beta0.01SegGptLosssmooth-L1的正则化因子drop_path_rate0.1随机深度DropPath线性插值的最大比率qkv_biasTrueQKV 线性层是否带偏置配置类还带有一条结构性校验validate_architecturemerge_index必须小于min(intermediate_hidden_state_indices)即特征合并必须发生在第一个被采集的中间层之前。默认配置2 5满足该约束自定义配置时若把merge_index调到 12 以上会直接抛出ValueError。官方文档给出的最小示例from transformers import SegGptConfig, SegGptModel configuration SegGptConfig() model SegGptModel(configuration) configuration model.configSegGptImageProcessor三类输入与提示掩码的两种格式SegGptImageProcessortorchvision 后端与SegGptImageProcessorPilPIL/numpy 后端接口一致默认参数为size{height: 448, width: 448}、do_resize/do_rescale/do_normalizeTrue、归一化使用 ImageNet 均值与标准差image_processing_seggpt.py。preprocess接受三类输入preprocessimages待分割的目标图像prompt_images提示图像与该图像对应的参考图prompt_masks提示掩码。三类输入中至少给一个否则抛ValueError。处理结果分别为pixel_values、prompt_pixel_values、prompt_masks三个张量形状均为(batch_size, 3, H, W)。提示掩码的两种合法格式这是文档中最强调、也最容易出错的一点原文档 Tips 第 2 条prompt_masks既可以是分割图2D 类别索引图也可以是RGB 图像3D 掩码图。源码中通过do_convert_rgb开关区分两种分支_preprocess_image_like_inputs2D 分割图默认do_convert_rgbTrue处理器把每个类别索引通过调色板上色为 3 通道 RGB已是 RGB 的 3 通道掩码图必须传do_convert_rgbFalse处理器才会按 3 通道直接处理否则会因维度不匹配而报错。两个分支中掩码的缩放采样方式都被强制替换为PILImageResampling.NEARESTL218避免边界类别被双三次插值染色。test_image_processing_seggpt.py 中的test_mask_equivalenceL151-L161专门验证同一掩码走灰度分割图路径与走RGB do_convert_rgbFalse路径输出的prompt_masks张量完全相等。num_labels 与调色板palette文档 Tips 第 3 条强烈建议使用segmentation_maps做前后处理时传入num_labels不含背景类。其原理在 build_palettedef build_palette(num_labels: int) - list[tuple[int, int, int]]: base int(num_labels ** (1 / 3)) 1 margin 256 // base # class_idx 0 is the background which is mapped to black color_list [(0, 0, 0)] for location in range(num_labels): num_seq_r location // base**2 num_seq_g (location % base**2) // base num_seq_b location % base R 255 - num_seq_r * margin G 255 - num_seq_g * margin B 255 - num_seq_b * margin color_list.append((R, G, B)) return color_list该调色板把类别 0 固定为黑色背景其余类别按类立方体方式在 RGB 空间中均匀取色保证互不相同且可逆。若不传num_labelsmask_to_rgb会把掩码直接在通道维上复制 3 份灰度复制此时后处理只能按通道均值转整型无法区分多类别。测试test_image_processor_palettetest_image_processing_seggpt.py断言调色板长度为num_labels 1且首项为(0, 0, 0)test_mask_to_rgb则验证单类别时灰度复制只产生(0,0,0)/(1,1,1)而调色板上色产生(0,0,0)/(255,255,255)。SegGptModel 编码器上下文拼接、merge 与特征融合SegGptModel.forward的完整签名见 modeling_seggpt.py核心输入为参数形状说明pixel_values(B, C, H, W)待分割图像prompt_pixel_values(B, C, H, W)提示图像prompt_masks(B, C, H, W)提示掩码bool_masked_pos(B, num_patches)布尔掩码位置1 表示该 patch 需被重建推理时可省略模型会自动构造feature_ensemblebool是否启用特征融合few-shot多个提示时推荐开启embedding_typestrsemantic或instance决定使用哪种任务类型嵌入labels(B, C, H, W)训练时的真实掩码前向流程的关键步骤上下文化拼接L704-L727pixel_values cat([prompt_pixel_values, pixel_values], dim2)沿高度拼接成(B, C, 2H, W)提示侧的伪像素则用cat([prompt_masks, prompt_masks], dim2)构造推理时或cat([prompt_masks, labels], dim2)训练时。注释明确指出推理时提示掩码中对应预测区域的部分不携带任何信息必须用bool_masked_pos屏蔽——若未提供模型自动把后半即待预测区的所有 patch 置为 1。嵌入构造SegGptEmbeddings对拼接后的图像与提示侧伪图像做 patch embedding将bool_masked_pos为 1 的位置替换为可学习的mask_token再加上segment_token_input/prompt、插值后的位置编码pretrain_image_size224起双三次插值适配实际 patch 网格与semantic/instance类型 token最后沿 batch 维把提示侧与输入侧串成一个更长的批。merge 机制L470-L473到达merge_index层时把提示侧特征与输入侧特征逐元素平均二者此后共享同一表征。特征融合feature ensembleSegGptLayer.forward当feature_ensembleTrue且提示数 ≥ 2 时即同一目标图像配多个提示每层注意力输出会把同一图像的多个提示的输入特征求平均再拼回去相当于跨提示的特征集成。这正是文档 Tips 第 4 条batch_size 1时可传feature_ensembleTrue的底层实现。中间层采集L475-L476intermediate_hidden_state_indices中每个索引层的输出经 LayerNorm 后存入intermediate_hidden_states。SegGptModel返回SegGptEncoderOutput其last_hidden_state形状为(B, patch_height, patch_width, hidden_size)如官方文档示例中 vit-large 配置下为[1, 56, 28, 1024]。SegGptForImageSegmentationRGB 解码与 smooth-L1 损失SegGptForImageSegmentation在编码器之上挂了一个轻量解码器SegGptDecoder把intermediate_hidden_state_indices个中间特征沿通道拼接后经decoder_embed线性层输出维度patch_size**2 * decoder_hidden_size投影_reshape_hidden_states将其重排为像素级特征图(B, decoder_hidden_size, H, W)SegGptDecoderHead3×3 卷积 channels-first LayerNorm 激活 1×1 卷积输出3 通道的pred_masks形状(B, 3, 2H, W)——即直接预测 RGB颜色掩码前一半高度是提示区无信息后一半是待分割区。训练时若提供labels会计算 SegGptLoss将pred_masks与cat([prompt_masks, labels], dim2)的 ground truth 做F.smooth_l1_lossbetaconfig.beta默认 0.01并仅对bool_masked_pos为 1 的 patch 区域求平均即只惩罚被屏蔽、需要重建的 patch。输出为SegGptImageSegmentationOutput(loss, pred_masks, hidden_states, attentions)。post_process_semantic_segmentation把 RGB 预测还原为类别图SegGptImageProcessor.post_process_semantic_segmentation实现PIL 后端在 image_processing_pil_seggpt.py 有等价实现把原始输出转成语义分割图流程如下切掉高度前一半提示区只保留待分割区masks[:, :, masks.shape[2] // 2 :, :]反归一化乘 std 加 mean通道置末位再置回、乘 255 并裁剪到[0, 255]若给定target_sizes长度须等于 batch 维否则抛错用nearest插值恢复到目标尺寸类别归属提供num_labels时对每个像素计算它与调色板num_labels 1种颜色的平方 L2 距离取最近颜色作为预测类别argmin不提供时退化为三通道均值取整仅适合单类别/灰度场景return_segmentation_scoresTrue时返回SemanticSegmentationPostProcessorOutput其segmentation为类别图(H, W)segmentation_scores为形状(num_labels1, H, W)的负平方 L2 距离分数默认返回纯list[torch.Tensor]类别图。注意num_labels在预处理与后处理中必须一致且都不含背景类别 0。测试 test_post_processing_semantic_segmentation 验证了后处理输出高度为size[height] // 2、宽度不变即提示区已被正确裁掉。完整实战示例one-shot 语义分割以下示例继承自官方模型文档 seggpt.md使用BAAI/seggpt-vit-large检查点以 Hugging Face 数据集EduardoPacheco/FoodSeg103103 个食物类别不含背景演示 one-shot 语义分割可直接复现import torch from datasets import load_dataset from transformers import SegGptForImageSegmentation, SegGptImageProcessor checkpoint BAAI/seggpt-vit-large image_processor SegGptImageProcessor.from_pretrained(checkpoint) model SegGptForImageSegmentation.from_pretrained(checkpoint, device_mapauto) dataset_id EduardoPacheco/FoodSeg103 ds load_dataset(dataset_id, splittrain) # Number of labels in FoodSeg103 (not including background) num_labels 103 image_input ds[4][image] # 待分割图像 ground_truth ds[4][label] # 真值掩码 image_prompt ds[29][image] # 提示图像 mask_prompt ds[29][label] # 提示掩码分割图2D inputs image_processor( imagesimage_input, prompt_imagesimage_prompt, segmentation_mapsmask_prompt, # 2D 分割图 num_labelsnum_labels, # 强烈建议传入用于构建调色板 return_tensorspt, ) with torch.no_grad(): outputs model(**inputs) target_sizes [image_input.size[::-1]] # PIL 的 size 为 (w, h)需翻转为 (h, w) mask image_processor.post_process_semantic_segmentation(outputs, target_sizes, num_labelsnum_labels)[0]说明两个细节image_input.size[::-1]是把 PIL 的(width, height)翻转为后处理期望的(height, width)文档中的segmentation_maps参数在preprocess的**kwargs通道中透传为prompt_masks的语义处理器对分割图按 2D 分支处理。若mask_prompt本身就是 RGB 图像则应改为传prompt_masksmask_prompt并加do_convert_rgbFalse见前文提示掩码的两种合法格式。模型 forward 层面还有一个独立的最小调用示例展示SegGptModel仅编码器的输出形状摘自 SegGptModel.forward 文档from transformers import SegGptImageProcessor, SegGptModel from PIL import Image import httpx from io import BytesIO # 从 Painter 仓库的示例图目标图 / 提示图 / 提示掩码灰度图下载 checkpoint BAAI/seggpt-vit-large model SegGptModel.from_pretrained(checkpoint) image_processor SegGptImageProcessor.from_pretrained(checkpoint) inputs image_processor(imagesimage_input, prompt_imagesimage_prompt, prompt_masksmask_prompt, return_tensorspt) outputs model(**inputs) list(outputs.last_hidden_state.shape) # - [1, 56, 28, 1024]测试依据与使用要点速查仓库测试为该模型的正确性提供了两层验证test_modeling_seggpt.pySegGptModelTester用image_size30、patch_size2等迷你配置验证SegGptModel与SegGptForImageSegmentation的输出形状last_hidden_state为(B, image_size/patch_size, image_size/patch_size, hidden_size)并单独覆盖SegGptLoss与feature_ensemble分支test_image_processing_seggpt.py除前文提到的掩码等价、调色板、后处理测试外test_prompt_mask_equivalenceL249-L321验证了 numpy / torch / PIL 三种输入、单张与批量 2D 分割图或 3D RGB 掩码之间输出完全一致test_backends_equivalence进一步断言 torchvision 与 PIL 两个后端的pixel_values、prompt_pixel_values、prompt_masks张量级等价。结合官方文档的 4 条 Tips 与源码细节实操要点可归纳为用SegGptImageProcessor统一准备图像、提示图与提示掩码文档 Tips 1提示掩码可以是分割图2D也可以是 RGB 图像后者必须传do_convert_rgbFalseTips 2对应源码 L191-L215 的两分支使用分割图做前后处理时务必传入num_labels不含背景使预处理上色与后处理反查使用同一调色板Tips 3build_palette/post_process_semantic_segmentation首尾呼应few-shot 推理同一目标图像配多个提示、batch_size 1时传feature_ensembleTrue源码在SegGptLayer中对提示侧输入特征做跨提示平均Tips 4L414-L423自定义配置时记住结构性约束merge_index min(intermediate_hidden_state_indices)且image_size高度默认是单张图像的 2 倍896 × 448 提示区 待分割区自定义image_size时需同时调整处理器size保证 patch 网格与位置编码插值自洽。该实现由 EduardoPacheco 贡献原始 PyTorch 代码位于 BAAI Painter 仓库的 SegGPT 目录见 convert_seggpt_to_hf.py 中的转换逻辑可将原始权重迁移到本仓库的 HF 格式。以上所有接口行为均以当前仓库源码为准适用前提为已安装 PyTorch 与图像依赖torch、PIL/torchvision。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考