initial 2.1 support
This commit is contained in:
+6
-8
@@ -41,14 +41,9 @@ def load_model(model_path, model_cfg_path, segmentor, dtype, device):
|
||||
fpn_interp_model=neck_config['fpn_interp_model']
|
||||
)
|
||||
|
||||
trunk = Hiera(
|
||||
embed_dim=trunk_config['embed_dim'],
|
||||
num_heads=trunk_config['num_heads'],
|
||||
stages=trunk_config['stages'],
|
||||
global_att_blocks=trunk_config['global_att_blocks'],
|
||||
window_pos_embed_bkg_spatial_size=trunk_config['window_pos_embed_bkg_spatial_size']
|
||||
|
||||
)
|
||||
keys_to_include = ['embed_dim', 'num_heads', 'global_att_blocks', 'window_pos_embed_bkg_spatial_size', 'stages']
|
||||
trunk_kwargs = {key: trunk_config[key] for key in keys_to_include if key in trunk_config}
|
||||
trunk = Hiera(**trunk_kwargs)
|
||||
|
||||
image_encoder = ImageEncoder(
|
||||
scalp=model_config['image_encoder']['scalp'],
|
||||
@@ -169,6 +164,9 @@ def load_model(model_path, model_cfg_path, segmentor, dtype, device):
|
||||
multimask_min_pt_num=model_config['multimask_min_pt_num'],
|
||||
multimask_max_pt_num=model_config['multimask_max_pt_num'],
|
||||
use_mlp_for_obj_ptr_proj=model_config['use_mlp_for_obj_ptr_proj'],
|
||||
proj_tpos_enc_in_obj_ptrs=model_config['proj_tpos_enc_in_obj_ptrs'],
|
||||
no_obj_embed_spatial=model_config['no_obj_embed_spatial'],
|
||||
use_signed_tpos_enc_to_obj_ptrs=model_config['use_signed_tpos_enc_to_obj_ptrs'],
|
||||
binarize_mask_from_pts_for_mem_enc=True if segmentor == 'video' else False,
|
||||
).to(dtype).to(device).eval()
|
||||
|
||||
|
||||
@@ -25,6 +25,10 @@ class DownloadAndLoadSAM2Model:
|
||||
'sam2_hiera_large.safetensors',
|
||||
'sam2_hiera_small.safetensors',
|
||||
'sam2_hiera_tiny.safetensors',
|
||||
'sam2.1_hiera_base_plus.safetensors',
|
||||
'sam2.1_hiera_large.safetensors',
|
||||
'sam2.1_hiera_small.safetensors',
|
||||
'sam2.1_hiera_tiny.safetensors',
|
||||
],),
|
||||
"segmentor": (
|
||||
['single_image','video', 'automaskgenerator'],
|
||||
@@ -32,7 +36,7 @@ class DownloadAndLoadSAM2Model:
|
||||
"device": (['cuda', 'cpu', 'mps'], ),
|
||||
"precision": ([ 'fp16','bf16','fp32'],
|
||||
{
|
||||
"default": 'bf16'
|
||||
"default": 'fp16'
|
||||
}),
|
||||
|
||||
},
|
||||
@@ -56,7 +60,11 @@ class DownloadAndLoadSAM2Model:
|
||||
device = {"cuda": torch.device("cuda"), "cpu": torch.device("cpu"), "mps": torch.device("mps")}[device]
|
||||
|
||||
download_path = os.path.join(folder_paths.models_dir, "sam2")
|
||||
if precision != 'fp32' and "2.1" in model:
|
||||
base_name, extension = model.rsplit('.', 1)
|
||||
model = f"{base_name}-fp16.{extension}"
|
||||
model_path = os.path.join(download_path, model)
|
||||
print("model_path: ", model_path)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading SAM2 model to: {model_path}")
|
||||
@@ -67,24 +75,36 @@ class DownloadAndLoadSAM2Model:
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
model_mapping = {
|
||||
"base": "sam2_hiera_b+.yaml",
|
||||
"large": "sam2_hiera_l.yaml",
|
||||
"small": "sam2_hiera_s.yaml",
|
||||
"tiny": "sam2_hiera_t.yaml"
|
||||
"2.0": {
|
||||
"base": "sam2_hiera_b+.yaml",
|
||||
"large": "sam2_hiera_l.yaml",
|
||||
"small": "sam2_hiera_s.yaml",
|
||||
"tiny": "sam2_hiera_t.yaml"
|
||||
},
|
||||
"2.1": {
|
||||
"base": "sam2.1_hiera_b+.yaml",
|
||||
"large": "sam2.1_hiera_l.yaml",
|
||||
"small": "sam2.1_hiera_s.yaml",
|
||||
"tiny": "sam2.1_hiera_t.yaml"
|
||||
}
|
||||
}
|
||||
version = "2.1" if "2.1" in model else "2.0"
|
||||
|
||||
model_cfg_path = next(
|
||||
(os.path.join(script_directory, "sam2_configs", cfg) for key, cfg in model_mapping.items() if key in model),
|
||||
(os.path.join(script_directory, "sam2_configs", cfg)
|
||||
for key, cfg in model_mapping[version].items() if key in model),
|
||||
None
|
||||
)
|
||||
)
|
||||
print(f"Using model config: {model_cfg_path}")
|
||||
|
||||
model =load_model(model_path, model_cfg_path, segmentor, dtype, device)
|
||||
model = load_model(model_path, model_cfg_path, segmentor, dtype, device)
|
||||
|
||||
sam2_model = {
|
||||
'model': model,
|
||||
'dtype': dtype,
|
||||
'device': device,
|
||||
'segmentor' : segmentor
|
||||
'segmentor' : segmentor,
|
||||
'version': version
|
||||
}
|
||||
|
||||
return (sam2_model,)
|
||||
@@ -197,8 +217,8 @@ class Sam2Segmentation:
|
||||
if segmentor == 'single_image' and B > 1:
|
||||
print("Segmenting batch of images with single_image segmentor")
|
||||
|
||||
if segmentor == 'video' and bboxes is not None:
|
||||
raise ValueError("Video segmentor doesn't support bboxes")
|
||||
if segmentor == 'video' and bboxes is not None and "2.1" not in sam2_model["version"]:
|
||||
raise ValueError("2.0 model doesn't support bboxes with video segmentor")
|
||||
|
||||
if segmentor == 'video': # video model needs images resized first thing
|
||||
model_input_image_size = model.image_size
|
||||
@@ -329,23 +349,35 @@ class Sam2Segmentation:
|
||||
if hasattr(self, 'inference_state'):
|
||||
model.reset_state(self.inference_state)
|
||||
self.inference_state = model.init_state(image.permute(0, 3, 1, 2).contiguous(), H, W, device=device)
|
||||
if bboxes is None:
|
||||
input_box = None
|
||||
else:
|
||||
input_box = bboxes[0]
|
||||
|
||||
if individual_objects and bboxes is not None:
|
||||
raise ValueError("bboxes not supported with individual_objects")
|
||||
|
||||
|
||||
if individual_objects:
|
||||
for i, (coord, label) in enumerate(zip(final_coords, final_labels)):
|
||||
_, out_obj_ids, out_mask_logits = model.add_new_points(
|
||||
_, out_obj_ids, out_mask_logits = model.add_new_points_or_box(
|
||||
inference_state=self.inference_state,
|
||||
frame_idx=0,
|
||||
obj_id=i,
|
||||
points=final_coords[i],
|
||||
labels=final_labels[i],
|
||||
clear_old_points=True,
|
||||
box=input_box
|
||||
)
|
||||
else:
|
||||
_, out_obj_ids, out_mask_logits = model.add_new_points(
|
||||
_, out_obj_ids, out_mask_logits = model.add_new_points_or_box(
|
||||
inference_state=self.inference_state,
|
||||
frame_idx=0,
|
||||
obj_id=1,
|
||||
points=final_coords,
|
||||
labels=final_labels,
|
||||
points=final_coords if coordinates_positive is not None else None,
|
||||
labels=final_labels if coordinates_positive is not None else None,
|
||||
clear_old_points=True,
|
||||
box=input_box
|
||||
)
|
||||
|
||||
pbar = ProgressBar(B)
|
||||
@@ -673,6 +705,48 @@ class Sam2AutoSegmentation:
|
||||
|
||||
mask_tensor = torch.stack(out_list, dim=0)
|
||||
return (mask_tensor.cpu().float(), segment_image_tensor.cpu().float(), bbox_list)
|
||||
|
||||
#WIP
|
||||
# class OwlV2Detector:
|
||||
# @classmethod
|
||||
# def INPUT_TYPES(s):
|
||||
# return {
|
||||
# "required": {
|
||||
# "image": ("IMAGE", ),
|
||||
# },
|
||||
# }
|
||||
|
||||
# RETURN_TYPES = ("MASK", )
|
||||
# RETURN_NAMES =("mask", )
|
||||
# FUNCTION = "segment"
|
||||
# CATEGORY = "SAM2"
|
||||
|
||||
# def segment(self, image):
|
||||
# from transformers import Owlv2Processor, Owlv2ForObjectDetection
|
||||
# device = mm.get_torch_device()
|
||||
# offload_device = mm.unet_offload_device()
|
||||
# processor = Owlv2Processor.from_pretrained("google/owlv2-base-patch16-ensemble")
|
||||
# model = Owlv2ForObjectDetection.from_pretrained("google/owlv2-base-patch16-ensemble")
|
||||
|
||||
# url = "http://images.cocodataset.org/val2017/000000039769.jpg"
|
||||
# image = Image.open(requests.get(url, stream=True).raw)
|
||||
# texts = [["a photo of a cat", "a photo of a dog"]]
|
||||
# inputs = processor(text=texts, images=image, return_tensors="pt")
|
||||
# outputs = model(**inputs)
|
||||
|
||||
# # Target image sizes (height, width) to rescale box predictions [batch_size, 2]
|
||||
# target_sizes = torch.Tensor([image.size[::-1]])
|
||||
# # Convert outputs (bounding boxes and class logits) to Pascal VOC Format (xmin, ymin, xmax, ymax)
|
||||
# results = processor.post_process_object_detection(outputs=outputs, target_sizes=target_sizes, threshold=0.1)
|
||||
# i = 0 # Retrieve predictions for the first image for the corresponding text queries
|
||||
# text = texts[i]
|
||||
# boxes, scores, labels = results[i]["boxes"], results[i]["scores"], results[i]["labels"]
|
||||
# for box, score, label in zip(boxes, scores, labels):
|
||||
# box = [round(i, 2) for i in box.tolist()]
|
||||
# print(f"Detected {text[label]} with confidence {round(score.item(), 3)} at location {box}")
|
||||
|
||||
|
||||
# return (mask_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadSAM2Model": DownloadAndLoadSAM2Model,
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import List, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from iopath.common.file_io import g_pathmgr
|
||||
|
||||
from ....sam2.modeling.backbones.utils import (
|
||||
PatchEmbed,
|
||||
@@ -46,11 +47,7 @@ class MultiScaleAttention(nn.Module):
|
||||
|
||||
self.dim = dim
|
||||
self.dim_out = dim_out
|
||||
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim_out // num_heads
|
||||
self.scale = head_dim**-0.5
|
||||
|
||||
self.q_pool = q_pool
|
||||
self.qkv = nn.Linear(dim, dim_out * 3)
|
||||
self.proj = nn.Linear(dim_out, dim_out)
|
||||
@@ -197,6 +194,7 @@ class Hiera(nn.Module):
|
||||
16,
|
||||
20,
|
||||
),
|
||||
weights_path=None,
|
||||
return_interm_layers=True, # return feats from every stage
|
||||
):
|
||||
super().__init__()
|
||||
@@ -266,6 +264,11 @@ class Hiera(nn.Module):
|
||||
else [self.blocks[-1].dim_out]
|
||||
)
|
||||
|
||||
if weights_path is not None:
|
||||
with g_pathmgr.open(weights_path, "rb") as f:
|
||||
chkpt = torch.load(f, map_location="cpu")
|
||||
logging.info("loading Hiera", self.load_state_dict(chkpt, strict=False))
|
||||
|
||||
def _get_pos_embed(self, hw: Tuple[int, int]) -> torch.Tensor:
|
||||
h, w = hw
|
||||
window_embed = self.pos_embed_window
|
||||
@@ -293,3 +296,21 @@ class Hiera(nn.Module):
|
||||
outputs.append(feats)
|
||||
|
||||
return outputs
|
||||
|
||||
def get_layer_id(self, layer_name):
|
||||
# https://github.com/microsoft/unilm/blob/master/beit/optim_factory.py#L33
|
||||
num_layers = self.get_num_layers()
|
||||
|
||||
if layer_name.find("rel_pos") != -1:
|
||||
return num_layers + 1
|
||||
elif layer_name.find("pos_embed") != -1:
|
||||
return 0
|
||||
elif layer_name.find("patch_embed") != -1:
|
||||
return 0
|
||||
elif layer_name.find("blocks") != -1:
|
||||
return int(layer_name.split("blocks")[1].split(".")[1]) + 1
|
||||
else:
|
||||
return num_layers + 1
|
||||
|
||||
def get_num_layers(self) -> int:
|
||||
return len(self.blocks)
|
||||
|
||||
@@ -71,6 +71,7 @@ class FpnNeck(nn.Module):
|
||||
self.position_encoding = position_encoding
|
||||
self.convs = nn.ModuleList()
|
||||
self.backbone_channel_list = backbone_channel_list
|
||||
self.d_model = d_model
|
||||
for dim in backbone_channel_list:
|
||||
current = nn.Sequential()
|
||||
current.add_module(
|
||||
|
||||
@@ -247,7 +247,7 @@ class MaskDecoder(nn.Module):
|
||||
def _get_stability_scores(self, mask_logits):
|
||||
"""
|
||||
Compute stability scores of the mask logits based on the IoU between upper and
|
||||
lower thresholds, similar to https://github.com/fairinternal/onevision/pull/568.
|
||||
lower thresholds.
|
||||
"""
|
||||
mask_logits = mask_logits.flatten(-2)
|
||||
stability_delta = self.dynamic_multimask_stability_delta
|
||||
|
||||
+154
-76
@@ -59,9 +59,6 @@ class SAM2Base(torch.nn.Module):
|
||||
# For r>1, the (self.num_maskmem - 1) non-conditioning memory frames consist of
|
||||
# (self.num_maskmem - 2) nearest frames from every r-th frames, plus the last frame.
|
||||
memory_temporal_stride_for_eval=1,
|
||||
# if `add_all_frames_to_correct_as_cond` is True, we also append to the conditioning frame list any frame that receives a later correction click
|
||||
# if `add_all_frames_to_correct_as_cond` is False, we conditioning frame list to only use those initial conditioning frames
|
||||
add_all_frames_to_correct_as_cond=False,
|
||||
# whether to apply non-overlapping constraints on the object masks in the memory encoder during evaluation (to avoid/alleviate superposing masks)
|
||||
non_overlap_masks_for_mem_enc=False,
|
||||
# whether to cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
@@ -73,6 +70,9 @@ class SAM2Base(torch.nn.Module):
|
||||
# whether to add an extra linear projection layer for the temporal positional encoding in the object pointers to avoid potential interference
|
||||
# with spatial positional encoding (only relevant when both `use_obj_ptrs_in_encoder=True` and `add_tpos_enc_to_obj_ptrs=True`)
|
||||
proj_tpos_enc_in_obj_ptrs=False,
|
||||
# whether to use signed distance (instead of unsigned absolute distance) in the temporal positional encoding in the object pointers
|
||||
# (only relevant when both `use_obj_ptrs_in_encoder=True` and `add_tpos_enc_to_obj_ptrs=True`)
|
||||
use_signed_tpos_enc_to_obj_ptrs=False,
|
||||
# whether to only attend to object pointers in the past (before the current frame) in the encoder during evaluation
|
||||
# (only relevant when `use_obj_ptrs_in_encoder=True`; this might avoid pointer information too far in the future to distract the initial tracking)
|
||||
only_obj_ptrs_in_the_past_for_eval=False,
|
||||
@@ -88,6 +88,8 @@ class SAM2Base(torch.nn.Module):
|
||||
# hope to make recovery easier if there is a mistake and mitigate accumulation of errors
|
||||
soft_no_obj_ptr: bool = False,
|
||||
use_mlp_for_obj_ptr_proj: bool = False,
|
||||
# add no obj embedding to spatial frames
|
||||
no_obj_embed_spatial: bool = False,
|
||||
# extra arguments used to construct the SAM mask decoder; if not None, it should be a dict of kwargs to be passed into `MaskDecoder` class.
|
||||
sam_mask_decoder_extra_args=None,
|
||||
compile_image_encoder: bool = False,
|
||||
@@ -110,12 +112,13 @@ class SAM2Base(torch.nn.Module):
|
||||
if proj_tpos_enc_in_obj_ptrs:
|
||||
assert add_tpos_enc_to_obj_ptrs # these options need to be used together
|
||||
self.proj_tpos_enc_in_obj_ptrs = proj_tpos_enc_in_obj_ptrs
|
||||
self.use_signed_tpos_enc_to_obj_ptrs = use_signed_tpos_enc_to_obj_ptrs
|
||||
self.only_obj_ptrs_in_the_past_for_eval = only_obj_ptrs_in_the_past_for_eval
|
||||
|
||||
# Part 2: memory attention to condition current frame's visual features
|
||||
# with memories (and obj ptrs) from past frames
|
||||
self.memory_attention = memory_attention
|
||||
self.hidden_dim = memory_attention.d_model
|
||||
self.hidden_dim = image_encoder.neck.d_model
|
||||
|
||||
# Part 3: memory encoder for the previous frame's outputs
|
||||
self.memory_encoder = memory_encoder
|
||||
@@ -170,9 +173,12 @@ class SAM2Base(torch.nn.Module):
|
||||
self.no_obj_ptr = torch.nn.Parameter(torch.zeros(1, self.hidden_dim))
|
||||
trunc_normal_(self.no_obj_ptr, std=0.02)
|
||||
self.use_mlp_for_obj_ptr_proj = use_mlp_for_obj_ptr_proj
|
||||
self.no_obj_embed_spatial = None
|
||||
if no_obj_embed_spatial:
|
||||
self.no_obj_embed_spatial = torch.nn.Parameter(torch.zeros(1, self.mem_dim))
|
||||
trunc_normal_(self.no_obj_embed_spatial, std=0.02)
|
||||
|
||||
self._build_sam_heads()
|
||||
self.add_all_frames_to_correct_as_cond = add_all_frames_to_correct_as_cond
|
||||
self.max_cond_frames_in_attn = max_cond_frames_in_attn
|
||||
|
||||
# Model compilation
|
||||
@@ -194,8 +200,8 @@ class SAM2Base(torch.nn.Module):
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Please use the corresponding methods in SAM2VideoPredictor for inference."
|
||||
"See notebooks/video_predictor_example.ipynb for an example."
|
||||
"Please use the corresponding methods in SAM2VideoPredictor for inference or SAM2Train for training/fine-tuning"
|
||||
"See notebooks/video_predictor_example.ipynb for an inference example."
|
||||
)
|
||||
|
||||
def _build_sam_heads(self):
|
||||
@@ -388,8 +394,6 @@ class SAM2Base(torch.nn.Module):
|
||||
if self.pred_obj_scores:
|
||||
# Allow *soft* no obj ptr, unlike for masks
|
||||
if self.soft_no_obj_ptr:
|
||||
# Only hard possible with gt
|
||||
assert not self.teacher_force_obj_scores_for_mem
|
||||
lambda_is_obj_appearing = object_score_logits.sigmoid()
|
||||
else:
|
||||
lambda_is_obj_appearing = is_obj_appearing.float()
|
||||
@@ -513,6 +517,7 @@ class SAM2Base(torch.nn.Module):
|
||||
return pix_feat
|
||||
|
||||
num_obj_ptr_tokens = 0
|
||||
tpos_sign_mul = -1 if track_in_reverse else 1
|
||||
# Step 1: condition the visual features of the current frame on previous memories
|
||||
if not is_init_cond_frame:
|
||||
# Retrieve the memories encoded with the maskmem backbone
|
||||
@@ -528,9 +533,9 @@ class SAM2Base(torch.nn.Module):
|
||||
t_pos_and_prevs = [(0, out) for out in selected_cond_outputs.values()]
|
||||
# Add last (self.num_maskmem - 1) frames before current frame for non-conditioning memory
|
||||
# the earliest one has t_pos=1 and the latest one has t_pos=self.num_maskmem-1
|
||||
# We also allow taking the memory frame non-consecutively (with r>1), in which case
|
||||
# we take (self.num_maskmem - 2) frames among every r-th frames plus the last frame.
|
||||
r = self.memory_temporal_stride_for_eval
|
||||
# We also allow taking the memory frame non-consecutively (with stride>1), in which case
|
||||
# we take (self.num_maskmem - 2) frames among every stride-th frames plus the last frame.
|
||||
stride = 1 if self.training else self.memory_temporal_stride_for_eval
|
||||
for t_pos in range(1, self.num_maskmem):
|
||||
t_rel = self.num_maskmem - t_pos # how many frames before current frame
|
||||
if t_rel == 1:
|
||||
@@ -546,15 +551,15 @@ class SAM2Base(torch.nn.Module):
|
||||
if not track_in_reverse:
|
||||
# first find the nearest frame among every r-th frames before this frame
|
||||
# for r=1, this would be (frame_idx - 2)
|
||||
prev_frame_idx = ((frame_idx - 2) // r) * r
|
||||
prev_frame_idx = ((frame_idx - 2) // stride) * stride
|
||||
# then seek further among every r-th frames
|
||||
prev_frame_idx = prev_frame_idx - (t_rel - 2) * r
|
||||
prev_frame_idx = prev_frame_idx - (t_rel - 2) * stride
|
||||
else:
|
||||
# first find the nearest frame among every r-th frames after this frame
|
||||
# for r=1, this would be (frame_idx + 2)
|
||||
prev_frame_idx = -(-(frame_idx + 2) // r) * r
|
||||
prev_frame_idx = -(-(frame_idx + 2) // stride) * stride
|
||||
# then seek further among every r-th frames
|
||||
prev_frame_idx = prev_frame_idx + (t_rel - 2) * r
|
||||
prev_frame_idx = prev_frame_idx + (t_rel - 2) * stride
|
||||
out = output_dict["non_cond_frame_outputs"].get(prev_frame_idx, None)
|
||||
if out is None:
|
||||
# If an unselected conditioning frame is among the last (self.num_maskmem - 1)
|
||||
@@ -593,7 +598,14 @@ class SAM2Base(torch.nn.Module):
|
||||
ptr_cond_outputs = selected_cond_outputs
|
||||
pos_and_ptrs = [
|
||||
# Temporal pos encoding contains how far away each pointer is from current frame
|
||||
(abs(frame_idx - t), out["obj_ptr"])
|
||||
(
|
||||
(
|
||||
(frame_idx - t) * tpos_sign_mul
|
||||
if self.use_signed_tpos_enc_to_obj_ptrs
|
||||
else abs(frame_idx - t)
|
||||
),
|
||||
out["obj_ptr"],
|
||||
)
|
||||
for t, out in ptr_cond_outputs.items()
|
||||
]
|
||||
# Add up to (max_obj_ptrs_in_encoder - 1) non-conditioning frames before current frame
|
||||
@@ -642,7 +654,7 @@ class SAM2Base(torch.nn.Module):
|
||||
pix_feat_with_mem = pix_feat_with_mem.permute(1, 2, 0).view(B, C, H, W)
|
||||
return pix_feat_with_mem
|
||||
|
||||
# Use a dummy token on the first frame (to avoid emtpy memory input to tranformer encoder)
|
||||
# Use a dummy token on the first frame (to avoid empty memory input to tranformer encoder)
|
||||
to_cat_memory = [self.no_mem_embed.expand(1, B, self.mem_dim)]
|
||||
to_cat_memory_pos_embed = [self.no_mem_pos_enc.expand(1, B, self.mem_dim)]
|
||||
|
||||
@@ -666,6 +678,7 @@ class SAM2Base(torch.nn.Module):
|
||||
current_vision_feats,
|
||||
feat_sizes,
|
||||
pred_masks_high_res,
|
||||
object_score_logits,
|
||||
is_mask_from_pts,
|
||||
):
|
||||
"""Encode the current image and its prediction into a memory feature."""
|
||||
@@ -698,9 +711,104 @@ class SAM2Base(torch.nn.Module):
|
||||
)
|
||||
maskmem_features = maskmem_out["vision_features"]
|
||||
maskmem_pos_enc = maskmem_out["vision_pos_enc"]
|
||||
# add a no-object embedding to the spatial memory to indicate that the frame
|
||||
# is predicted to be occluded (i.e. no object is appearing in the frame)
|
||||
if self.no_obj_embed_spatial is not None:
|
||||
is_obj_appearing = (object_score_logits > 0).float()
|
||||
maskmem_features += (
|
||||
1 - is_obj_appearing[..., None, None]
|
||||
) * self.no_obj_embed_spatial[..., None, None].expand(
|
||||
*maskmem_features.shape
|
||||
)
|
||||
|
||||
return maskmem_features, maskmem_pos_enc
|
||||
|
||||
def _track_step(
|
||||
self,
|
||||
frame_idx,
|
||||
is_init_cond_frame,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
point_inputs,
|
||||
mask_inputs,
|
||||
output_dict,
|
||||
num_frames,
|
||||
track_in_reverse,
|
||||
prev_sam_mask_logits,
|
||||
):
|
||||
current_out = {"point_inputs": point_inputs, "mask_inputs": mask_inputs}
|
||||
# High-resolution feature maps for the SAM head, reshape (HW)BC => BCHW
|
||||
if len(current_vision_feats) > 1:
|
||||
high_res_features = [
|
||||
x.permute(1, 2, 0).view(x.size(1), x.size(2), *s)
|
||||
for x, s in zip(current_vision_feats[:-1], feat_sizes[:-1])
|
||||
]
|
||||
else:
|
||||
high_res_features = None
|
||||
if mask_inputs is not None and self.use_mask_input_as_output_without_sam:
|
||||
# When use_mask_input_as_output_without_sam=True, we directly output the mask input
|
||||
# (see it as a GT mask) without using a SAM prompt encoder + mask decoder.
|
||||
pix_feat = current_vision_feats[-1].permute(1, 2, 0)
|
||||
pix_feat = pix_feat.view(-1, self.hidden_dim, *feat_sizes[-1])
|
||||
sam_outputs = self._use_mask_as_output(
|
||||
pix_feat, high_res_features, mask_inputs
|
||||
)
|
||||
else:
|
||||
# fused the visual feature with previous memory features in the memory bank
|
||||
pix_feat = self._prepare_memory_conditioned_features(
|
||||
frame_idx=frame_idx,
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
current_vision_feats=current_vision_feats[-1:],
|
||||
current_vision_pos_embeds=current_vision_pos_embeds[-1:],
|
||||
feat_sizes=feat_sizes[-1:],
|
||||
output_dict=output_dict,
|
||||
num_frames=num_frames,
|
||||
track_in_reverse=track_in_reverse,
|
||||
)
|
||||
# apply SAM-style segmentation head
|
||||
# here we might feed previously predicted low-res SAM mask logits into the SAM mask decoder,
|
||||
# e.g. in demo where such logits come from earlier interaction instead of correction sampling
|
||||
# (in this case, any `mask_inputs` shouldn't reach here as they are sent to _use_mask_as_output instead)
|
||||
if prev_sam_mask_logits is not None:
|
||||
assert point_inputs is not None and mask_inputs is None
|
||||
mask_inputs = prev_sam_mask_logits
|
||||
multimask_output = self._use_multimask(is_init_cond_frame, point_inputs)
|
||||
sam_outputs = self._forward_sam_heads(
|
||||
backbone_features=pix_feat,
|
||||
point_inputs=point_inputs,
|
||||
mask_inputs=mask_inputs,
|
||||
high_res_features=high_res_features,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
|
||||
return current_out, sam_outputs, high_res_features, pix_feat
|
||||
|
||||
def _encode_memory_in_output(
|
||||
self,
|
||||
current_vision_feats,
|
||||
feat_sizes,
|
||||
point_inputs,
|
||||
run_mem_encoder,
|
||||
high_res_masks,
|
||||
object_score_logits,
|
||||
current_out,
|
||||
):
|
||||
if run_mem_encoder and self.num_maskmem > 0:
|
||||
high_res_masks_for_mem_enc = high_res_masks
|
||||
maskmem_features, maskmem_pos_enc = self._encode_new_memory(
|
||||
current_vision_feats=current_vision_feats,
|
||||
feat_sizes=feat_sizes,
|
||||
pred_masks_high_res=high_res_masks_for_mem_enc,
|
||||
object_score_logits=object_score_logits,
|
||||
is_mask_from_pts=(point_inputs is not None),
|
||||
)
|
||||
current_out["maskmem_features"] = maskmem_features
|
||||
current_out["maskmem_pos_enc"] = maskmem_pos_enc
|
||||
else:
|
||||
current_out["maskmem_features"] = None
|
||||
current_out["maskmem_pos_enc"] = None
|
||||
|
||||
def track_step(
|
||||
self,
|
||||
frame_idx,
|
||||
@@ -722,50 +830,20 @@ class SAM2Base(torch.nn.Module):
|
||||
# The previously predicted SAM mask logits (which can be fed together with new clicks in demo).
|
||||
prev_sam_mask_logits=None,
|
||||
):
|
||||
current_out = {"point_inputs": point_inputs, "mask_inputs": mask_inputs}
|
||||
# High-resolution feature maps for the SAM head, reshape (HW)BC => BCHW
|
||||
if len(current_vision_feats) > 1:
|
||||
high_res_features = [
|
||||
x.permute(1, 2, 0).view(x.size(1), x.size(2), *s)
|
||||
for x, s in zip(current_vision_feats[:-1], feat_sizes[:-1])
|
||||
]
|
||||
else:
|
||||
high_res_features = None
|
||||
if mask_inputs is not None and self.use_mask_input_as_output_without_sam:
|
||||
# When use_mask_input_as_output_without_sam=True, we directly output the mask input
|
||||
# (see it as a GT mask) without using a SAM prompt encoder + mask decoder.
|
||||
pix_feat = current_vision_feats[-1].permute(1, 2, 0)
|
||||
pix_feat = pix_feat.view(-1, self.hidden_dim, *feat_sizes[-1])
|
||||
sam_outputs = self._use_mask_as_output(
|
||||
pix_feat, high_res_features, mask_inputs
|
||||
)
|
||||
else:
|
||||
# fused the visual feature with previous memory features in the memory bank
|
||||
pix_feat_with_mem = self._prepare_memory_conditioned_features(
|
||||
frame_idx=frame_idx,
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
current_vision_feats=current_vision_feats[-1:],
|
||||
current_vision_pos_embeds=current_vision_pos_embeds[-1:],
|
||||
feat_sizes=feat_sizes[-1:],
|
||||
output_dict=output_dict,
|
||||
num_frames=num_frames,
|
||||
track_in_reverse=track_in_reverse,
|
||||
)
|
||||
# apply SAM-style segmentation head
|
||||
# here we might feed previously predicted low-res SAM mask logits into the SAM mask decoder,
|
||||
# e.g. in demo where such logits come from earlier interaction instead of correction sampling
|
||||
# (in this case, any `mask_inputs` shouldn't reach here as they are sent to _use_mask_as_output instead)
|
||||
if prev_sam_mask_logits is not None:
|
||||
assert point_inputs is not None and mask_inputs is None
|
||||
mask_inputs = prev_sam_mask_logits
|
||||
multimask_output = self._use_multimask(is_init_cond_frame, point_inputs)
|
||||
sam_outputs = self._forward_sam_heads(
|
||||
backbone_features=pix_feat_with_mem,
|
||||
point_inputs=point_inputs,
|
||||
mask_inputs=mask_inputs,
|
||||
high_res_features=high_res_features,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
current_out, sam_outputs, _, _ = self._track_step(
|
||||
frame_idx,
|
||||
is_init_cond_frame,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
point_inputs,
|
||||
mask_inputs,
|
||||
output_dict,
|
||||
num_frames,
|
||||
track_in_reverse,
|
||||
prev_sam_mask_logits,
|
||||
)
|
||||
|
||||
(
|
||||
_,
|
||||
_,
|
||||
@@ -773,28 +851,28 @@ class SAM2Base(torch.nn.Module):
|
||||
low_res_masks,
|
||||
high_res_masks,
|
||||
obj_ptr,
|
||||
_,
|
||||
object_score_logits,
|
||||
) = sam_outputs
|
||||
|
||||
current_out["pred_masks"] = low_res_masks
|
||||
current_out["pred_masks_high_res"] = high_res_masks
|
||||
current_out["obj_ptr"] = obj_ptr
|
||||
if not self.training:
|
||||
# Only add this in inference (to avoid unused param in activation checkpointing;
|
||||
# it's mainly used in the demo to encode spatial memories w/ consolidated masks)
|
||||
current_out["object_score_logits"] = object_score_logits
|
||||
|
||||
# Finally run the memory encoder on the predicted mask to encode
|
||||
# it into a new memory feature (that can be used in future frames)
|
||||
if run_mem_encoder and self.num_maskmem > 0:
|
||||
high_res_masks_for_mem_enc = high_res_masks
|
||||
maskmem_features, maskmem_pos_enc = self._encode_new_memory(
|
||||
current_vision_feats=current_vision_feats,
|
||||
feat_sizes=feat_sizes,
|
||||
pred_masks_high_res=high_res_masks_for_mem_enc,
|
||||
is_mask_from_pts=(point_inputs is not None),
|
||||
)
|
||||
current_out["maskmem_features"] = maskmem_features
|
||||
current_out["maskmem_pos_enc"] = maskmem_pos_enc
|
||||
else:
|
||||
current_out["maskmem_features"] = None
|
||||
current_out["maskmem_pos_enc"] = None
|
||||
self._encode_memory_in_output(
|
||||
current_vision_feats,
|
||||
feat_sizes,
|
||||
point_inputs,
|
||||
run_mem_encoder,
|
||||
high_res_masks,
|
||||
object_score_logits,
|
||||
current_out,
|
||||
)
|
||||
|
||||
return current_out
|
||||
|
||||
|
||||
@@ -6,11 +6,15 @@
|
||||
|
||||
|
||||
import copy
|
||||
from typing import Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..utils.misc import mask_to_box
|
||||
|
||||
|
||||
def select_closest_cond_frames(frame_idx, cond_frame_outputs, max_cond_frame_num):
|
||||
"""
|
||||
@@ -147,3 +151,173 @@ class LayerNorm2d(nn.Module):
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
||||
return x
|
||||
|
||||
|
||||
def sample_box_points(
|
||||
masks: torch.Tensor,
|
||||
noise: float = 0.1, # SAM default
|
||||
noise_bound: int = 20, # SAM default
|
||||
top_left_label: int = 2,
|
||||
bottom_right_label: int = 3,
|
||||
) -> Tuple[np.array, np.array]:
|
||||
"""
|
||||
Sample a noised version of the top left and bottom right corners of a given `bbox`
|
||||
|
||||
Inputs:
|
||||
- masks: [B, 1, H,W] boxes, dtype=torch.Tensor
|
||||
- noise: noise as a fraction of box width and height, dtype=float
|
||||
- noise_bound: maximum amount of noise (in pure pixesl), dtype=int
|
||||
|
||||
Returns:
|
||||
- box_coords: [B, num_pt, 2], contains (x, y) coordinates of top left and bottom right box corners, dtype=torch.float
|
||||
- box_labels: [B, num_pt], label 2 is reserverd for top left and 3 for bottom right corners, dtype=torch.int32
|
||||
"""
|
||||
device = masks.device
|
||||
box_coords = mask_to_box(masks)
|
||||
B, _, H, W = masks.shape
|
||||
box_labels = torch.tensor(
|
||||
[top_left_label, bottom_right_label], dtype=torch.int, device=device
|
||||
).repeat(B)
|
||||
if noise > 0.0:
|
||||
if not isinstance(noise_bound, torch.Tensor):
|
||||
noise_bound = torch.tensor(noise_bound, device=device)
|
||||
bbox_w = box_coords[..., 2] - box_coords[..., 0]
|
||||
bbox_h = box_coords[..., 3] - box_coords[..., 1]
|
||||
max_dx = torch.min(bbox_w * noise, noise_bound)
|
||||
max_dy = torch.min(bbox_h * noise, noise_bound)
|
||||
box_noise = 2 * torch.rand(B, 1, 4, device=device) - 1
|
||||
box_noise = box_noise * torch.stack((max_dx, max_dy, max_dx, max_dy), dim=-1)
|
||||
|
||||
box_coords = box_coords + box_noise
|
||||
img_bounds = (
|
||||
torch.tensor([W, H, W, H], device=device) - 1
|
||||
) # uncentered pixel coords
|
||||
box_coords.clamp_(torch.zeros_like(img_bounds), img_bounds) # In place clamping
|
||||
|
||||
box_coords = box_coords.reshape(-1, 2, 2) # always 2 points
|
||||
box_labels = box_labels.reshape(-1, 2)
|
||||
return box_coords, box_labels
|
||||
|
||||
|
||||
def sample_random_points_from_errors(gt_masks, pred_masks, num_pt=1):
|
||||
"""
|
||||
Sample `num_pt` random points (along with their labels) independently from the error regions.
|
||||
|
||||
Inputs:
|
||||
- gt_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool
|
||||
- pred_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool or None
|
||||
- num_pt: int, number of points to sample independently for each of the B error maps
|
||||
|
||||
Outputs:
|
||||
- points: [B, num_pt, 2], dtype=torch.float, contains (x, y) coordinates of each sampled point
|
||||
- labels: [B, num_pt], dtype=torch.int32, where 1 means positive clicks and 0 means
|
||||
negative clicks
|
||||
"""
|
||||
if pred_masks is None: # if pred_masks is not provided, treat it as empty
|
||||
pred_masks = torch.zeros_like(gt_masks)
|
||||
assert gt_masks.dtype == torch.bool and gt_masks.size(1) == 1
|
||||
assert pred_masks.dtype == torch.bool and pred_masks.shape == gt_masks.shape
|
||||
assert num_pt >= 0
|
||||
|
||||
B, _, H_im, W_im = gt_masks.shape
|
||||
device = gt_masks.device
|
||||
|
||||
# false positive region, a new point sampled in this region should have
|
||||
# negative label to correct the FP error
|
||||
fp_masks = ~gt_masks & pred_masks
|
||||
# false negative region, a new point sampled in this region should have
|
||||
# positive label to correct the FN error
|
||||
fn_masks = gt_masks & ~pred_masks
|
||||
# whether the prediction completely match the ground-truth on each mask
|
||||
all_correct = torch.all((gt_masks == pred_masks).flatten(2), dim=2)
|
||||
all_correct = all_correct[..., None, None]
|
||||
|
||||
# channel 0 is FP map, while channel 1 is FN map
|
||||
pts_noise = torch.rand(B, num_pt, H_im, W_im, 2, device=device)
|
||||
# sample a negative new click from FP region or a positive new click
|
||||
# from FN region, depend on where the maximum falls,
|
||||
# and in case the predictions are all correct (no FP or FN), we just
|
||||
# sample a negative click from the background region
|
||||
pts_noise[..., 0] *= fp_masks | (all_correct & ~gt_masks)
|
||||
pts_noise[..., 1] *= fn_masks
|
||||
pts_idx = pts_noise.flatten(2).argmax(dim=2)
|
||||
labels = (pts_idx % 2).to(torch.int32)
|
||||
pts_idx = pts_idx // 2
|
||||
pts_x = pts_idx % W_im
|
||||
pts_y = pts_idx // W_im
|
||||
points = torch.stack([pts_x, pts_y], dim=2).to(torch.float)
|
||||
return points, labels
|
||||
|
||||
|
||||
def sample_one_point_from_error_center(gt_masks, pred_masks, padding=True):
|
||||
"""
|
||||
Sample 1 random point (along with its label) from the center of each error region,
|
||||
that is, the point with the largest distance to the boundary of each error region.
|
||||
This is the RITM sampling method from https://github.com/saic-vul/ritm_interactive_segmentation/blob/master/isegm/inference/clicker.py
|
||||
|
||||
Inputs:
|
||||
- gt_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool
|
||||
- pred_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool or None
|
||||
- padding: if True, pad with boundary of 1 px for distance transform
|
||||
|
||||
Outputs:
|
||||
- points: [B, 1, 2], dtype=torch.float, contains (x, y) coordinates of each sampled point
|
||||
- labels: [B, 1], dtype=torch.int32, where 1 means positive clicks and 0 means negative clicks
|
||||
"""
|
||||
import cv2
|
||||
|
||||
if pred_masks is None:
|
||||
pred_masks = torch.zeros_like(gt_masks)
|
||||
assert gt_masks.dtype == torch.bool and gt_masks.size(1) == 1
|
||||
assert pred_masks.dtype == torch.bool and pred_masks.shape == gt_masks.shape
|
||||
|
||||
B, _, _, W_im = gt_masks.shape
|
||||
device = gt_masks.device
|
||||
|
||||
# false positive region, a new point sampled in this region should have
|
||||
# negative label to correct the FP error
|
||||
fp_masks = ~gt_masks & pred_masks
|
||||
# false negative region, a new point sampled in this region should have
|
||||
# positive label to correct the FN error
|
||||
fn_masks = gt_masks & ~pred_masks
|
||||
|
||||
fp_masks = fp_masks.cpu().numpy()
|
||||
fn_masks = fn_masks.cpu().numpy()
|
||||
points = torch.zeros(B, 1, 2, dtype=torch.float)
|
||||
labels = torch.ones(B, 1, dtype=torch.int32)
|
||||
for b in range(B):
|
||||
fn_mask = fn_masks[b, 0]
|
||||
fp_mask = fp_masks[b, 0]
|
||||
if padding:
|
||||
fn_mask = np.pad(fn_mask, ((1, 1), (1, 1)), "constant")
|
||||
fp_mask = np.pad(fp_mask, ((1, 1), (1, 1)), "constant")
|
||||
# compute the distance of each point in FN/FP region to its boundary
|
||||
fn_mask_dt = cv2.distanceTransform(fn_mask.astype(np.uint8), cv2.DIST_L2, 0)
|
||||
fp_mask_dt = cv2.distanceTransform(fp_mask.astype(np.uint8), cv2.DIST_L2, 0)
|
||||
if padding:
|
||||
fn_mask_dt = fn_mask_dt[1:-1, 1:-1]
|
||||
fp_mask_dt = fp_mask_dt[1:-1, 1:-1]
|
||||
|
||||
# take the point in FN/FP region with the largest distance to its boundary
|
||||
fn_mask_dt_flat = fn_mask_dt.reshape(-1)
|
||||
fp_mask_dt_flat = fp_mask_dt.reshape(-1)
|
||||
fn_argmax = np.argmax(fn_mask_dt_flat)
|
||||
fp_argmax = np.argmax(fp_mask_dt_flat)
|
||||
is_positive = fn_mask_dt_flat[fn_argmax] > fp_mask_dt_flat[fp_argmax]
|
||||
pt_idx = fn_argmax if is_positive else fp_argmax
|
||||
points[b, 0, 0] = pt_idx % W_im # x
|
||||
points[b, 0, 1] = pt_idx // W_im # y
|
||||
labels[b, 0] = int(is_positive)
|
||||
|
||||
points = points.to(device)
|
||||
labels = labels.to(device)
|
||||
return points, labels
|
||||
|
||||
|
||||
def get_next_point(gt_masks, pred_masks, method):
|
||||
if method == "uniform":
|
||||
return sample_random_points_from_errors(gt_masks, pred_masks)
|
||||
elif method == "center":
|
||||
return sample_one_point_from_error_center(gt_masks, pred_masks)
|
||||
else:
|
||||
raise ValueError(f"unknown sampling method {method}")
|
||||
|
||||
+263
-15
@@ -4,10 +4,11 @@
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import warnings
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from ..sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
||||
@@ -27,6 +28,9 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
clear_non_cond_mem_around_input=False,
|
||||
# whether to also clear non-conditioning memory of the surrounding frames (only effective when `clear_non_cond_mem_around_input` is True).
|
||||
clear_non_cond_mem_for_multi_obj=False,
|
||||
# if `add_all_frames_to_correct_as_cond` is True, we also append to the conditioning frame list any frame that receives a later correction click
|
||||
# if `add_all_frames_to_correct_as_cond` is False, we conditioning frame list to only use those initial conditioning frames
|
||||
add_all_frames_to_correct_as_cond=False,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
@@ -34,6 +38,7 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
self.non_overlap_masks = non_overlap_masks
|
||||
self.clear_non_cond_mem_around_input = clear_non_cond_mem_around_input
|
||||
self.clear_non_cond_mem_for_multi_obj = clear_non_cond_mem_for_multi_obj
|
||||
self.add_all_frames_to_correct_as_cond = add_all_frames_to_correct_as_cond
|
||||
|
||||
@torch.inference_mode()
|
||||
def init_state(
|
||||
@@ -149,34 +154,66 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
return len(inference_state["obj_idx_to_id"])
|
||||
|
||||
@torch.inference_mode()
|
||||
def add_new_points(
|
||||
def add_new_points_or_box(
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
obj_id,
|
||||
points,
|
||||
labels,
|
||||
points=None,
|
||||
labels=None,
|
||||
clear_old_points=True,
|
||||
normalize_coords=True,
|
||||
box=None,
|
||||
):
|
||||
"""Add new points to a frame."""
|
||||
obj_idx = self._obj_id_to_idx(inference_state, obj_id)
|
||||
point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
|
||||
mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
|
||||
|
||||
if not isinstance(points, torch.Tensor):
|
||||
if isinstance(points, list) and all(isinstance(p, np.ndarray) for p in points):
|
||||
points = np.array(points)
|
||||
points = torch.tensor(points, dtype=torch.float32)
|
||||
if (points is not None) != (labels is not None):
|
||||
raise ValueError("points and labels must be provided together")
|
||||
if points is None and box is None:
|
||||
raise ValueError("at least one of points or box must be provided as input")
|
||||
|
||||
if not isinstance(labels, torch.Tensor):
|
||||
if isinstance(labels, list) and all(isinstance(l, np.ndarray) for l in labels):
|
||||
labels = np.array(labels)
|
||||
if points is None:
|
||||
points = torch.zeros(0, 2, dtype=torch.float32)
|
||||
elif not isinstance(points, torch.Tensor):
|
||||
points = torch.tensor(points, dtype=torch.float32)
|
||||
if labels is None:
|
||||
labels = torch.zeros(0, dtype=torch.int32)
|
||||
elif not isinstance(labels, torch.Tensor):
|
||||
labels = torch.tensor(labels, dtype=torch.int32)
|
||||
if points.dim() == 2:
|
||||
points = points.unsqueeze(0) # add batch dimension
|
||||
if labels.dim() == 1:
|
||||
labels = labels.unsqueeze(0) # add batch dimension
|
||||
|
||||
# If `box` is provided, we add it as the first two points with labels 2 and 3
|
||||
# along with the user-provided points (consistent with how SAM 2 is trained).
|
||||
if box is not None:
|
||||
if not clear_old_points:
|
||||
raise ValueError(
|
||||
"cannot add box without clearing old points, since "
|
||||
"box prompt must be provided before any point prompt "
|
||||
"(please use clear_old_points=True instead)"
|
||||
)
|
||||
if inference_state["tracking_has_started"]:
|
||||
warnings.warn(
|
||||
"You are adding a box after tracking starts. SAM 2 may not always be "
|
||||
"able to incorporate a box prompt for *refinement*. If you intend to "
|
||||
"use box prompt as an *initial* input before tracking, please call "
|
||||
"'reset_state' on the inference state to restart from scratch.",
|
||||
category=UserWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if not isinstance(box, torch.Tensor):
|
||||
box = torch.tensor(box, dtype=torch.float32, device=points.device)
|
||||
box_coords = box.reshape(1, 2, 2)
|
||||
box_labels = torch.tensor([2, 3], dtype=torch.int32, device=labels.device)
|
||||
box_labels = box_labels.reshape(1, 2)
|
||||
points = torch.cat([box_coords, points], dim=1)
|
||||
labels = torch.cat([box_labels, labels], dim=1)
|
||||
|
||||
if normalize_coords:
|
||||
video_H = inference_state["video_height"]
|
||||
video_W = inference_state["video_width"]
|
||||
@@ -259,6 +296,10 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
)
|
||||
return frame_idx, obj_ids, video_res_masks
|
||||
|
||||
def add_new_points(self, *args, **kwargs):
|
||||
"""Deprecated method. Please use `add_new_points_or_box` instead."""
|
||||
return self.add_new_points_or_box(*args, **kwargs)
|
||||
|
||||
@torch.inference_mode()
|
||||
def add_new_mask(
|
||||
self,
|
||||
@@ -414,6 +455,14 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
dtype=torch.float32,
|
||||
device=inference_state["device"],
|
||||
),
|
||||
"object_score_logits": torch.full(
|
||||
size=(batch_size, 1),
|
||||
# default to 10.0 for object_score_logits, i.e. assuming the object is
|
||||
# present as sigmoid(10)=1, same as in `predict_masks` of `MaskDecoder`
|
||||
fill_value=10.0,
|
||||
dtype=torch.float32,
|
||||
device=inference_state["device"],
|
||||
),
|
||||
}
|
||||
empty_mask_ptr = None
|
||||
for obj_idx in range(batch_size):
|
||||
@@ -458,6 +507,9 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
)
|
||||
consolidated_pred_masks[obj_idx : obj_idx + 1] = resized_obj_mask
|
||||
consolidated_out["obj_ptr"][obj_idx : obj_idx + 1] = out["obj_ptr"]
|
||||
consolidated_out["object_score_logits"][obj_idx : obj_idx + 1] = out[
|
||||
"object_score_logits"
|
||||
]
|
||||
|
||||
# Optionally, apply non-overlapping constraints on the consolidated scores
|
||||
# and rerun the memory encoder
|
||||
@@ -476,6 +528,7 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
frame_idx=frame_idx,
|
||||
batch_size=batch_size,
|
||||
high_res_masks=high_res_masks,
|
||||
object_score_logits=consolidated_out["object_score_logits"],
|
||||
is_mask_from_pts=True, # these frames are what the user interacted with
|
||||
)
|
||||
consolidated_out["maskmem_features"] = maskmem_features
|
||||
@@ -535,16 +588,16 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
# to `propagate_in_video_preflight`).
|
||||
consolidated_frame_inds = inference_state["consolidated_frame_inds"]
|
||||
for is_cond in [False, True]:
|
||||
# Separately consolidate conditioning and non-conditioning temp outptus
|
||||
# Separately consolidate conditioning and non-conditioning temp outputs
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
# Find all the frames that contain temporary outputs for any objects
|
||||
# (these should be the frames that have just received clicks for mask inputs
|
||||
# via `add_new_points` or `add_new_mask`)
|
||||
# via `add_new_points_or_box` or `add_new_mask`)
|
||||
temp_frame_inds = set()
|
||||
for obj_temp_output_dict in temp_output_dict_per_obj.values():
|
||||
temp_frame_inds.update(obj_temp_output_dict[storage_key].keys())
|
||||
consolidated_frame_inds[storage_key].update(temp_frame_inds)
|
||||
# consolidate the temprary output across all objects on this frame
|
||||
# consolidate the temporary output across all objects on this frame
|
||||
for frame_idx in temp_frame_inds:
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state, frame_idx, is_cond=is_cond, run_mem_encoder=True
|
||||
@@ -695,6 +748,7 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
"maskmem_pos_enc": None,
|
||||
"pred_masks": current_out["pred_masks"][obj_slice],
|
||||
"obj_ptr": current_out["obj_ptr"][obj_slice],
|
||||
"object_score_logits": current_out["object_score_logits"][obj_slice],
|
||||
}
|
||||
if maskmem_features is not None:
|
||||
obj_out["maskmem_features"] = maskmem_features[obj_slice]
|
||||
@@ -702,6 +756,77 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
obj_out["maskmem_pos_enc"] = [x[obj_slice] for x in maskmem_pos_enc]
|
||||
obj_output_dict[storage_key][frame_idx] = obj_out
|
||||
|
||||
@torch.inference_mode()
|
||||
def clear_all_prompts_in_frame(
|
||||
self, inference_state, frame_idx, obj_id, need_output=True
|
||||
):
|
||||
"""Remove all input points or mask in a specific frame for a given object."""
|
||||
obj_idx = self._obj_id_to_idx(inference_state, obj_id)
|
||||
|
||||
# Clear the conditioning information on the given frame
|
||||
inference_state["point_inputs_per_obj"][obj_idx].pop(frame_idx, None)
|
||||
inference_state["mask_inputs_per_obj"][obj_idx].pop(frame_idx, None)
|
||||
|
||||
temp_output_dict_per_obj = inference_state["temp_output_dict_per_obj"]
|
||||
temp_output_dict_per_obj[obj_idx]["cond_frame_outputs"].pop(frame_idx, None)
|
||||
temp_output_dict_per_obj[obj_idx]["non_cond_frame_outputs"].pop(frame_idx, None)
|
||||
|
||||
# Check and see if there are still any inputs left on this frame
|
||||
batch_size = self._get_obj_num(inference_state)
|
||||
frame_has_input = False
|
||||
for obj_idx2 in range(batch_size):
|
||||
if frame_idx in inference_state["point_inputs_per_obj"][obj_idx2]:
|
||||
frame_has_input = True
|
||||
break
|
||||
if frame_idx in inference_state["mask_inputs_per_obj"][obj_idx2]:
|
||||
frame_has_input = True
|
||||
break
|
||||
|
||||
# If this frame has no remaining inputs for any objects, we further clear its
|
||||
# conditioning frame status
|
||||
if not frame_has_input:
|
||||
output_dict = inference_state["output_dict"]
|
||||
consolidated_frame_inds = inference_state["consolidated_frame_inds"]
|
||||
consolidated_frame_inds["cond_frame_outputs"].discard(frame_idx)
|
||||
consolidated_frame_inds["non_cond_frame_outputs"].discard(frame_idx)
|
||||
# Remove the frame's conditioning output (possibly downgrading it to non-conditioning)
|
||||
out = output_dict["cond_frame_outputs"].pop(frame_idx, None)
|
||||
if out is not None:
|
||||
# The frame is not a conditioning frame anymore since it's not receiving inputs,
|
||||
# so we "downgrade" its output (if exists) to a non-conditioning frame output.
|
||||
output_dict["non_cond_frame_outputs"][frame_idx] = out
|
||||
inference_state["frames_already_tracked"].pop(frame_idx, None)
|
||||
# Similarly, do it for the sliced output on each object.
|
||||
for obj_idx2 in range(batch_size):
|
||||
obj_output_dict = inference_state["output_dict_per_obj"][obj_idx2]
|
||||
obj_out = obj_output_dict["cond_frame_outputs"].pop(frame_idx, None)
|
||||
if obj_out is not None:
|
||||
obj_output_dict["non_cond_frame_outputs"][frame_idx] = obj_out
|
||||
|
||||
# If all the conditioning frames have been removed, we also clear the tracking outputs
|
||||
if len(output_dict["cond_frame_outputs"]) == 0:
|
||||
self._reset_tracking_results(inference_state)
|
||||
|
||||
if not need_output:
|
||||
return
|
||||
# Finally, output updated masks per object (after removing the inputs above)
|
||||
obj_ids = inference_state["obj_ids"]
|
||||
is_cond = any(
|
||||
frame_idx in obj_temp_output_dict["cond_frame_outputs"]
|
||||
for obj_temp_output_dict in temp_output_dict_per_obj.values()
|
||||
)
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond=is_cond,
|
||||
run_mem_encoder=False,
|
||||
consolidate_at_video_res=True,
|
||||
)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, consolidated_out["pred_masks_video_res"]
|
||||
)
|
||||
return frame_idx, obj_ids, video_res_masks
|
||||
|
||||
@torch.inference_mode()
|
||||
def reset_state(self, inference_state):
|
||||
"""Remove all input points or mask in all frames throughout the video."""
|
||||
@@ -823,17 +948,25 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
maskmem_pos_enc = self._get_maskmem_pos_enc(inference_state, current_out)
|
||||
# object pointer is a small tensor, so we always keep it on GPU memory for fast access
|
||||
obj_ptr = current_out["obj_ptr"]
|
||||
object_score_logits = current_out["object_score_logits"]
|
||||
# make a compact version of this frame's output to reduce the state size
|
||||
compact_current_out = {
|
||||
"maskmem_features": maskmem_features,
|
||||
"maskmem_pos_enc": maskmem_pos_enc,
|
||||
"pred_masks": pred_masks,
|
||||
"obj_ptr": obj_ptr,
|
||||
"object_score_logits": object_score_logits,
|
||||
}
|
||||
return compact_current_out, pred_masks_gpu
|
||||
|
||||
def _run_memory_encoder(
|
||||
self, inference_state, frame_idx, batch_size, high_res_masks, is_mask_from_pts
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
batch_size,
|
||||
high_res_masks,
|
||||
object_score_logits,
|
||||
is_mask_from_pts,
|
||||
):
|
||||
"""
|
||||
Run the memory encoder on `high_res_masks`. This is usually after applying
|
||||
@@ -848,6 +981,7 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
current_vision_feats=current_vision_feats,
|
||||
feat_sizes=feat_sizes,
|
||||
pred_masks_high_res=high_res_masks,
|
||||
object_score_logits=object_score_logits,
|
||||
is_mask_from_pts=is_mask_from_pts,
|
||||
)
|
||||
|
||||
@@ -886,6 +1020,120 @@ class SAM2VideoPredictor(SAM2Base):
|
||||
expanded_maskmem_pos_enc = None
|
||||
return expanded_maskmem_pos_enc
|
||||
|
||||
@torch.inference_mode()
|
||||
def remove_object(self, inference_state, obj_id, strict=False, need_output=True):
|
||||
"""
|
||||
Remove an object id from the tracking state. If strict is True, we check whether
|
||||
the object id actually exists and raise an error if it doesn't exist.
|
||||
"""
|
||||
old_obj_idx_to_rm = inference_state["obj_id_to_idx"].get(obj_id, None)
|
||||
updated_frames = []
|
||||
# Check whether this object_id to remove actually exists and possibly raise an error.
|
||||
if old_obj_idx_to_rm is None:
|
||||
if not strict:
|
||||
return inference_state["obj_ids"], updated_frames
|
||||
raise RuntimeError(
|
||||
f"Cannot remove object id {obj_id} as it doesn't exist. "
|
||||
f"All existing object ids: {inference_state['obj_ids']}."
|
||||
)
|
||||
|
||||
# If this is the only remaining object id, we simply reset the state.
|
||||
if len(inference_state["obj_id_to_idx"]) == 1:
|
||||
self.reset_state(inference_state)
|
||||
return inference_state["obj_ids"], updated_frames
|
||||
|
||||
# There are still remaining objects after removing this object id. In this case,
|
||||
# we need to delete the object storage from inference state tensors.
|
||||
# Step 0: clear the input on those frames where this object id has point or mask input
|
||||
# (note that this step is required as it might downgrade conditioning frames to
|
||||
# non-conditioning ones)
|
||||
obj_input_frames_inds = set()
|
||||
obj_input_frames_inds.update(
|
||||
inference_state["point_inputs_per_obj"][old_obj_idx_to_rm]
|
||||
)
|
||||
obj_input_frames_inds.update(
|
||||
inference_state["mask_inputs_per_obj"][old_obj_idx_to_rm]
|
||||
)
|
||||
for frame_idx in obj_input_frames_inds:
|
||||
self.clear_all_prompts_in_frame(
|
||||
inference_state, frame_idx, obj_id, need_output=False
|
||||
)
|
||||
|
||||
# Step 1: Update the object id mapping (note that it must be done after Step 0,
|
||||
# since Step 0 still requires the old object id mappings in inference_state)
|
||||
old_obj_ids = inference_state["obj_ids"]
|
||||
old_obj_inds = list(range(len(old_obj_ids)))
|
||||
remain_old_obj_inds = old_obj_inds.copy()
|
||||
remain_old_obj_inds.remove(old_obj_idx_to_rm)
|
||||
new_obj_ids = [old_obj_ids[old_idx] for old_idx in remain_old_obj_inds]
|
||||
new_obj_inds = list(range(len(new_obj_ids)))
|
||||
# build new mappings
|
||||
old_idx_to_new_idx = dict(zip(remain_old_obj_inds, new_obj_inds))
|
||||
inference_state["obj_id_to_idx"] = dict(zip(new_obj_ids, new_obj_inds))
|
||||
inference_state["obj_idx_to_id"] = dict(zip(new_obj_inds, new_obj_ids))
|
||||
inference_state["obj_ids"] = new_obj_ids
|
||||
|
||||
# Step 2: For per-object tensor storage, we shift their obj_idx in the dict keys.
|
||||
# (note that "consolidated_frame_inds" doesn't need to be updated in this step as
|
||||
# it's already handled in Step 0)
|
||||
def _map_keys(container):
|
||||
new_kvs = []
|
||||
for k in old_obj_inds:
|
||||
v = container.pop(k)
|
||||
if k in old_idx_to_new_idx:
|
||||
new_kvs.append((old_idx_to_new_idx[k], v))
|
||||
container.update(new_kvs)
|
||||
|
||||
_map_keys(inference_state["point_inputs_per_obj"])
|
||||
_map_keys(inference_state["mask_inputs_per_obj"])
|
||||
_map_keys(inference_state["output_dict_per_obj"])
|
||||
_map_keys(inference_state["temp_output_dict_per_obj"])
|
||||
|
||||
# Step 3: For packed tensor storage, we index the remaining ids and rebuild the per-object slices.
|
||||
def _slice_state(output_dict, storage_key):
|
||||
for frame_idx, out in output_dict[storage_key].items():
|
||||
out["maskmem_features"] = out["maskmem_features"][remain_old_obj_inds]
|
||||
out["maskmem_pos_enc"] = [
|
||||
x[remain_old_obj_inds] for x in out["maskmem_pos_enc"]
|
||||
]
|
||||
# "maskmem_pos_enc" is the same across frames, so we only need to store one copy of it
|
||||
out["maskmem_pos_enc"] = self._get_maskmem_pos_enc(inference_state, out)
|
||||
out["pred_masks"] = out["pred_masks"][remain_old_obj_inds]
|
||||
out["obj_ptr"] = out["obj_ptr"][remain_old_obj_inds]
|
||||
out["object_score_logits"] = out["object_score_logits"][
|
||||
remain_old_obj_inds
|
||||
]
|
||||
# also update the per-object slices
|
||||
self._add_output_per_object(
|
||||
inference_state, frame_idx, out, storage_key
|
||||
)
|
||||
|
||||
_slice_state(inference_state["output_dict"], "cond_frame_outputs")
|
||||
_slice_state(inference_state["output_dict"], "non_cond_frame_outputs")
|
||||
|
||||
# Step 4: Further collect the outputs on those frames in `obj_input_frames_inds`, which
|
||||
# could show an updated mask for objects previously occluded by the object being removed
|
||||
if need_output:
|
||||
temp_output_dict_per_obj = inference_state["temp_output_dict_per_obj"]
|
||||
for frame_idx in obj_input_frames_inds:
|
||||
is_cond = any(
|
||||
frame_idx in obj_temp_output_dict["cond_frame_outputs"]
|
||||
for obj_temp_output_dict in temp_output_dict_per_obj.values()
|
||||
)
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond=is_cond,
|
||||
run_mem_encoder=False,
|
||||
consolidate_at_video_res=True,
|
||||
)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, consolidated_out["pred_masks_video_res"]
|
||||
)
|
||||
updated_frames.append((frame_idx, video_res_masks))
|
||||
|
||||
return inference_state["obj_ids"], updated_frames
|
||||
|
||||
def _clear_non_cond_mem_around_input(self, inference_state, frame_idx):
|
||||
"""
|
||||
Remove the non-conditioning memory around the input frame. When users provide
|
||||
|
||||
+124
-13
@@ -68,7 +68,7 @@ def mask_to_box(masks: torch.Tensor):
|
||||
compute bounding box given an input mask
|
||||
|
||||
Inputs:
|
||||
- masks: [B, 1, H, W] boxes, dtype=torch.Tensor
|
||||
- masks: [B, 1, H, W] masks, dtype=torch.Tensor
|
||||
|
||||
Returns:
|
||||
- box_coords: [B, 1, 4], contains (x, y) coordinates of top left and bottom right box corners, dtype=torch.Tensor
|
||||
@@ -106,19 +106,28 @@ class AsyncVideoFrameLoader:
|
||||
A list of video frames to be load asynchronously without blocking session start.
|
||||
"""
|
||||
|
||||
def __init__(self, img_paths, image_size, offload_video_to_cpu, img_mean, img_std):
|
||||
def __init__(
|
||||
self,
|
||||
img_paths,
|
||||
image_size,
|
||||
offload_video_to_cpu,
|
||||
img_mean,
|
||||
img_std,
|
||||
compute_device,
|
||||
):
|
||||
self.img_paths = img_paths
|
||||
self.image_size = image_size
|
||||
self.offload_video_to_cpu = offload_video_to_cpu
|
||||
self.img_mean = img_mean
|
||||
self.img_std = img_std
|
||||
# items in `self._images` will be loaded asynchronously
|
||||
# items in `self.images` will be loaded asynchronously
|
||||
self.images = [None] * len(img_paths)
|
||||
# catch and raise any exceptions in the async loading thread
|
||||
self.exception = None
|
||||
# video_height and video_width be filled when loading the first image
|
||||
self.video_height = None
|
||||
self.video_width = None
|
||||
self.compute_device = compute_device
|
||||
|
||||
# load the first frame to fill video_height and video_width and also
|
||||
# to cache it (since it's most likely where the user will click)
|
||||
@@ -152,7 +161,7 @@ class AsyncVideoFrameLoader:
|
||||
img -= self.img_mean
|
||||
img /= self.img_std
|
||||
if not self.offload_video_to_cpu:
|
||||
img = img.cuda(non_blocking=True)
|
||||
img = img.to(self.compute_device, non_blocking=True)
|
||||
self.images[index] = img
|
||||
return img
|
||||
|
||||
@@ -167,6 +176,48 @@ def load_video_frames(
|
||||
img_mean=(0.485, 0.456, 0.406),
|
||||
img_std=(0.229, 0.224, 0.225),
|
||||
async_loading_frames=False,
|
||||
compute_device=torch.device("cuda"),
|
||||
):
|
||||
"""
|
||||
Load the video frames from video_path. The frames are resized to image_size as in
|
||||
the model and are loaded to GPU if offload_video_to_cpu=False. This is used by the demo.
|
||||
"""
|
||||
is_bytes = isinstance(video_path, bytes)
|
||||
is_str = isinstance(video_path, str)
|
||||
is_mp4_path = is_str and os.path.splitext(video_path)[-1] in [".mp4", ".MP4"]
|
||||
if is_bytes or is_mp4_path:
|
||||
return load_video_frames_from_video_file(
|
||||
video_path=video_path,
|
||||
image_size=image_size,
|
||||
offload_video_to_cpu=offload_video_to_cpu,
|
||||
img_mean=img_mean,
|
||||
img_std=img_std,
|
||||
compute_device=compute_device,
|
||||
)
|
||||
elif is_str and os.path.isdir(video_path):
|
||||
return load_video_frames_from_jpg_images(
|
||||
video_path=video_path,
|
||||
image_size=image_size,
|
||||
offload_video_to_cpu=offload_video_to_cpu,
|
||||
img_mean=img_mean,
|
||||
img_std=img_std,
|
||||
async_loading_frames=async_loading_frames,
|
||||
compute_device=compute_device,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"Only MP4 video and JPEG folder are supported at this moment"
|
||||
)
|
||||
|
||||
|
||||
def load_video_frames_from_jpg_images(
|
||||
video_path,
|
||||
image_size,
|
||||
offload_video_to_cpu,
|
||||
img_mean=(0.485, 0.456, 0.406),
|
||||
img_std=(0.229, 0.224, 0.225),
|
||||
async_loading_frames=False,
|
||||
compute_device=torch.device("cuda"),
|
||||
):
|
||||
"""
|
||||
Load the video frames from a directory of JPEG files ("<frame_index>.jpg" format).
|
||||
@@ -179,7 +230,15 @@ def load_video_frames(
|
||||
if isinstance(video_path, str) and os.path.isdir(video_path):
|
||||
jpg_folder = video_path
|
||||
else:
|
||||
raise NotImplementedError("Only JPEG frames are supported at this moment")
|
||||
raise NotImplementedError(
|
||||
"Only JPEG frames are supported at this moment. For video files, you may use "
|
||||
"ffmpeg (https://ffmpeg.org/) to extract frames into a folder of JPEG files, such as \n"
|
||||
"```\n"
|
||||
"ffmpeg -i <your_video>.mp4 -q:v 2 -start_number 0 <output_dir>/'%05d.jpg'\n"
|
||||
"```\n"
|
||||
"where `-q:v` generates high-quality JPEG frames and `-start_number 0` asks "
|
||||
"ffmpeg to start the JPEG file from 00000.jpg."
|
||||
)
|
||||
|
||||
frame_names = [
|
||||
p
|
||||
@@ -196,7 +255,12 @@ def load_video_frames(
|
||||
|
||||
if async_loading_frames:
|
||||
lazy_images = AsyncVideoFrameLoader(
|
||||
img_paths, image_size, offload_video_to_cpu, img_mean, img_std
|
||||
img_paths,
|
||||
image_size,
|
||||
offload_video_to_cpu,
|
||||
img_mean,
|
||||
img_std,
|
||||
compute_device,
|
||||
)
|
||||
return lazy_images, lazy_images.video_height, lazy_images.video_width
|
||||
|
||||
@@ -204,9 +268,41 @@ def load_video_frames(
|
||||
for n, img_path in enumerate(tqdm(img_paths, desc="frame loading (JPEG)")):
|
||||
images[n], video_height, video_width = _load_img_as_tensor(img_path, image_size)
|
||||
if not offload_video_to_cpu:
|
||||
images = images.cuda()
|
||||
img_mean = img_mean.cuda()
|
||||
img_std = img_std.cuda()
|
||||
images = images.to(compute_device)
|
||||
img_mean = img_mean.to(compute_device)
|
||||
img_std = img_std.to(compute_device)
|
||||
# normalize by mean and std
|
||||
images -= img_mean
|
||||
images /= img_std
|
||||
return images, video_height, video_width
|
||||
|
||||
|
||||
def load_video_frames_from_video_file(
|
||||
video_path,
|
||||
image_size,
|
||||
offload_video_to_cpu,
|
||||
img_mean=(0.485, 0.456, 0.406),
|
||||
img_std=(0.229, 0.224, 0.225),
|
||||
compute_device=torch.device("cuda"),
|
||||
):
|
||||
"""Load the video frames from a video file."""
|
||||
import decord
|
||||
|
||||
img_mean = torch.tensor(img_mean, dtype=torch.float32)[:, None, None]
|
||||
img_std = torch.tensor(img_std, dtype=torch.float32)[:, None, None]
|
||||
# Get the original video height and width
|
||||
decord.bridge.set_bridge("torch")
|
||||
video_height, video_width, _ = decord.VideoReader(video_path).next().shape
|
||||
# Iterate over all frames in the video
|
||||
images = []
|
||||
for frame in decord.VideoReader(video_path, width=image_size, height=image_size):
|
||||
images.append(frame.permute(2, 0, 1))
|
||||
|
||||
images = torch.stack(images, dim=0).float() / 255.0
|
||||
if not offload_video_to_cpu:
|
||||
images = images.to(compute_device)
|
||||
img_mean = img_mean.to(compute_device)
|
||||
img_std = img_std.to(compute_device)
|
||||
# normalize by mean and std
|
||||
images -= img_mean
|
||||
images /= img_std
|
||||
@@ -220,10 +316,25 @@ def fill_holes_in_mask_scores(mask, max_area):
|
||||
# Holes are those connected components in background with area <= self.max_area
|
||||
# (background regions are those with mask scores <= 0)
|
||||
assert max_area > 0, "max_area must be positive"
|
||||
labels, areas = get_connected_components(mask <= 0)
|
||||
is_hole = (labels > 0) & (areas <= max_area)
|
||||
# We fill holes with a small positive mask score (0.1) to change them to foreground.
|
||||
mask = torch.where(is_hole, 0.1, mask)
|
||||
|
||||
input_mask = mask
|
||||
try:
|
||||
labels, areas = get_connected_components(mask <= 0)
|
||||
is_hole = (labels > 0) & (areas <= max_area)
|
||||
# We fill holes with a small positive mask score (0.1) to change them to foreground.
|
||||
mask = torch.where(is_hole, 0.1, mask)
|
||||
except Exception as e:
|
||||
# Skip the post-processing step on removing small holes if the CUDA kernel fails
|
||||
warnings.warn(
|
||||
f"{e}\n\nSkipping the post-processing step due to the error above. You can "
|
||||
"still use SAM 2 and it's OK to ignore the error above, although some post-processing "
|
||||
"functionality may be limited (which doesn't affect the results in most cases; see "
|
||||
"https://github.com/facebookresearch/sam2/blob/main/INSTALL.md).",
|
||||
category=UserWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
mask = input_mask
|
||||
|
||||
return mask
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 112
|
||||
num_heads: 2
|
||||
neck:
|
||||
_target_: sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [896, 448, 224, 112]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: true
|
||||
proj_tpos_enc_in_obj_ptrs: true
|
||||
use_signed_tpos_enc_to_obj_ptrs: true
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,120 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 144
|
||||
num_heads: 2
|
||||
stages: [2, 6, 36, 4]
|
||||
global_att_blocks: [23, 33, 43]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
window_spec: [8, 4, 16, 8]
|
||||
neck:
|
||||
_target_: sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [1152, 576, 288, 144]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: true
|
||||
proj_tpos_enc_in_obj_ptrs: true
|
||||
use_signed_tpos_enc_to_obj_ptrs: true
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,119 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 96
|
||||
num_heads: 1
|
||||
stages: [1, 2, 11, 2]
|
||||
global_att_blocks: [7, 10, 13]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
neck:
|
||||
_target_: sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [768, 384, 192, 96]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: true
|
||||
proj_tpos_enc_in_obj_ptrs: true
|
||||
use_signed_tpos_enc_to_obj_ptrs: true
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,121 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 96
|
||||
num_heads: 1
|
||||
stages: [1, 2, 7, 2]
|
||||
global_att_blocks: [5, 7, 9]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
neck:
|
||||
_target_: sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [768, 384, 192, 96]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
# SAM decoder
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: true
|
||||
proj_tpos_enc_in_obj_ptrs: true
|
||||
use_signed_tpos_enc_to_obj_ptrs: true
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
# HieraT does not currently support compilation, should always be set to False
|
||||
compile_image_encoder: False
|
||||
@@ -92,6 +92,7 @@ model:
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: false
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
@@ -101,6 +102,8 @@ model:
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
proj_tpos_enc_in_obj_ptrs: false
|
||||
use_signed_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
|
||||
@@ -93,6 +93,7 @@ model:
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: false
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
@@ -102,6 +103,8 @@ model:
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
proj_tpos_enc_in_obj_ptrs: false
|
||||
use_signed_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
|
||||
@@ -92,6 +92,7 @@ model:
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: false
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
@@ -101,6 +102,8 @@ model:
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
proj_tpos_enc_in_obj_ptrs: false
|
||||
use_signed_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
|
||||
@@ -93,6 +93,7 @@ model:
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
no_obj_embed_spatial: false
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
@@ -102,6 +103,8 @@ model:
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
proj_tpos_enc_in_obj_ptrs: false
|
||||
use_signed_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
|
||||
Reference in New Issue
Block a user