ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Transformers 中 SAM-HQ 模型使用指南:高质量可提示图像分割的原理、参数与实操

Transformers 中 SAM-HQ 模型使用指南:高质量可提示图像分割的原理、参数与实操 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/transformersSAM-HQSegment 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 贡献进入 TransformersSAM-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_embeddingsmask 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_mapauto) 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, streamTrue).raw).convert(RGB) input_points [[[450, 600]]] # 2D location of a window in the image inputs processor(raw_image, input_pointsinput_points, return_tensorspt).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_scoresmasks与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_mapauto) 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, streamTrue).raw).convert(RGB) mask_url https://huggingface.co/ybelkada/segment-anything/resolve/main/assets/car.png segmentation_map Image.open(requests.get(mask_url, streamTrue).raw).convert(1) input_points [[[450, 600]]] # 2D location of a window in the image inputs processor( raw_image, input_pointsinput_points, segmentation_mapssegmentation_map, return_tensorspt ).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__的完整输入契约参数结构说明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.py1表示点位于目标物体上0表示点不在物体上-1表示背景Transformers 额外定义了-10表示 padding 点会被 prompt encoder 忽略且这部分由处理器自动完成。若只传点不传标签forward会自动将标签置为全 1默认全部视为前景点。坐标归一化_normalize_coordinatesprocessing_sam_hq.py#L198-L217按new_w / old_w、new_h / old_h的缩放比把原图坐标映射到预处理后的尺寸缩放由image_processor._get_preprocess_shape(original_size, longest_edgetarget_size)决定——因此你始终可以按原始像素坐标写提示无需自己换算。变长点自动 padding_pad_points_and_labelsprocessing_sam_hq.py#L182-L196会把同一批次的点/标签补齐到最长序列padding 坐标设为point_pad_value默认 -10。张量维度input_points最终是 4D 张量input_labels是 3Dinput_boxes是 3D这与SamHQModel.forward中的形状校验一致点为(batch_size, num_points, 2)起步、框为(batch_size, num_boxes, 4)见 forward 校验逻辑。五、SamHQModel.forward核心参数解析SamHQModel.forwardmodular_sam_hq.py#L417-L584除了常规的pixel_values、input_points、input_labels、input_boxes、input_masks外还有两个 SAM-HQ 特有或值得强调的参数multimask_output默认TrueSAM 论文中每个提示输出 3 个掩码设为False时只输出对应的“最佳”单掩码。实现上SamHQMaskDecoder 在multimask_outputTrue时会按 IoU 分数降序排序多掩码输出。hq_token_only默认FalseSAM-HQ 特有的开关。False时最终掩码为标准 SAM 掩码与 HQ 掩码之和masks masks_sam masks_hqTrue时只取 HQ token 路径输出的掩码。模型 docstring 中的示例演示了两种用法 # Get high-quality segmentation mask outputs model(**inputs) # For high-quality mask only outputs model(**inputs, hq_token_onlyTrue)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-L190model_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_ratioSamHQMaskDecoderConfigconfiguration_sam_hq.py#L117-L148字段默认值说明hidden_size256decoder 隐层维度num_hidden_layers/num_attention_heads2 / 8双向 transformer 的层数/头数attention_downsample_rate2注意力下采样率num_multimask_outputs3多掩码输出数对应multimask_outputTrue时的 3 个掩码iou_head_depth/iou_head_hidden_dim3 / 256IoU 预测头深度/隐层维度vit_dim768参与特征融合的 ViT 维度SAM-HQ 相对 SAM 的新增字段决定compress_vit_conv1的输入通道SamHQPromptEncoderConfigconfiguration_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_bbase默认SamHQVisionConfig()vit_dim 768sam_hq_vit_llargehidden_size1024, num_hidden_layers24, num_attention_heads16, global_attn_indexes[5, 11, 17, 23]vit_dim 1024sam_hq_vit_hhugehidden_size1280, num_hidden_layers32, num_attention_heads16, 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_MAPPINGconvert_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-basedocstring 示例另用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 中以SamHQModelSamHQProcessor提供完整的点/框/掩码提示分割能力用法与 SAM 高度一致两个最有价值的 HQ 专属开关是hq_token_only仅取 HQ 路径掩码与多掩码排序输出提示坐标一律按原图像素书写处理器负责归一化与 padding显存敏感场景可先用get_image_embeddings缓存图像嵌入记得同时传递intermediate_embeddingsbase/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),仅供参考
返回列表