Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b4ae8f44bb | ||
|
|
0c6f8b321c | ||
|
|
5465385d48 | ||
|
|
0564c80e07 | ||
|
|
2a59fe59b3 | ||
|
|
8b00984581 | ||
|
|
8db624cfe1 | ||
|
|
6258558746 | ||
|
|
a4c053dae8 | ||
|
|
6b2c03f8bf | ||
|
|
0bbc1e0efa | ||
|
|
60c60c5c2b | ||
|
|
4f587443fb | ||
|
|
37aa0d4c89 | ||
|
|
843ca3e733 | ||
|
|
f4e56bd733 | ||
|
|
de1fb0ab2a | ||
|
|
b29709c50d | ||
|
|
0479c98a11 | ||
|
|
3839d938d6 |
@@ -31,7 +31,7 @@ class DownloadAndLoadSAM2RealtimeModel:
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
"model": ([
|
"model": ([
|
||||||
'sam2_hiera_tiny.pt',
|
'sam2_hiera_tiny.pt', 'sam2_hiera_small.pt',
|
||||||
],),
|
],),
|
||||||
"segmentor": (
|
"segmentor": (
|
||||||
['realtime'],
|
['realtime'],
|
||||||
@@ -70,7 +70,8 @@ class DownloadAndLoadSAM2RealtimeModel:
|
|||||||
|
|
||||||
if not os.path.exists(model_path):
|
if not os.path.exists(model_path):
|
||||||
print(f"Downloading SAM2 model to: {model_path}")
|
print(f"Downloading SAM2 model to: {model_path}")
|
||||||
url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/sam2_hiera_tiny.pt"
|
base_url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/"
|
||||||
|
url = f"{base_url}{model}"
|
||||||
response = requests.get(url, stream=True)
|
response = requests.get(url, stream=True)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
|
|
||||||
@@ -83,8 +84,9 @@ class DownloadAndLoadSAM2RealtimeModel:
|
|||||||
|
|
||||||
config_dir = os.path.join(script_directory, "sam2_configs")
|
config_dir = os.path.join(script_directory, "sam2_configs")
|
||||||
|
|
||||||
|
model_cfg = model.replace(".pt", ".yaml")
|
||||||
|
|
||||||
# Code ripped out of sam2.build_sam.build_sam2_camera_predictor to appease Hydra
|
# Code ripped out of sam2.build_sam.build_sam2_camera_predictor to appease Hydra
|
||||||
model_cfg = "sam2_hiera_t.yaml" #TODO: remove hardcoded config and path
|
|
||||||
with initialize_config_dir(config_dir=config_dir, version_base=None):
|
with initialize_config_dir(config_dir=config_dir, version_base=None):
|
||||||
cfg = compose(config_name=model_cfg)
|
cfg = compose(config_name=model_cfg)
|
||||||
|
|
||||||
@@ -140,12 +142,12 @@ class Sam2RealtimeSegmentation:
|
|||||||
"required": {
|
"required": {
|
||||||
"images": ("IMAGE",),
|
"images": ("IMAGE",),
|
||||||
"sam2_model": ("SAM2MODEL",),
|
"sam2_model": ("SAM2MODEL",),
|
||||||
"reset_tracking": ("BOOLEAN", {"default": False}),
|
|
||||||
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"coordinates_positive": ("STRING", ),
|
"coordinates_positive": ("STRING", ),
|
||||||
"coordinates_negative": ("STRING", ),
|
"coordinates_negative": ("STRING", ),
|
||||||
|
"reset_tracking": ("BOOLEAN", {"default": False}),
|
||||||
# "bboxes": ("BBOX", ),
|
# "bboxes": ("BBOX", ),
|
||||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||||
# "mask": ("MASK", ),
|
# "mask": ("MASK", ),
|
||||||
@@ -192,9 +194,9 @@ class Sam2RealtimeSegmentation:
|
|||||||
images,
|
images,
|
||||||
sam2_model,
|
sam2_model,
|
||||||
# keep_model_loaded,
|
# keep_model_loaded,
|
||||||
reset_tracking,
|
|
||||||
coordinates_positive=None,
|
coordinates_positive=None,
|
||||||
coordinates_negative=None,
|
coordinates_negative=None,
|
||||||
|
reset_tracking=False,
|
||||||
#point_labels=None,
|
#point_labels=None,
|
||||||
# bboxes=None,
|
# bboxes=None,
|
||||||
# individual_objects=False,
|
# individual_objects=False,
|
||||||
@@ -250,6 +252,7 @@ class Sam2RealtimeSegmentation:
|
|||||||
|
|
||||||
# Create colored overlay for processed frames
|
# Create colored overlay for processed frames
|
||||||
mask_colored = torch.stack([mask] * 3, dim=2)
|
mask_colored = torch.stack([mask] * 3, dim=2)
|
||||||
|
|
||||||
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
|
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
|
||||||
|
|
||||||
processed_frames.append(overlayed_frame)
|
processed_frames.append(overlayed_frame)
|
||||||
|
|||||||
@@ -1,6 +1,29 @@
|
|||||||
|
[project]
|
||||||
|
name = "sam2_realtime_forktest"
|
||||||
|
description = "This extension provides object segmentation capabilities for ComfyUI workflows"
|
||||||
|
version = "0.0.4"
|
||||||
|
license = { file = "LICENSE" }
|
||||||
|
dependencies = [
|
||||||
|
"sam2_realtime @ git+https://github.com/pschroedl/ComfyUI-SAM2-Realtime.git@main",
|
||||||
|
"pyyaml>6.0.2",
|
||||||
|
"numpy>=1.24.4",
|
||||||
|
"tqdm>=4.66.1",
|
||||||
|
"hydra-core>=1.3.2",
|
||||||
|
"iopath>=0.1.10",
|
||||||
|
"pillow>=9.4.0"
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
Repository = "https://github.com/eliteprox/ComfyUI-SAM2-Realtime"
|
||||||
|
|
||||||
[build-system]
|
[build-system]
|
||||||
requires = [
|
requires = [
|
||||||
"setuptools>=61.0",
|
"setuptools>=61.0",
|
||||||
"torch>=2.3.1",
|
"torch>=2.3.1",
|
||||||
]
|
]
|
||||||
build-backend = "setuptools.build_meta"
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[tool.comfy]
|
||||||
|
PublisherId = "eliteprox"
|
||||||
|
DisplayName = "ComfyUI-SAM2-Realtime-TEST"
|
||||||
|
Icon = ""
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
# @package _global_
|
||||||
|
|
||||||
|
# Model
|
||||||
|
model:
|
||||||
|
_target_: sam2_realtime.modeling.sam2_base.SAM2Base
|
||||||
|
image_encoder:
|
||||||
|
_target_: sam2_realtime.modeling.backbones.image_encoder.ImageEncoder
|
||||||
|
scalp: 1
|
||||||
|
trunk:
|
||||||
|
_target_: sam2_realtime.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_realtime.modeling.backbones.image_encoder.FpnNeck
|
||||||
|
position_encoding:
|
||||||
|
_target_: sam2_realtime.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_realtime.modeling.memory_attention.MemoryAttention
|
||||||
|
d_model: 256
|
||||||
|
pos_enc_at_input: true
|
||||||
|
layer:
|
||||||
|
_target_: sam2_realtime.modeling.memory_attention.MemoryAttentionLayer
|
||||||
|
activation: relu
|
||||||
|
dim_feedforward: 2048
|
||||||
|
dropout: 0.1
|
||||||
|
pos_enc_at_attn: false
|
||||||
|
self_attention:
|
||||||
|
_target_: sam2_realtime.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_realtime.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_realtime.modeling.memory_encoder.MemoryEncoder
|
||||||
|
out_dim: 64
|
||||||
|
position_encoding:
|
||||||
|
_target_: sam2_realtime.modeling.position_encoding.PositionEmbeddingSine
|
||||||
|
num_pos_feats: 64
|
||||||
|
normalize: true
|
||||||
|
scale: null
|
||||||
|
temperature: 10000
|
||||||
|
mask_downsampler:
|
||||||
|
_target_: sam2_realtime.modeling.memory_encoder.MaskDownSampler
|
||||||
|
kernel_size: 3
|
||||||
|
stride: 2
|
||||||
|
padding: 1
|
||||||
|
fuser:
|
||||||
|
_target_: sam2_realtime.modeling.memory_encoder.Fuser
|
||||||
|
layer:
|
||||||
|
_target_: sam2_realtime.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: 512
|
||||||
|
# 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
|
||||||
|
# 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: false
|
||||||
|
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
|
||||||
@@ -85,7 +85,7 @@ model:
|
|||||||
num_layers: 2
|
num_layers: 2
|
||||||
|
|
||||||
num_maskmem: 7
|
num_maskmem: 7
|
||||||
image_size: 1024
|
image_size: 512
|
||||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||||
# SAM decoder
|
# SAM decoder
|
||||||
sigmoid_scale_for_mem_enc: 20.0
|
sigmoid_scale_for_mem_enc: 20.0
|
||||||
@@ -85,7 +85,7 @@ def build_sam2_camera_predictor(
|
|||||||
apply_postprocessing=True,
|
apply_postprocessing=True,
|
||||||
):
|
):
|
||||||
if GlobalHydra.instance().is_initialized():
|
if GlobalHydra.instance().is_initialized():
|
||||||
GlobalHydra.instance().clear()
|
GlobalHydra.instance().clear()
|
||||||
|
|
||||||
# Initialize Hydra to load the configuration
|
# Initialize Hydra to load the configuration
|
||||||
config_path = "sam2_configs"
|
config_path = "sam2_configs"
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ from tqdm import tqdm
|
|||||||
from sam2_realtime.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
from sam2_realtime.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
||||||
from sam2_realtime.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
|
from sam2_realtime.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
|
||||||
|
|
||||||
|
|
||||||
class SAM2TensorPredictor(SAM2Base):
|
class SAM2TensorPredictor(SAM2Base):
|
||||||
"""The predictor class to handle user interactions and manage inference states."""
|
"""The predictor class to handle user interactions and manage inference states."""
|
||||||
|
|
||||||
@@ -55,7 +54,9 @@ class SAM2TensorPredictor(SAM2Base):
|
|||||||
img = img.float()
|
img = img.float()
|
||||||
else:
|
else:
|
||||||
raise ValueError("Input must be a numpy array or a PyTorch tensor")
|
raise ValueError("Input must be a numpy array or a PyTorch tensor")
|
||||||
|
#save original height/width
|
||||||
|
orig_h, orig_w = img.shape[1:]
|
||||||
|
|
||||||
# Resize to the target size (supports tensor resizing)
|
# Resize to the target size (supports tensor resizing)
|
||||||
img = torch.nn.functional.interpolate(
|
img = torch.nn.functional.interpolate(
|
||||||
img.unsqueeze(0), size=(image_size, image_size), mode="bilinear", align_corners=False
|
img.unsqueeze(0), size=(image_size, image_size), mode="bilinear", align_corners=False
|
||||||
@@ -68,28 +69,29 @@ class SAM2TensorPredictor(SAM2Base):
|
|||||||
img /= img_std
|
img /= img_std
|
||||||
|
|
||||||
height, width = img.shape[1:] # CHW format
|
height, width = img.shape[1:] # CHW format
|
||||||
return img, width, height
|
return img, width, height, orig_w, orig_h
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
def load_first_frame(self, img):
|
def load_first_frame(self, img):
|
||||||
if isinstance(img, torch.Tensor):
|
if isinstance(img, torch.Tensor):
|
||||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||||
|
|
||||||
self.condition_state = self._init_state(
|
self.condition_state = self._init_state(
|
||||||
offload_video_to_cpu=False, offload_state_to_cpu=False
|
offload_video_to_cpu=False, offload_state_to_cpu=False
|
||||||
)
|
)
|
||||||
img, width, height = self.prepare_data(img, image_size=self.image_size)
|
img, width, height, orig_w, orig_h = self.prepare_data(img, image_size=self.image_size)
|
||||||
|
self._orig_hw = (orig_w, orig_h)
|
||||||
self.condition_state["images"] = [img]
|
self.condition_state["images"] = [img]
|
||||||
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
||||||
self.condition_state["video_height"] = height
|
self.condition_state["video_height"] = height
|
||||||
self.condition_state["video_width"] = width
|
self.condition_state["video_width"] = width
|
||||||
self._get_image_feature(frame_idx=0, batch_size=1)
|
self._get_image_feature(frame_idx=0, batch_size=1)
|
||||||
|
|
||||||
def add_conditioning_frame(self, img):
|
def add_conditioning_frame(self, img):
|
||||||
if isinstance(img, torch.Tensor):
|
if isinstance(img, torch.Tensor):
|
||||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||||
|
|
||||||
img, width, height = self.prepare_data(img, image_size=self.image_size)
|
img, width, height, _, _ = self.prepare_data(img, image_size=self.image_size)
|
||||||
self.condition_state["images"].append(img)
|
self.condition_state["images"].append(img)
|
||||||
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
||||||
self._get_image_feature(
|
self._get_image_feature(
|
||||||
@@ -235,14 +237,15 @@ class SAM2TensorPredictor(SAM2Base):
|
|||||||
points = torch.cat([box_coords, points], dim=1)
|
points = torch.cat([box_coords, points], dim=1)
|
||||||
labels = torch.cat([box_labels, labels], dim=1)
|
labels = torch.cat([box_labels, labels], dim=1)
|
||||||
if normalize_coords:
|
if normalize_coords:
|
||||||
video_H = self.condition_state["video_height"]
|
#video_H = self.condition_state["video_height"]
|
||||||
video_W = self.condition_state["video_width"]
|
#video_W = self.condition_state["video_width"]
|
||||||
points = points / torch.tensor([video_W, video_H]).to(points.device)
|
orig_w, orig_h = self._orig_hw
|
||||||
|
|
||||||
|
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
|
||||||
# scale the (normalized) coordinates by the model's internal image size
|
# scale the (normalized) coordinates by the model's internal image size
|
||||||
points = points * self.image_size
|
points = points * self.image_size
|
||||||
points = points.to(self.condition_state["device"])
|
points = points.to(self.condition_state["device"])
|
||||||
labels = labels.to(self.condition_state["device"])
|
labels = labels.to(self.condition_state["device"])
|
||||||
|
|
||||||
if not clear_old_points:
|
if not clear_old_points:
|
||||||
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
||||||
else:
|
else:
|
||||||
@@ -342,14 +345,16 @@ class SAM2TensorPredictor(SAM2Base):
|
|||||||
if labels.dim() == 1:
|
if labels.dim() == 1:
|
||||||
labels = labels.unsqueeze(0) # add batch dimension
|
labels = labels.unsqueeze(0) # add batch dimension
|
||||||
if normalize_coords:
|
if normalize_coords:
|
||||||
video_H = self.condition_state["video_height"]
|
#video_H = self.condition_state["video_height"]
|
||||||
video_W = self.condition_state["video_width"]
|
#video_W = self.condition_state["video_width"]
|
||||||
points = points / torch.tensor([video_W, video_H]).to(points.device)
|
orig_w, orig_h = self._orig_hw
|
||||||
|
|
||||||
|
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
|
||||||
|
|
||||||
# scale the (normalized) coordinates by the model's internal image size
|
# scale the (normalized) coordinates by the model's internal image size
|
||||||
points = points * self.image_size
|
points = points * self.image_size
|
||||||
points = points.to(self.condition_state["device"])
|
points = points.to(self.condition_state["device"])
|
||||||
labels = labels.to(self.condition_state["device"])
|
labels = labels.to(self.condition_state["device"])
|
||||||
|
|
||||||
if not clear_old_points:
|
if not clear_old_points:
|
||||||
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
||||||
else:
|
else:
|
||||||
@@ -769,7 +774,7 @@ class SAM2TensorPredictor(SAM2Base):
|
|||||||
if isinstance(img, torch.Tensor):
|
if isinstance(img, torch.Tensor):
|
||||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||||
|
|
||||||
img, _, _ = self.prepare_data(img, image_size=self.image_size)
|
img, _, _ , _, _ = self.prepare_data(img, image_size=self.image_size)
|
||||||
|
|
||||||
output_dict = self.condition_state["output_dict"]
|
output_dict = self.condition_state["output_dict"]
|
||||||
obj_ids = self.condition_state["obj_ids"]
|
obj_ids = self.condition_state["obj_ids"]
|
||||||
|
|||||||
Reference in New Issue
Block a user