Transformers 文档

SAM-HQ

Hugging Face's logo
加入 Hugging Face 社区

并获得增强的文档体验

开始使用

该模型于 2023 年 6 月 2 日发布在 HF papers 上,并于 2025 年 4 月 28 日贡献给 Hugging Face Transformers。

SAM-HQ

概述

SAM-HQ(高质量 Segment Anything 模型)由 Lei Ke、Mingqiao Ye、Martin Danelljan、Yifan Liu、Yu-Wing Tai、Chi-Keung Tang 和 Fisher Yu 在 Segment Anything in High Quality 一文中提出。

该模型是对原始 SAM 模型的改进,它能在保持 SAM 原有的提示设计、效率和零样本泛化能力的同时,生成质量显著更高的分割掩码。

example image

与原始 SAM 模型相比,SAM-HQ 引入了几个关键改进:

  1. 高质量输出标记(High-Quality Output Token):注入到 SAM 掩码解码器中的可学习标记,用于实现更高质量的掩码预测。
  2. 全局-局部特征融合(Global-local Feature Fusion):结合了模型不同阶段的特征,以改善掩码细节。
  3. 训练数据:使用了一个精心筛选的 4.4 万个高质量掩码数据集,而非 SA-1B。
  4. 效率:仅增加了 0.5% 的额外参数,但显著提升了掩码质量。
  5. 零样本能力:在提高精度的同时,保持了 SAM 强大的零样本性能。

论文摘要如下:

近期的 Segment Anything Model (SAM) 代表了分割模型规模化的一大飞跃,实现了强大的零样本能力和灵活的提示功能。尽管使用了 11 亿个掩码进行训练,SAM 的掩码预测质量在许多情况下仍显不足,特别是在处理具有复杂结构的对象时。我们提出了 HQ-SAM,使 SAM 能够准确分割任何物体,同时保持其原始的提示设计、效率和零样本泛化能力。我们精心的设计重用并保留了 SAM 的预训练模型权重,只引入了极少的额外参数和计算量。我们设计了一个可学习的高质量输出标记,将其注入到 SAM 的掩码解码器中,负责预测高质量掩码。我们没有仅仅将其应用于掩码解码器特征,而是首先将它们与早期和最终的 ViT 特征进行融合,以改善掩码细节。为了训练我们引入的可学习参数,我们整理了一个包含 4.4 万个精细掩码的数据集。HQ-SAM 仅在所引入的 4.4 万个掩码数据集上进行训练,在 8 张 GPU 上仅需 4 小时即可完成。

技巧

  • SAM-HQ 生成的掩码质量高于原始 SAM 模型,特别适用于具有复杂结构和精细细节的对象。
  • 该模型预测出的二值掩码边界更准确,且能更好地处理细长结构。
  • 与 SAM 一样,该模型在使用输入 2D 点和/或输入边界框时表现更佳。
  • 您可以为同一图像提示多个点,并预测出单个高质量掩码。
  • 该模型保持了 SAM 的零样本泛化能力。
  • 与 SAM 相比,SAM-HQ 仅增加了约 0.5% 的额外参数。
  • 目前尚不支持对该模型进行微调。

该模型由 sushmanth 贡献。原始代码可以在这里找到。

以下是给定图像和 2D 点进行掩码生成的示例。

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

您还可以在处理器中处理您自己的掩码,并将其与输入图像一起传递给模型。

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

资源

以下是官方 Hugging Face 和社区(以 🌎 表示)资源列表,帮助您上手 SAM-HQ。

SamHQConfig

class transformers.SamHQConfig

< >

(省略代码参数列表)

参数

  • vision_config (Union[dict, ~configuration_utils.PreTrainedConfig], 可选) — 视觉主干网络的配置对象或字典。
  • prompt_encoder_config (Union[dict, SamHQPromptEncoderConfig], 可选) — 用于初始化 SamHQPromptEncoderConfig 的配置选项字典。
  • mask_decoder_config (Union[dict, SamHQMaskDecoderConfig], 可选) — 用于初始化 SamHQMaskDecoderConfig 的配置选项字典。
  • initializer_range (float, 可选, 默认为 0.02) — 用于初始化所有权重矩阵的 truncated_normal_initializer 的标准差。
  • tie_word_embeddings (bool, 可选, 默认为 True) — 是否根据模型的 tied_weights_keys 映射来绑定权重嵌入。

