Transformers 中 SAM-HQ 模型使用指南:高质量可提示图像分割的原理、参数与实操
2026/9/8 16:37:19 网站建设 项目流程

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 已将其完整集成,提供SamHQModelSamHQProcessor及一整套配置类,可直接用于点提示(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 小时。

文档列出的五大改进点,均可在源码中找到对应实现:

  1. High-Quality Output Token:在 mask decoder 中注入可学习 token 以提升掩码质量。对应 modular_sam_hq.py 中的self.hq_token = nn.Embedding(1, self.hidden_size)与配套的hq_mask_mlp
  2. 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)。
  3. 训练数据:使用 44K 高质量掩码数据集,而非 SA-1B。
  4. 效率:仅新增约 0.5% 参数(文档声明)。
  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_masksiou_scoresmasksiou_scores的形状为(batch_size, point_batch_size, num_masks, height, width)(batch_size, point_batch_size, num_masks)
  • 后处理必须传入original_sizesreshaped_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 中,SamHQImagesKwargssegmentation_maps的说明是:这些真值分割图会与输入图像一起处理,用于训练或评估目的,会被缩放并归一化以匹配处理后图像的维度。

四、SamHQProcessor:输入结构与坐标归一化

处理器是 SAM-HQ 提示工程的核心。从 processing_sam_hq.py 可以看到SamHQProcessor.__call__的完整输入契约:

参数结构说明
imagesImageInput输入图像(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_mapsImageInput与图像一同处理的分割图(训练/评估场景)
point_pad_valueint,默认None变长点序列批处理时的填充值;为None时使用处理器配置默认值
mask_size/mask_pad_sizedict[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_wnew_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_valuesinput_pointsinput_labelsinput_boxesinput_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_embeddingsget_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_size768ViT 隐层维度(base 规格)
num_hidden_layers/num_attention_heads12 / 12编码器层数与注意力头数
image_size/patch_size1024 / 16输入尺寸与 patch 大小
output_channels256patch encoder 输出通道数
use_rel_posTrue是否使用相对位置编码
window_size14相对位置编码窗口大小
global_attn_indexes(2, 5, 8, 11)全局注意力层的索引——这些层正是 SAM-HQ 用来收集中间特征的“非窗口”层
mlp_dimNone未指定时自动取hidden_size * mlp_ratio

SamHQMaskDecoderConfig(configuration_sam_hq.py#L117-L148):

字段默认值说明
hidden_size256decoder 隐层维度
num_hidden_layers/num_attention_heads2 / 8双向 transformer 的层数/头数
attention_downsample_rate2注意力下采样率
num_multimask_outputs3多掩码输出数(对应multimask_output=True时的 3 个掩码)
iou_head_depth/iou_head_hidden_dim3 / 256IoU 预测头深度/隐层维度
vit_dim768参与特征融合的 ViT 维度(SAM-HQ 相对 SAM 的新增字段,决定compress_vit_conv1的输入通道)

SamHQPromptEncoderConfig(configuration_sam_hq.py#L27-L49):

字段默认值说明
hidden_size256点/框提示嵌入维度
mask_input_channels16喂给 mask decoder 的掩码通道数
num_point_embeddings4点嵌入数量
image_size/patch_size1024 / 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_tokencompress_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_valuesimage_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),仅供参考

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

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

立即咨询