Transformers 中 SAM-HQ 模型使用指南:高质量可提示图像分割的原理、参数与实操
【免费下载链接】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
SAM-HQ(Segment Anything Model in High Quality)是在原始 SAM 基础上通过“高质量输出 Token + 全局-局部特征融合”实现更精细分割掩码的增强模型。Hugging Face Transformers 已将其完整集成,提供SamHQModel、SamHQProcessor及一整套配置类,可直接用于点提示(point prompt)、框提示(box prompt)甚至掩码输入的高精度分割。本文基于仓库文档 sam_hq.md 展开,并结合 sam_hq 源码目录 深入讲解每个关键参数、处理器输入结构与底层实现细节。
一、SAM-HQ 是什么:在 SAM 之上做最小侵入式增强
根据官方模型文档(该模型论文于 2023-06-02 发布,2025-04-28 贡献进入 Transformers),SAM-HQ 是一个对原始 SAM 的增强模型:在保持 SAM 原有可提示设计、效率与零样本泛化能力的前提下,产生显著更高质量的分割掩码。
论文摘要(引自文档)说明了其核心设计思想:
SAM 用 11 亿掩码训练,仍具备强大的零样本能力与灵活提示,但在处理结构精细的物体时掩码质量仍有不足。HQ-SAM 复用并保留了 SAM 的预训练权重,仅引入极少量的额外参数与计算:设计一个可学习的High-Quality Output Token(高质量输出 Token),注入 SAM 的 mask decoder,负责预测高质量掩码;并且不只在 mask-decoder 特征上操作,而是先与 ViT 的早期与最终特征融合,以改善掩码细节。训练所用的可学习参数来自一个由多个来源组成的 44K 细粒度掩码数据集——整个训练仅需 8 块 GPU 约 4 小时。
文档列出的五大改进点,均可在源码中找到对应实现:
- High-Quality Output Token:在 mask decoder 中注入可学习 token 以提升掩码质量。对应 modular_sam_hq.py 中的
self.hq_token = nn.Embedding(1, self.hidden_size)与配套的hq_mask_mlp。 - Global-local Feature Fusion(全局-局部特征融合):融合模型不同阶段的特征以改善掩码细节。实现上,SamHQVisionEncoder 在 forward 中会收集非窗口注意力层(
window_size == 0)的中间嵌入intermediate_embeddings,mask decoder 再用compress_vit_conv1/2压缩 ViT 特征后与 decoder 特征相加(hq_features = embed_encode + compressed_vit_features)。 - 训练数据:使用 44K 高质量掩码数据集,而非 SA-1B。
- 效率:仅新增约 0.5% 参数(文档声明)。
- 零样本能力:保持 SAM 的强零样本泛化,同时提升精度。
文档还给出了几条实用提示(Tips),值得在使用前记住:
- 对结构精细、细节丰富的物体,SAM-HQ 生成的掩码质量高于原始 SAM;
- 模型预测二值掩码,边界更准确,对细薄结构(thin structures)处理更好;
- 与 SAM 一样,输入 2D 点和/或输入框的效果更好;
- 可以对同一张图片提示多个点,模型预测出单个高质量掩码;
- 保持 SAM 的零样本泛化能力;
- 相比 SAM 仅增加约 0.5% 参数;
- 目前尚不支持微调(fine-tuning)。
二、快速上手:图像 + 2D 点提示生成掩码
文档给出的最小可用示例如下,使用syscv-community/sam-hq-vit-base检查点:
import requests import torch from PIL import Image from transformers import SamHQModel, SamHQProcessor model = SamHQModel.from_pretrained("syscv-community/sam-hq-vit-base", device_map="auto") processor = SamHQProcessor.from_pretrained("syscv-community/sam-hq-vit-base") img_url = "https://huggingface.co/ybelkada/segment-anything/resolve/main/assets/car.png" raw_image = Image.open(requests.get(img_url, stream=True).raw).convert("RGB") input_points = [[[450, 600]]] # 2D location of a window in the image inputs = processor(raw_image, input_points=input_points, return_tensors="pt").to(model.device) with torch.no_grad(): outputs = model(**inputs) masks = processor.image_processor.post_process_masks( outputs.pred_masks.cpu(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu() ) scores = outputs.iou_scores几个值得注意的细节:
input_points是三层嵌套列表:[图片批次, 掩码批次, 每个掩码的点数, [x, y]]。示例中[[[450, 600]]]表示 1 张图、1 个掩码、1 个点;处理器会将这些原图坐标归一化到模型目标尺寸(见下文处理器部分)。- 输出为
SamHQImageSegmentationOutput,包含pred_masks与iou_scores;masks与iou_scores的形状为(batch_size, point_batch_size, num_masks, height, width)与(batch_size, point_batch_size, num_masks)。 - 后处理必须传入
original_sizes与reshaped_input_sizes(处理器输出中自带这两个键),它们用于把低分辨率预测掩码还原/裁剪回原图尺寸。 - 也可以直接通过
AutoModel/AutoProcessor加载(模型 docstring 示例中使用的是sushmanth/sam_hq_vit_b检查点,见 modeling 文档字符串)。
三、掩码输入:把已有分割图一起喂给处理器
文档的第二个示例展示了将自定义掩码与图像一起输入的能力——处理器接受segmentation_maps参数:
import requests import torch from PIL import Image from transformers import SamHQModel, SamHQProcessor model = SamHQModel.from_pretrained("syscv-community/sam-hq-vit-base", device_map="auto") processor = SamHQProcessor.from_pretrained("syscv-community/sam-hq-vit-base") img_url = "https://huggingface.co/ybelkada/segment-anything/resolve/main/assets/car.png" raw_image = Image.open(requests.get(img_url, stream=True).raw).convert("RGB") mask_url = "https://huggingface.co/ybelkada/segment-anything/resolve/main/assets/car.png" segmentation_map = Image.open(requests.get(mask_url, stream=True).raw).convert("1") input_points = [[[450, 600]]] # 2D location of a window in the image inputs = processor( raw_image, input_points=input_points, segmentation_maps=segmentation_map, return_tensors="pt" ).to(model.device) with torch.no_grad(): outputs = model(**inputs) masks = processor.image_processor.post_process_masks( outputs.pred_masks.cpu(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu() ) scores = outputs.iou_scores注意segmentation_map被转换为"1"(二值)模式。在 processing_sam_hq.py 中,SamHQImagesKwargs对segmentation_maps的说明是:这些真值分割图会与输入图像一起处理,用于训练或评估目的,会被缩放并归一化以匹配处理后图像的维度。
四、SamHQProcessor:输入结构与坐标归一化
处理器是 SAM-HQ 提示工程的核心。从 processing_sam_hq.py 可以看到SamHQProcessor.__call__的完整输入契约:
| 参数 | 结构 | 说明 |
|---|---|---|
images | ImageInput | 输入图像(PIL/NumPy/路径均可) |
input_points | [image_level, object_level, point_level, [x, y]] | 原图坐标空间的点提示;处理器会归一化到目标尺寸 |
input_labels | [image_level, object_level, point_level] | 每个点的标签,结构须与input_points(去掉坐标维)一致 |
input_boxes | [image_level, box_level, [x1, y1, x2, y2]] | 原图坐标空间的框提示,x1/y1/x2/y2分别为左上、右下 |
segmentation_maps | ImageInput | 与图像一同处理的分割图(训练/评估场景) |
point_pad_value | int,默认None | 变长点序列批处理时的填充值;为None时使用处理器配置默认值 |
mask_size/mask_pad_size | dict[str, int] | 控制输出掩码的目标尺寸与批处理时的 padding |
关键点:
- 标签语义(来自模型
forward文档字符串,modular_sam_hq.py):1表示点位于目标物体上,0表示点不在物体上,-1表示背景;Transformers 额外定义了-10表示 padding 点,会被 prompt encoder 忽略,且这部分由处理器自动完成。若只传点不传标签,forward会自动将标签置为全 1(默认全部视为前景点)。 - 坐标归一化:
_normalize_coordinates(processing_sam_hq.py#L198-L217)按new_w / old_w、new_h / old_h的缩放比把原图坐标映射到预处理后的尺寸,缩放由image_processor._get_preprocess_shape(original_size, longest_edge=target_size)决定——因此你始终可以按原始像素坐标写提示,无需自己换算。 - 变长点自动 padding:
_pad_points_and_labels(processing_sam_hq.py#L182-L196)会把同一批次的点/标签补齐到最长序列,padding 坐标设为point_pad_value(默认 -10)。 - 张量维度:
input_points最终是 4D 张量,input_labels是 3D,input_boxes是 3D;这与SamHQModel.forward中的形状校验一致(点为(batch_size, num_points, 2)起步、框为(batch_size, num_boxes, 4),见 forward 校验逻辑)。
五、SamHQModel.forward:核心参数解析
SamHQModel.forward(modular_sam_hq.py#L417-L584)除了常规的pixel_values、input_points、input_labels、input_boxes、input_masks外,还有两个 SAM-HQ 特有或值得强调的参数:
multimask_output(默认True):SAM 论文中每个提示输出 3 个掩码;设为False时只输出对应的“最佳”单掩码。实现上,SamHQMaskDecoder 在multimask_output=True时会按 IoU 分数降序排序多掩码输出。hq_token_only(默认False):SAM-HQ 特有的开关。False时最终掩码为标准 SAM 掩码与 HQ 掩码之和(masks = masks_sam + masks_hq);True时只取 HQ token 路径输出的掩码。模型 docstring 中的示例演示了两种用法:>>> # Get high-quality segmentation mask >>> outputs = model(**inputs) >>> # For high-quality mask only >>> outputs = model(**inputs, hq_token_only=True)image_embeddings与get_image_embeddings:为节省显存,可以先调用model.get_image_embeddings(pixel_values)(modular_sam_hq.py#L400-L415)预计算图像嵌入,再把嵌入喂给forward替代pixel_values。注意两者互斥:同时传会抛出ValueError。由于 HQ 特征融合依赖中间嵌入,此路径下还需把get_image_embeddings返回的第二项intermediate_embeddings一并传入forward(文档字符串明确要求)。attention_similarity/target_embedding:可选的个性化(PerSAM 论文引入)参数,用于目标引导注意力与目标语义提示,一般场景不需要。
六、配置类:SamHQConfig 及其三个子配置
SamHQConfig是复合配置(configuration_sam_hq.py#L151-L190),model_type = "sam_hq",聚合三个子配置:
sub_configs = { "prompt_encoder_config": SamHQPromptEncoderConfig, "mask_decoder_config": SamHQMaskDecoderConfig, "vision_config": SamHQVisionConfig, }各子配置的关键字段与默认值(可直接用于从零构建配置):
SamHQVisionConfig(视觉编码器,configuration_sam_hq.py#L52-L114):
| 字段 | 默认值 | 说明 |
|---|---|---|
hidden_size | 768 | ViT 隐层维度(base 规格) |
num_hidden_layers/num_attention_heads | 12 / 12 | 编码器层数与注意力头数 |
image_size/patch_size | 1024 / 16 | 输入尺寸与 patch 大小 |
output_channels | 256 | patch encoder 输出通道数 |
use_rel_pos | True | 是否使用相对位置编码 |
window_size | 14 | 相对位置编码窗口大小 |
global_attn_indexes | (2, 5, 8, 11) | 全局注意力层的索引——这些层正是 SAM-HQ 用来收集中间特征的“非窗口”层 |
mlp_dim | None | 未指定时自动取hidden_size * mlp_ratio |
SamHQMaskDecoderConfig(configuration_sam_hq.py#L117-L148):
| 字段 | 默认值 | 说明 |
|---|---|---|
hidden_size | 256 | decoder 隐层维度 |
num_hidden_layers/num_attention_heads | 2 / 8 | 双向 transformer 的层数/头数 |
attention_downsample_rate | 2 | 注意力下采样率 |
num_multimask_outputs | 3 | 多掩码输出数(对应multimask_output=True时的 3 个掩码) |
iou_head_depth/iou_head_hidden_dim | 3 / 256 | IoU 预测头深度/隐层维度 |
vit_dim | 768 | 参与特征融合的 ViT 维度(SAM-HQ 相对 SAM 的新增字段,决定compress_vit_conv1的输入通道) |
SamHQPromptEncoderConfig(configuration_sam_hq.py#L27-L49):
| 字段 | 默认值 | 说明 |
|---|---|---|
hidden_size | 256 | 点/框提示嵌入维度 |
mask_input_channels | 16 | 喂给 mask decoder 的掩码通道数 |
num_point_embeddings | 4 | 点嵌入数量 |
image_size/patch_size | 1024 / 16 | 与视觉编码器对齐,__post_init__中会计算image_embedding_size = image_size // patch_size |
从源码结构看,这套配置类完全继承自 SAM 的对应配置(见 modular_sam_hq.py#L44-L79 中SamHQPromptEncoderConfig(SamPromptEncoderConfig)等定义),SAM-HQ 的增量配置只有 mask decoder 的vit_dim一项,印证了“最小侵入式增强”的设计。
七、模型规格与检查点转换
仓库自带官方的检查点转换脚本 convert_samhq_to_hf.py,可用于把原始仓库的.pth权重(来自lkeab/hq-sam)转为 HF 格式,支持三个规格:
sam_hq_vit_b:base,默认SamHQVisionConfig(),vit_dim = 768;sam_hq_vit_l:large,hidden_size=1024, num_hidden_layers=24, num_attention_heads=16, global_attn_indexes=[5, 11, 17, 23],vit_dim = 1024;sam_hq_vit_h:huge,hidden_size=1280, num_hidden_layers=32, num_attention_heads=16, global_attn_indexes=[7, 15, 23, 31],vit_dim = 1280。
脚本用法(参数见 脚本主入口):
python src/transformers/models/sam_hq/convert_samhq_to_hf.py \ --model_name sam_hq_vit_b \ --checkpoint_path /path/to/sam_hq_vit_b.pth \ --pytorch_dump_folder_path ./out # 可选:--push_to_hub --hub_path <user>脚本内部通过KEYS_TO_MODIFY_MAPPING(convert_samhq_to_hf.py#L66-L116)做权重键名映射,其中 HQ 特有部分的映射(hf_token → hq_token、compress_vit_feat → compress_vit_conv*、embedding_encoder → encoder_conv*、embedding_maskfeature → mask_conv*、hf_mlp → hq_mask_mlp)恰好对应上文源码分析中的特征融合与 HQ token 模块。转换后默认使用SamImageProcessor构造SamHQProcessor。
文档推荐的日常使用检查点为syscv-community/sam-hq-vit-base(docstring 示例另用sushmanth/sam_hq_vit_b)。
八、质量验证与限制
- 测试:模型与处理器的测试分别位于 test_modeling_sam_hq.py 与 test_processing_sam_hq.py,其中建模测试覆盖了视觉编码器形状校验、注意力输出维度等,可作为行为基准参考。
- 限制:文档明确说明暂不支持微调;
forward中若同时提供pixel_values与image_embeddings、或点/框批次维度不一致,都会抛出ValueError。 - 适用前提:需要 PyTorch 环境;图像输入经
SamImageProcessor预处理为 1024 边长的标准尺寸(由SamHQProcessor.__init__读取image_processor.size["longest_edge"]作为target_size)。
小结
- SAM-HQ 在 Transformers 中以
SamHQModel+SamHQProcessor提供完整的点/框/掩码提示分割能力,用法与 SAM 高度一致; - 两个最有价值的 HQ 专属开关是
hq_token_only(仅取 HQ 路径掩码)与多掩码排序输出; - 提示坐标一律按原图像素书写,处理器负责归一化与 padding;
- 显存敏感场景可先用
get_image_embeddings缓存图像嵌入(记得同时传递intermediate_embeddings); - base/large/huge 三档规格可通过仓库内置脚本从原始检查点转换。
【免费下载链接】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),仅供参考