这是用于存储 SamHQModel 配置的类。它用于根据指定的参数实例化 Sam Hq 模型,定义模型架构。使用默认值实例化配置将产生与 syscv-community/sam-hq-vit-base 类似的配置。

配置对象继承自 PreTrainedConfig,可用于控制模型输出。阅读 PreTrainedConfig 的文档以获取更多信息。

SamHQVisionConfig

class transformers.SamHQVisionConfig

(省略源码链接)

(省略代码参数列表)

参数

  • hidden_size (int, 可选, 默认为 768) — 隐藏表示的维度。
  • output_channels (int, 可选, 默认为 256) — Patch Encoder 中输出通道的维度。
  • num_hidden_layers (int, 可选, 默认为 12) — Transformer 解码器中的隐藏层数量。
  • num_attention_heads (int, 可选, 默认为 12) — Transformer 解码器中每个注意力层的注意力头数。
  • num_channels (int, 可选, 默认为 3) — 输入通道的数量。
  • image_size (Union[int, list[int], tuple[int, int]], 可选, 默认为 1024) — 每个图像的大小(分辨率)。
  • patch_size (Union[int, list[int], tuple[int, int]], 可选, 默认为 16) — 每个补丁的大小(分辨率)。
  • hidden_act (str, 可选, 默认为 gelu) — 解码器中的非线性激活函数(函数或字符串)。例如:"gelu", "relu", "silu" 等。
  • layer_norm_eps (float, 可选, 默认为 1e-06) — 层归一化层使用的 epsilon 值。
  • attention_dropout (Union[float, int], 可选, 默认为 0.0) — 注意力概率的 dropout 比率。
  • initializer_range (float, 可选, 默认为 1e-10) — 用于初始化所有权重矩阵的 truncated_normal_initializer 的标准差。
  • qkv_bias (bool, 可选, 默认为 True) — 是否为查询、键和值添加偏置。
  • mlp_ratio (float, optional, 默认为 4.0) — MLP 隐藏层维度与嵌入维度的比率。
  • use_abs_pos (bool, optional, 默认为 True) — 是否使用绝对位置嵌入。
  • use_rel_pos (bool, optional, 默认为 True) — 是否使用相对位置嵌入。
  • window_size (int, optional, 默认为 14) — 相对位置的窗口大小。
  • global_attn_indexes (list[int], optional, 默认为 [2, 5, 8, 11]) — 全局注意力层的索引。
  • num_pos_feats (int, optional, 默认为 128) — 位置嵌入的维度。
  • mlp_dim (int, optional) — Transformer 编码器中 MLP 层的维度。如果为 None,则默认为 mlp_ratio * hidden_size

这是用于存储 SamHQModel 配置的类。它用于根据指定的参数实例化 Sam Hq 模型,定义模型架构。使用默认值实例化配置将产生与 syscv-community/sam-hq-vit-base 类似的配置。

配置对象继承自 PreTrainedConfig,可用于控制模型输出。阅读 PreTrainedConfig 的文档以获取更多信息。

示例

>>> from transformers import (
...     SamHQVisionConfig,
...     SamHQVisionModel,
... )

>>> # Initializing a SamHQVisionConfig with `"facebook/sam_hq-vit-huge"` style configuration
>>> configuration = SamHQVisionConfig()

>>> # Initializing a SamHQVisionModel (with random weights) from the `"facebook/sam_hq-vit-huge"` style configuration
>>> model = SamHQVisionModel(configuration)

>>> # Accessing the model configuration
>>> configuration = model.config

SamHQMaskDecoderConfig

class transformers.SamHQMaskDecoderConfig

< >

