Transformers 文档
SAM-HQ
并获得增强的文档体验
开始使用
该模型于 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 原有的提示设计、效率和零样本泛化能力的同时,生成质量显著更高的分割掩码。

与原始 SAM 模型相比,SAM-HQ 引入了几个关键改进:
- 高质量输出标记(High-Quality Output Token):注入到 SAM 掩码解码器中的可学习标记,用于实现更高质量的掩码预测。
- 全局-局部特征融合(Global-local Feature Fusion):结合了模型不同阶段的特征,以改善掩码细节。
- 训练数据:使用了一个精心筛选的 4.4 万个高质量掩码数据集,而非 SA-1B。
- 效率:仅增加了 0.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。
- 使用该模型的演示笔记本(即将推出)
- 论文实现和代码:SAM-HQ GitHub 仓库
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.configSamHQMaskDecoderConfig
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 )
构建一个将图像处理器包装在单个处理器中的 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 (
str或 TensorType, 可选) — 如果设置,将返回特定框架的张量。可接受的值为:'pt': 返回 PyTorchtorch.Tensor对象。'np': 返回 NumPynp.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]orset[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] ) → SamHQVisionEncoderOutput 或 tuple(torch.FloatTensor)
参数
- pixel_values (形状为
(batch_size, num_channels, image_size, image_size)的torch.FloatTensor, 可选) — 对应于输入图像的张量。像素值可以使用 SamImageProcessor 获取。有关详细信息,请参阅SamImageProcessor.__call__()(SamHQProcessor 使用 SamImageProcessor 处理图像)。
返回
SamHQVisionEncoderOutput 或 tuple(torch.FloatTensor)
一个 SamHQVisionEncoderOutput 或 torch.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=True或config.output_hidden_states=True时返回) —torch.FloatTensor的元组(如果模型有嵌入层,则第一个为嵌入输出,其余为每一层的输出),形状为(batch_size, sequence_length, hidden_size)。模型在每个层输出的隐藏状态以及可选的初始嵌入输出。
attentions (
tuple[torch.FloatTensor, ...],可选,当传递output_attentions=True或config.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"]
... )