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/transformers
SegGPT 是 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.03284,Xinlong Wang 等人),模型文档注明该模型于 2023-04-06 发布于 HF papers、2024-02-26 合入 Transformers。其论文摘要指出:SegGPT 将各类分割任务统一为一个通用的上下文学习框架,把不同形式的分割数据转换为相同格式的图像,训练被形式化为一个"上下文填色问题"(in-context coloring problem),每个数据样本使用随机的颜色映射;训练目标是根据上下文完成任务,而不是依赖特定颜色。训练完成后,它可以在图像或视频上执行任意图上下文推理任务,如对象实例、stuff、部件、轮廓和文字分割,并覆盖 few-shot 语义分割、视频目标分割、语义分割、全景分割等多种任务。官方文档中给出的代表结果为:COCO-20 上 56.1 mIoU(one-shot)、FSS-1000 上 85.6 mIoU。
从 Transformers 的实现看,这套思想落在三段式结构上:
- 上下文拼接:把提示图像与待分割图像沿高度方向拼接成一张"高为两倍"的伪图像,输入同一个 ViT 编码器;
- 中间层特征采集:在编码器的若干中间层(默认第 5、11、17、23 层)取出特征,作为解码器输入;
- RGB 空间解码:解码器输出 3 通道的"颜色图",即直接在 RGB 像素空间预测掩码,再通过调色板(palette)映射回类别索引。
仓库内该模型的完整文件布局如下,均为相对仓库根目录的路径:
| 文件 | 作用 |
|---|---|
| configuration_seggpt.py | SegGptConfig,架构超参数定义 |
| modeling_seggpt.py | SegGptModel(编码器)、SegGptForImageSegmentation(含解码器与损失) |
| image_processing_seggpt.py | SegGptImageProcessor(torchvision 后端) |
| image_processing_pil_seggpt.py | SegGptImageProcessorPil(纯 PIL/numpy 后端) |
| convert_seggpt_to_hf.py | 将原始 Painter 仓库权重转换为 HF 格式 |
| test_modeling_seggpt.py、test_image_processing_seggpt.py | 模型与图像处理测试 |
推荐加载的官方检查点为BAAI/seggpt-vit-large。
SegGptConfig:关键参数与默认值
SegGptConfig继承自PreTrainedConfig,model_type = "seggpt"。以下参数与默认值均直接取自 configuration_seggpt.py:
| 参数 | 默认值 | 含义 |
|---|---|---|
hidden_size | 1024 | Transformer 隐藏维度 |
num_hidden_layers | 24 | 编码器层数 |
num_attention_heads | 16 | 注意力头数 |
hidden_act | "gelu" | 激活函数 |
image_size | (896, 448) | 输入图像尺寸(高度为提示与图像拼接后的总高) |
patch_size | 16 | patch 大小 |
mlp_dim | None(回退为hidden_size * 4) | MLP 维度,在__post_init__中若为None则置为hidden_size * 4 |
pretrain_image_size | 224 | 绝对位置编码的预训练尺寸,用于双三次插值 |
use_relative_position_embeddings | True | 注意力中是否使用分解式相对位置编码 |
merge_index | 2 | 提示特征与输入特征"合并"(取平均)的编码器层索引 |
intermediate_hidden_state_indices | (5, 11, 17, 23) | 供解码器使用的中间层索引 |
decoder_hidden_size | 64 | 解码器内部特征维度 |
beta | 0.01 | SegGptLoss(smooth-L1)的正则化因子 |
drop_path_rate | 0.1 | 随机深度(DropPath)线性插值的最大比率 |
qkv_bias | True | QKV 线性层是否带偏置 |
配置类还带有一条结构性校验(validate_architecture):merge_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:三类输入与提示掩码的两种格式
SegGptImageProcessor(torchvision 后端)与SegGptImageProcessorPil(PIL/numpy 后端)接口一致,默认参数为size={"height": 448, "width": 448}、do_resize/do_rescale/do_normalize=True、归一化使用 ImageNet 均值与标准差(image_processing_seggpt.py)。
preprocess接受三类输入(preprocess):
images:待分割的目标图像;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_inputs):
- 2D 分割图(默认,
do_convert_rgb=True):处理器把每个类别索引通过调色板"上色"为 3 通道 RGB; - 已是 RGB 的 3 通道掩码图:必须传
do_convert_rgb=False,处理器才会按 3 通道直接处理,否则会因维度不匹配而报错。
两个分支中,掩码的缩放采样方式都被强制替换为PILImageResampling.NEAREST(L218),避免边界类别被双三次插值"染色"。test_image_processing_seggpt.py 中的test_mask_equivalence(L151-L161)专门验证:同一掩码走"灰度分割图"路径与走"RGB + do_convert_rgb=False"路径,输出的prompt_masks张量完全相等。
num_labels 与调色板(palette)
文档 Tips 第 3 条强烈建议:使用segmentation_maps做前后处理时传入num_labels(不含背景类)。其原理在 build_palette:
def 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_labels,mask_to_rgb会把掩码直接在通道维上复制 3 份(灰度复制),此时后处理只能按通道均值转整型,无法区分多类别。测试test_image_processor_palette(test_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_ensemble | bool | 是否启用特征融合(few-shot,多个提示时推荐开启) |
embedding_type | str | "semantic"或"instance",决定使用哪种任务类型嵌入 |
labels | (B, C, H, W) | 训练时的真实掩码 |
前向流程的关键步骤:
- 上下文化拼接(L704-L727):
pixel_values = cat([prompt_pixel_values, pixel_values], dim=2),沿高度拼接成(B, C, 2H, W);提示侧的"伪像素"则用cat([prompt_masks, prompt_masks], dim=2)构造(推理时)或cat([prompt_masks, labels], dim=2)(训练时)。注释明确指出:推理时提示掩码中对应预测区域的部分不携带任何信息,必须用bool_masked_pos屏蔽——若未提供,模型自动把后半(即待预测区)的所有 patch 置为 1。 - 嵌入构造(SegGptEmbeddings):对拼接后的图像与"提示侧伪图像"做 patch embedding,将
bool_masked_pos为 1 的位置替换为可学习的mask_token,再加上segment_token_input/prompt、插值后的位置编码(pretrain_image_size=224起,双三次插值适配实际 patch 网格)与semantic/instance类型 token,最后沿 batch 维把提示侧与输入侧串成一个更长的"批"。 - merge 机制(L470-L473):到达
merge_index层时,把提示侧特征与输入侧特征逐元素平均,二者此后共享同一表征。 - 特征融合(feature ensemble)(SegGptLayer.forward):当
feature_ensemble=True且提示数 ≥ 2 时(即同一目标图像配多个提示),每层注意力输出会把同一图像的多个提示的输入特征求平均再拼回去,相当于跨提示的特征集成。这正是文档 Tips 第 4 条"batch_size > 1时可传feature_ensemble=True"的底层实现。 - 中间层采集(L475-L476):
intermediate_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])。
SegGptForImageSegmentation:RGB 解码与 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);SegGptDecoderHead(3×3 卷积 + channels-first LayerNorm + 激活 + 1×1 卷积)输出3 通道的pred_masks,形状(B, 3, 2H, W)——即直接预测 RGB"颜色掩码",前一半高度是提示区(无信息),后一半是待分割区。
训练时若提供labels,会计算 SegGptLoss:将pred_masks与cat([prompt_masks, labels], dim=2)的 ground truth 做F.smooth_l1_loss(beta=config.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_scores=True时返回SemanticSegmentationPostProcessorOutput,其segmentation为类别图(H, W),segmentation_scores为形状(num_labels+1, 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/FoodSeg103(103 个食物类别,不含背景)演示 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_map="auto") dataset_id = "EduardoPacheco/FoodSeg103" ds = load_dataset(dataset_id, split="train") # 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( images=image_input, prompt_images=image_prompt, segmentation_maps=mask_prompt, # 2D 分割图 num_labels=num_labels, # 强烈建议传入,用于构建调色板 return_tensors="pt", ) 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_labels=num_labels)[0]说明两个细节:image_input.size[::-1]是把 PIL 的(width, height)翻转为后处理期望的(height, width);文档中的segmentation_maps参数在preprocess的**kwargs通道中透传为prompt_masks的语义(处理器对分割图按 2D 分支处理)。若mask_prompt本身就是 RGB 图像,则应改为传prompt_masks=mask_prompt并加do_convert_rgb=False(见前文"提示掩码的两种合法格式")。
模型 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(images=image_input, prompt_images=image_prompt, prompt_masks=mask_prompt, return_tensors="pt") outputs = model(**inputs) list(outputs.last_hidden_state.shape) # -> [1, 56, 28, 1024]测试依据与使用要点速查
仓库测试为该模型的正确性提供了两层验证:
- test_modeling_seggpt.py:
SegGptModelTester用image_size=30、patch_size=2等迷你配置验证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_equivalence(L249-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_rgb=False(Tips 2,对应源码 L191-L215 的两分支); - 使用分割图做前后处理时务必传入
num_labels(不含背景),使预处理上色与后处理反查使用同一调色板(Tips 3,build_palette/post_process_semantic_segmentation首尾呼应); - few-shot 推理(同一目标图像配多个提示、
batch_size > 1)时传feature_ensemble=True,源码在SegGptLayer中对提示侧输入特征做跨提示平均(Tips 4,L414-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),仅供参考