( transformers_version: str | None = None architectures: list[str] | None = None output_hidden_states: bool | None = False return_dict: bool | None = True dtype: typing.Union[str, ForwardRef('torch.dtype'), NoneType] = None chunk_size_feed_forward: int = 0 is_encoder_decoder: bool = False id2label: dict[int, str] | dict[str, str] | None = None label2id: dict[str, int] | dict[str, str] | None = None problem_type: typing.Optional[typing.Literal['regression', 'single_label_classification', 'multi_label_classification']] = None hidden_size: int = 256 hidden_act: str = 'relu' mlp_dim: int = 2048 num_hidden_layers: int = 2 num_attention_heads: int = 8 attention_downsample_rate: int = 2 num_multimask_outputs: int = 3 iou_head_depth: int = 3 iou_head_hidden_dim: int = 256 layer_norm_eps: float = 1e-06 vit_dim: int = 768 )

参数

  • hidden_size (int, optional, 默认为 256) — 隐藏表示的维度。
  • hidden_act (str, optional, 默认为 relu) — 解码器中的非线性激活函数(函数或字符串)。例如:"gelu", "relu", "silu" 等。
  • mlp_dim (int, optional, 默认为 2048) — Transformer 编码器中“中间”(即前馈)层的维度。
  • num_hidden_layers (int, optional, 默认为 2) — Transformer 解码器中隐藏层的数量。
  • num_attention_heads (int, optional, 默认为 8) — Transformer 解码器中每个注意力层的注意力头数。
  • attention_downsample_rate (int, optional, 默认为 2) — 注意力层的下采样率。
  • num_multimask_outputs (int, optional, 默认为 3) — SamMaskDecoder 模块的输出数量。在 Segment Anything 论文中,此值设为 3。
  • iou_head_depth (int, optional, 默认为 3) — IoU 头模块中的层数。
  • iou_head_hidden_dim (int, optional, 默认为 256) — IoU 头模块中隐藏状态的维度。
  • layer_norm_eps (float, optional, 默认为 1e-06) — 层归一化层使用的 epsilon 值。
  • vit_dim (int, optional, 默认为 768) — SamHQMaskDecoder 模块中使用的视觉 Transformer (ViT) 的维度。

这是用于存储 SamHQModel 配置的类。它用于根据指定的参数实例化 Sam Hq 模型,定义模型架构。使用默认值实例化配置将产生与 syscv-community/sam-hq-vit-base 类似的配置。

配置对象继承自 PreTrainedConfig,可用于控制模型输出。阅读 PreTrainedConfig 的文档以获取更多信息。

SamHQPromptEncoderConfig

class transformers.SamHQPromptEncoderConfig

< >

( transformers_version: str | None = None architectures: list[str] | None = None output_hidden_states: bool | None = False return_dict: bool | None = True dtype: typing.Union[str, ForwardRef('torch.dtype'), NoneType] = None chunk_size_feed_forward: int = 0 is_encoder_decoder: bool = False id2label: dict[int, str] | dict[str, str] | None = None label2id: dict[str, int] | dict[str, str] | None = None problem_type: typing.Optional[typing.Literal['regression', 'single_label_classification', 'multi_label_classification']] = None hidden_size: int = 256 image_size: int | list[int] | tuple[int, int] = 1024 patch_size: int | list[int] | tuple[int, int] = 16 mask_input_channels: int = 16 num_point_embeddings: int = 4 hidden_act: str = 'gelu' layer_norm_eps: float = 1e-06 )

参数

  • hidden_size (int, optional, 默认为 256) — 隐藏表示的维度。
  • image_size (Union[int, list[int], tuple[int, int]], optional, 默认为 1024) — 每张图像的大小(分辨率)。
  • patch_size (Union[int, list[int], tuple[int, int]], optional, 默认为 16) — 每个切片(patch)的大小(分辨率)。
  • mask_input_channels (int, 可选, 默认为 16) — 馈送到 MaskDecoder 模块的通道数。
  • num_point_embeddings (int, 可选, 默认为 4) — 使用的点嵌入(point embeddings)数量。
  • hidden_act (str, 可选, 默认为 gelu) — 解码器中的非线性激活函数(函数或字符串)。例如:"gelu", "relu", "silu" 等。
  • layer_norm_eps (float, 可选, 默认为 1e-06) — 层归一化层使用的 epsilon 值。

这是用于存储 SamHQModel 配置的类。它用于根据指定的参数实例化 Sam Hq 模型,定义模型架构。使用默认值实例化配置将产生与 syscv-community/sam-hq-vit-base 类似的配置。

配置对象继承自 PreTrainedConfig,可用于控制模型输出。阅读 PreTrainedConfig 的文档以获取更多信息。

SamHQProcessor

class transformers.SamHQProcessor

< >

( image_processor )

参数

  • image_processor (SamImageProcessor) — 图像处理器是必须的输入项。

构建一个将图像处理器包装在单个处理器中的 SamHQProcessor。

SamHQProcessor 提供了 SamImageProcessor 的所有功能。有关更多信息,请参阅 ~SamImageProcessor

__call__

< >

( images: typing.Union[ForwardRef('PIL.Image.Image'), numpy.ndarray, ForwardRef('torch.Tensor'), list['PIL.Image.Image'], list[numpy.ndarray], list['torch.Tensor'], NoneType] = None **kwargs: typing_extensions.Unpack[transformers.models.sam_hq.processing_sam_hq.SamHQProcessorKwargs] ) ~feature_extraction_utils.BatchFeature

参数

  • images (Union[PIL.Image.Image, numpy.ndarray, torch.Tensor, list[PIL.Image.Image], list[numpy.ndarray], list[torch.Tensor]], 可选) — 待预处理的图像。期望输入单个图像或一批图像,像素值范围在 0 到 255 之间。如果传入像素值在 0 到 1 之间的图像,请设置 do_rescale=False
  • segmentation_maps (ImageInput, kwargs, 可选) — 与输入图像一起处理的地面真实(ground truth)分割图。这些图用于训练或评估目的,会被调整大小并归一化以匹配处理后的图像尺寸。
  • input_points (NestedList, kwargs, 可选) — 用于基于提示(prompt-based)分割的输入点。应为一个嵌套列表,结构为 [image_level, object_level, point_level, [x, y]],其中每个点指定为原始图像空间中的 [x, y] 坐标。点在传给模型之前会归一化到目标图像尺寸。
  • input_labels (NestedList, kwargs, 可选) — 输入点的标签,指示每个点是前景(1)还是背景(0)点。应为一个结构为 [image_level, object_level, point_level] 的嵌套列表。其结构必须与 input_points 相同(不包括坐标维度)。
  • input_boxes (NestedList, kwargs, 可选) — 用于基于提示分割的边界框。应为一个结构为 [image_level, box_level, [x1, y1, x2, y2]] 的嵌套列表,其中每个框指定为原始图像空间中的 [x1, y1, x2, y2] 坐标。框在传给模型之前会归一化到目标图像尺寸。
  • point_pad_value (int, kwargs, 可选, 默认为 None) — 在批量处理不同长度的序列时,用于填充输入点的值。该值标记了填充位置,并在坐标归一化期间予以保留,以区分真实点和填充点。如果为 None,则使用处理器配置中的默认填充值。
  • mask_size (dict[str, *kwargs*, int], 可选) — 指定目标掩码大小的字典,包含键 "height""width"。这决定了模型生成的输出分割掩码的分辨率。
  • mask_pad_size (dict[str, *kwargs*, int], 可选) — 指定掩码填充大小的字典,包含键 "height""width"。在批量处理不同尺寸的掩码时使用,以确保维度一致。
  • return_tensors (strTensorType, 可选) — 如果设置,将返回特定框架的张量。可接受的值为:

    • 'pt': 返回 PyTorch torch.Tensor 对象。
    • 'np': 返回 NumPy np.ndarray 对象。
  • **kwargs (ProcessingKwargs, 可选) — 每个模态(文本、图像、视频、音频)的附加处理选项。模型特定的参数列在上面;请参阅 TypedDict 类以获取支持参数的完整列表。

返回

~feature_extraction_utils.BatchFeature

  • data (dict, optional) — 由 **call** / pad 方法返回的列表/数组/张量字典(“input_values”、“attention_mask”等)。
  • tensor_type (Union[None, str, TensorType], optional) — 您可以在此处提供 tensor_type 以在初始化时将整数列表转换为 PyTorch/Numpy 张量。
  • skip_tensor_conversion (list[str] or set[str], optional) — 不应转换为张量的键列表或集合,即使指定了 tensor_type 也是如此。

SamHQVisionModel

class transformers.SamHQVisionModel

< >

( config: SamHQVisionConfig )

参数

  • config (SamHQVisionConfig) — 模型配置类,包含模型的所有参数。使用配置文件初始化模型不会加载与模型相关的权重,仅加载配置。请查看 from_pretrained() 方法以加载模型权重。

来自 SamHQ 的视觉模型,顶部没有任何头部(head)或投影层。

该模型继承自 PreTrainedModel。请查看超类文档以了解该库为所有模型实现的通用方法(例如下载或保存、调整输入嵌入大小、剪枝头部等)。

此模型也是一个 PyTorch torch.nn.Module 子类。像普通的 PyTorch Module 一样使用它,并参考 PyTorch 文档了解一般用法和行为的所有相关信息。

forward

< >

( pixel_values: torch.FloatTensor | None = None **kwargs: typing_extensions.Unpack[transformers.utils.generic.TransformersKwargs] ) SamHQVisionEncoderOutputtuple(torch.FloatTensor)

参数

  • pixel_values (形状为 (batch_size, num_channels, image_size, image_size)torch.FloatTensor, 可选) — 对应于输入图像的张量。像素值可以使用 SamImageProcessor 获取。有关详细信息,请参阅 SamImageProcessor.__call__()SamHQProcessor 使用 SamImageProcessor 处理图像)。

返回

SamHQVisionEncoderOutputtuple(torch.FloatTensor)

一个 SamHQVisionEncoderOutputtorch.FloatTensor 元组(如果传递了 return_dict=False 或当 config.return_dict=False 时),根据配置(SamHQConfig)和输入,包含各种元素。

SamHQVisionModel 的前向传播方法,覆盖了 __call__ 特殊方法。

虽然 forward pass 的实现需要在此函数中定义,但你应该在之后调用 Module 实例而不是这个,因为前者负责运行预处理和后处理步骤,而后者会静默地忽略它们。

  • image_embeds (torch.FloatTensor, shape (batch_size, output_dim), 当模型初始化时 with_projection=True 返回可选) — 通过将投影层应用于 pooler_output 得到的图像嵌入。

  • last_hidden_state (形状为 (batch_size, sequence_length, hidden_size)torch.FloatTensor, 可选,默认为 None) — 模型最后一层输出的隐藏状态序列。

  • hidden_states (tuple[torch.FloatTensor, ...]可选,当传递 output_hidden_states=Trueconfig.output_hidden_states=True 时返回) — torch.FloatTensor 的元组(如果模型有嵌入层,则第一个为嵌入输出,其余为每一层的输出),形状为 (batch_size, sequence_length, hidden_size)

    模型在每个层输出的隐藏状态以及可选的初始嵌入输出。

  • attentions (tuple[torch.FloatTensor, ...]可选,当传递 output_attentions=Trueconfig.output_attentions=True 时返回) — torch.FloatTensor 的元组(每层一个),形状为 (batch_size, num_heads, sequence_length, sequence_length)

    注意力 softmax 后的注意力权重,用于计算自注意力头中的加权平均值。

  • intermediate_embeddings (list(torch.FloatTensor), 可选) — 从模型内特定块(通常是不带窗口注意力机制的块)收集的中间嵌入列表。列表中的每个元素形状为 (batch_size, sequence_length, hidden_size)。这是 SAM-HQ 特有的,基础 SAM 中没有此项。

SamHQModel

class transformers.SamHQModel

< >

( config model_args: ~utils.generic.ModelArgs | None = None adapter_args: ~utils.generic.AdapterArgs | None = None lora_args: ~utils.generic.LoRAArgs | None = None tokenizer_args: ~utils.generic.TokenizerArgs | None = None dataset_args: ~utils.generic.DatasetArgs | None = None data_args: ~utils.generic.DataArgs | None = None training_args: ~utils.generic.TrainingArgs | None = None generation_args: ~utils.generic.GenerationArgs | None = None vision_tower_args: ~utils.generic.VisionTowerArgs | None = None qlora_args: ~utils.generic.QLoRAArgs | None = None vision_tower_template_args: ~utils.generic.VisionTowerTemplateArgs | None = None video_tower_args: ~utils.generic.VideoTowerArgs | None = None vision_config: ~utils.generic.VisionConfig | None = None video_config: ~utils.generic.VideoConfig | None = None load_dataset: bool | None = None load_data_collator: bool | None = None load_processor: bool | None = None load_lora_adapter: bool | None = None load_adapter: bool | None = None load_qlora_adapter: bool | None = None **kwargs: typing_extensions.Unpack[transformers.modeling_utils.PreTrainedModelKwargs] )

参数

  • config (SamHQModel) — 模型配置类,包含模型的所有参数。使用配置文件初始化模型不会加载与模型相关的权重,仅加载配置。请查看 from_pretrained() 方法以加载模型权重。

Segment Anything Model HQ (SAM-HQ),用于根据输入图像以及可选的二维位置和边界框生成掩码。

该模型继承自 PreTrainedModel。请查看超类文档以了解该库为所有模型实现的通用方法(例如下载或保存、调整输入嵌入大小、剪枝头部等)。

此模型也是一个 PyTorch torch.nn.Module 子类。像普通的 PyTorch Module 一样使用它,并参考 PyTorch 文档了解一般用法和行为的所有相关信息。

forward

< >

( pixel_values: torch.FloatTensor | None = None input_points: torch.FloatTensor | None = None input_labels: torch.LongTensor | None = None input_boxes: torch.FloatTensor | None = None input_masks: torch.LongTensor | None = None image_embeddings: torch.FloatTensor | None = None multimask_output: bool = True hq_token_only: bool = False attention_similarity: torch.FloatTensor | None = None target_embedding: torch.FloatTensor | None = None intermediate_embeddings: list[torch.FloatTensor] | None = None **kwargs: typing_extensions.Unpack[transformers.utils.generic.TransformersKwargs] )

参数

  • pixel_values (torch.FloatTensor,形状为 (batch_size, num_channels, image_size, image_size)可选) — 与输入图像对应的张量。像素值可以使用 SamImageProcessor 获取。有关详细信息,请参阅 SamImageProcessor.__call__()SamHQProcessor 使用 SamImageProcessor 来处理图像)。
  • input_points (torch.FloatTensor,形状为 (batch_size, num_points, 2)) — 输入的二维空间点,供提示编码器 (prompt encoder) 用于对提示进行编码。通常能带来更好的效果。这些点可以通过向处理器传入嵌套列表来获取,处理器将创建对应的四维 torch 张量。第一维是图像批次大小,第二维是点批次大小(即我们希望模型为每个输入点预测多少个分割掩码),第三维是每个分割掩码的点数(可以为单个掩码传入多个点),最后一维是点的 x(垂直)和 y(水平)坐标。如果每张图像或每个掩码传入的点数不同,处理器将创建“PAD”点(对应 (0, 0) 坐标),且计算嵌入时会通过标签跳过这些点。
  • input_labels (torch.LongTensor,形状为 (batch_size, point_batch_size, num_points)) — 点的输入标签,供提示编码器用于对提示进行编码。根据官方实现,标签有 3 种类型:

    • 1:该点是包含目标对象的点
    • 0:该点是不包含目标对象的点
    • -1:该点对应于背景

    我们添加了以下标签:

    • -10:该点是填充点,因此应被提示编码器忽略

    填充标签应由处理器自动处理。

  • input_boxes (torch.FloatTensor,形状为 (batch_size, num_boxes, 4)) — 点的输入框,供提示编码器用于对提示进行编码。通常能生成更好的掩码。框可以通过向处理器传入嵌套列表来获取,该处理器将生成一个 torch 张量,各维度分别对应图像批次大小、每张图像的框数以及框的左上角和右下角坐标。顺序为 (x1, y1, x2, y2):

    • x1:输入框左上角的 x 坐标
    • y1:输入框左上角的 y 坐标
    • x2:输入框右下角的 x 坐标
    • y2:输入框右下角的 y 坐标
  • input_masks (torch.FloatTensor,形状为 (batch_size, image_size, image_size)) — SAM_HQ 模型也接受分割掩码作为输入。掩码将由提示编码器嵌入以生成相应的嵌入,之后输入掩码解码器。这些掩码需要由用户手动输入,并且形状必须为 (batch_size, image_size, image_size)。
  • image_embeddings (torch.FloatTensor,形状为 (batch_size, output_channels, window_size, window_size)) — 图像嵌入,由掩码解码器用于生成掩码和 IOU 分数。为了获得更高效的内存计算,用户可以首先使用 get_image_embeddings 方法获取图像嵌入,然后将其输入到 forward 方法中,而不是输入 pixel_values
  • multimask_output (bool可选) — 在原始实现和论文中,模型始终为每张图像(或相关的每个点/每个边界框)输出 3 个掩码。但是,可以通过指定 multimask_output=False 只输出单个对应“最佳”掩码的掩码。
  • hq_token_only (bool可选,默认为 False) — 是否仅使用 HQ token 路径进行掩码生成。当为 False 时,会同时结合标准路径和 HQ 路径。这是 SAM-HQ 架构所特有的。
  • attention_similarity (torch.FloatTensor可选) — 注意力相似度张量,在模型用于 PerSAM 中介绍的个性化场景时,提供给掩码解码器以进行目标引导的注意力计算。
  • target_embedding (torch.FloatTensor可选) — 目标概念的嵌入,在模型用于 PerSAM 中介绍的个性化场景时,提供给掩码解码器以进行目标语义提示。
  • intermediate_embeddings (List[torch.FloatTensor]可选) — 来自视觉编码器非窗口化块的中间嵌入,被 SAM-HQ 用于增强掩码质量。当提供预计算的 image_embeddings 而不是 pixel_values 时,此项是必须的。

SamHQModel 的 forward 方法,重写了 __call__ 特殊方法。

虽然 forward pass 的实现需要在此函数中定义,但你应该在之后调用 Module 实例而不是这个,因为前者负责运行预处理和后处理步骤,而后者会静默地忽略它们。

示例

>>> from PIL import Image
>>> import httpx
>>> from io import BytesIO
>>> from transformers import AutoModel, AutoProcessor

>>> model = AutoModel.from_pretrained("sushmanth/sam_hq_vit_b")
>>> processor = AutoProcessor.from_pretrained("sushmanth/sam_hq_vit_b")

>>> url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/model_doc/sam-car.png"
>>> with httpx.stream("GET", url) as response:
...     image = Image.open(BytesIO(response.read())).convert("RGB")
>>> input_points = [[[400, 650]]]  # 2D location of a window on the car
>>> inputs = processor(images=image, input_points=input_points, return_tensors="pt")

>>> # Get high-quality segmentation mask
>>> outputs = model(**inputs)

>>> # For high-quality mask only
>>> outputs = model(**inputs, hq_token_only=True)

>>> # Postprocess masks
>>> masks = processor.post_process_masks(
...     outputs.pred_masks, inputs["original_sizes"], inputs["reshaped_input_sizes"]
... )
在 GitHub 上更新

© . This site is unofficial and not affiliated with Hugging Face, Inc.