Author SHA1 Message Date
PSchroedl 8e6213cfdd Merge pull request #8 from eliteprox/revert-patch1
Revert "update pyproject.toml"
2025-01-20 21:29:03 -08:00
Elite Encoder 6a68c7b11d Revert "update pyproject.toml"
This reverts commit 5465385d48.
2025-01-21 00:16:36 -05:00
John Mull 5465385d48 update pyproject.toml 2025-01-17 16:16:55 +00:00
PSchroedl 0564c80e07 Merge pull request #7 from pschroedl/revert-6-fix-install 2025-01-13 20:41:31 -08:00
John | Elite Encoder 2a59fe59b3 Revert "Fix node install from conda environments" 2025-01-13 23:10:45 -05:00
PSchroedl 8b00984581 Merge pull request #6 from eliteprox/fix-install 2025-01-09 19:41:11 -08:00
John | Elite Encoder 8db624cfe1 Update README.md 2025-01-09 21:49:26 -05:00
John | Elite Encoder 6258558746 Update README.md 2025-01-09 21:48:56 -05:00
John | Elite Encoder a4c053dae8 Update README.md 2025-01-09 21:34:35 -05:00
John | Elite Encoder 6b2c03f8bf Create pyproject.toml 2025-01-08 23:15:00 -05:00
John | Elite Encoder 0bbc1e0efa Create README.md
add install notes
2025-01-08 23:14:36 -05:00
John | Elite Encoder 60c60c5c2b remove self install from requirements.txt
Resolves issue with conda windows environments losing system environment variable context (e.g. CUDA_HOME error)
2025-01-08 22:56:46 -05:00
PSchroedl 4f587443fb Merge pull request #5 from pschroedl/unbreak_old_workflows_and_speedup
Unbreak old workflows and speedup
2024-12-09 23:58:21 -08:00
Peter Schroedl 37aa0d4c89 add small model option and config 2024-12-10 08:42:30 +01:00
Peter Schroedl 843ca3e733 reduce resolution internally to 512 2024-12-10 08:28:06 +01:00
Peter Schroedl f4e56bd733 make reset_tracking optional 2024-12-10 08:26:06 +01:00
PSchroedl de1fb0ab2a Merge pull request #4 from pschroedl/fix_point_coords
fix: update point coords scaling for mask
2024-12-09 21:47:43 -08:00
Brad P b29709c50d remove import not needed 2024-12-09 19:57:55 -06:00
Brad P 0479c98a11 fix point coords scaling for mask 2024-12-09 19:52:14 -06:00
PSchroedl 3839d938d6 Merge pull request #3 from pschroedl/relax_python_3_10
Relax python 3.11.15 -> 3.10.15
2024-12-05 15:48:17 -08:00
Peter Schroedl 4115029965 update pip install to use main 2024-12-06 00:44:12 +01:00
Peter Schroedl 2dd9fc7ac5 temp use branch to test relaxed python 2024-12-05 23:49:42 +01:00
Peter Schroedl 38723ebb07 add ARCH_LIST and install package for cuda extensions 2024-12-05 23:37:43 +01:00
Peter Schroedl fa428ca4d9 relax python version requirement to >= 3.10 2024-12-05 22:57:43 +01:00
Peter Schroedl 8873c23a14 build cuda extension with pip install 2024-12-02 23:09:24 -08:00
Peter Schroedl 67d96c93fc create sam2 folder, fix extension path 2024-12-03 07:22:47 +01:00
PSchroedl 966b023733 Merge pull request #1 from ryanontheinside/type-mask
Refactor
2024-12-01 17:18:48 -08:00
ryanontheinstide cd16f6e98c refactor
readability, efficiency
2024-11-30 16:38:50 -05:00
ryanontheinstide 524050b695 Handle negative coordinates 2024-11-30 16:27:18 -05:00
ryanontheinstide a664d8b6d3 Reset tracking
Add reset tracking option to allow for manually updating coordinates in workflow
2024-11-30 15:30:00 -05:00
ryanontheinstide 40f9f971c0 Convert return value to type MASK 2024-11-30 14:27:28 -05:00
7 changed files with 228 additions and 80 deletions
+80 -58
View File
@@ -31,7 +31,7 @@ class DownloadAndLoadSAM2RealtimeModel:
def INPUT_TYPES(s):
return {"required": {
"model": ([
'sam2_hiera_tiny.pt',
'sam2_hiera_tiny.pt', 'sam2_hiera_small.pt',
],),
"segmentor": (
['realtime'],
@@ -64,12 +64,14 @@ class DownloadAndLoadSAM2RealtimeModel:
download_path = os.path.join(folder_paths.models_dir, "sam2")
model_path = os.path.join(download_path, model)
print("model_path: ", model_path)
url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/sam2_hiera_tiny.pt"
if not os.path.exists(download_path):
os.makedirs(download_path)
if not os.path.exists(model_path):
print(f"Downloading SAM2 model to: {model_path}")
base_url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/"
url = f"{base_url}{model}"
response = requests.get(url, stream=True)
response.raise_for_status()
@@ -82,8 +84,9 @@ class DownloadAndLoadSAM2RealtimeModel:
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
model_cfg = "sam2_hiera_t.yaml" #TODO: remove hardcoded config and path
with initialize_config_dir(config_dir=config_dir, version_base=None):
cfg = compose(config_name=model_cfg)
@@ -141,18 +144,18 @@ class Sam2RealtimeSegmentation:
"sam2_model": ("SAM2MODEL",),
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
},
"optional": {
"coordinates_positive": ("STRING", {"forceInput": True}),
"point_labels": ("STRING", {"forceInput": True}),
# "coordinates_negative": ("STRING", {"forceInput": True}),
"optional": {
"coordinates_positive": ("STRING", ),
"coordinates_negative": ("STRING", ),
"reset_tracking": ("BOOLEAN", {"default": False}),
# "bboxes": ("BBOX", ),
# "individual_objects": ("BOOLEAN", {"default": False}),
# "mask": ("MASK", ),
},
}
RETURN_NAMES = ("PROCESSED_IMAGES","MASK",)
RETURN_TYPES = ("IMAGE", "IMAGE",)
RETURN_NAMES = ("PROCESSED_IMAGES", "MASK",)
RETURN_TYPES = ("IMAGE", "MASK",)
FUNCTION = "segment_images"
CATEGORY = "SAM2-Realtime"
@@ -160,93 +163,112 @@ class Sam2RealtimeSegmentation:
self.predictor = None
self.if_init = False
def _process_coordinate_input(self, coordinates, label):
"""Helper function to process coordinate inputs safely"""
if not coordinates:
return [], []
try:
coord_list = ast.literal_eval(coordinates)
points = [tuple(map(int, point)) for point in coord_list]
labels = [label] * len(points)
return points, labels
except (ValueError, SyntaxError) as e:
print(f"Error processing coordinates: {e}")
return [], []
def _process_mask_logits(self, out_mask_logits, frame_shape, device):
"""Helper function to process mask logits"""
if out_mask_logits.shape[0] > 0:
mask = (out_mask_logits[0, 0] > 0.5).byte()
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0).float(),
size=frame_shape[:2],
mode='nearest'
).squeeze().byte().to(device)
else:
mask = torch.ones(frame_shape[:2], device=device, dtype=torch.uint8)
return mask
def segment_images(
self,
images,
sam2_model,
# keep_model_loaded,
coordinates_positive=None,
# coordinates_negative=None,
point_labels=None,
coordinates_negative=None,
reset_tracking=False,
#point_labels=None,
# bboxes=None,
# individual_objects=False,
# mask=None,
):
model = sam2_model["model"]
device = sam2_model["device"]
device = torch.device("cuda")
model.to(device)
processed_frames = []
mask_list = []
# The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
if self.predictor is None:
self.predictor = model
if reset_tracking:
self.if_init = False
self.predictor = None
def process_frame(frame, frame_idx):
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
frame = frame.to(device).float() # Keep everything in torch
if self.predictor is None:
self.predictor = model
# Process coordinates once, outside the frame loop
pos_points, pos_labels = self._process_coordinate_input(coordinates_positive, 1)
neg_points, neg_labels = self._process_coordinate_input(coordinates_negative, 0)
all_points = pos_points + neg_points
all_labels = pos_labels + neg_labels
if all_points:
points_tensor = torch.tensor([all_points], device=device)
labels_tensor = torch.tensor([all_labels], device=device)
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
for frame_idx, frame in enumerate(images):
frame = frame.to(device).float()
if not self.if_init:
self.predictor.load_first_frame(frame)
self.if_init = True
# obj_id = 1
# point = [256, 256]
# points = [point]
# labels = [1]
coordinates_positive_list = ast.literal_eval(coordinates_positive)
point_labels_list = ast.literal_eval(point_labels)
point_labels_list = list(map(int, point_labels_list))
for idx, point in enumerate(coordinates_positive_list):
point_tuple = tuple(map(int, point))
if all_points:
_, _, out_mask_logits = self.predictor.add_new_prompt(
frame_idx=0,
obj_id=idx + 1,
points=[point_tuple],
labels=[point_labels_list[idx]]
frame_idx=0,
obj_id=1,
points=points_tensor,
labels=labels_tensor,
)
# _, _, _ = self.predictor.add_new_prompt(frame_idx, obj_id, points=points, labels=labels)
else:
out_mask_logits = torch.zeros((0,), device=device)
else:
out_obj_ids, out_mask_logits = self.predictor.track(frame)
if out_mask_logits.shape[0] > 0:
# Ensure out_mask_logits is on the same device
out_mask_logits = out_mask_logits.to(device)
mask = (out_mask_logits[0, 0] > 0.5).byte()
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0).float(),
size=(frame.shape[0], frame.shape[1]),
mode='nearest'
).squeeze(0).squeeze(0).byte().to(device) # Move the interpolated mask to the correct device
else:
mask = torch.ones((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8)
# Process mask logits
mask = self._process_mask_logits(out_mask_logits, frame.shape, device)
# Ensure frame is on the same device
frame = frame.to(device)
# Create colored overlay for processed frames
mask_colored = torch.stack([mask] * 3, dim=2)
mask_colored = torch.stack([mask] * 3, dim=2).to(device) # Create 3-channel mask and move to device
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3).to(device)
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
processed_frames.append(overlayed_frame)
mask_list.append(mask)
constructed_mask = torch.add(frame * 0.1, mask_colored * 0.9).to(device)
mask_list.append(constructed_mask)
for frame_idx, img in enumerate(images):
process_frame(img, frame_idx)
# Stack masks and frames
stacked_masks = torch.stack(mask_list, dim=0)
stacked_frames = torch.stack(processed_frames, dim=0)
stacked_frames = torch.stack(processed_frames, dim=0)
return (stacked_frames, stacked_masks)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel,
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
+2 -1
View File
@@ -3,4 +3,5 @@ numpy>=1.24.4
tqdm>=4.66.1
hydra-core>=1.3.2
iopath>=0.1.10
pillow>=9.4.0
pillow>=9.4.0
git+https://github.com/pschroedl/ComfyUI-SAM2-Realtime.git@main#egg=sam2_realtime
+116
View File
@@ -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_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
# SAM decoder
sigmoid_scale_for_mem_enc: 20.0
+21 -16
View File
@@ -14,7 +14,6 @@ from tqdm import tqdm
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
class SAM2TensorPredictor(SAM2Base):
"""The predictor class to handle user interactions and manage inference states."""
@@ -55,7 +54,9 @@ class SAM2TensorPredictor(SAM2Base):
img = img.float()
else:
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)
img = torch.nn.functional.interpolate(
img.unsqueeze(0), size=(image_size, image_size), mode="bilinear", align_corners=False
@@ -68,28 +69,29 @@ class SAM2TensorPredictor(SAM2Base):
img /= img_std
height, width = img.shape[1:] # CHW format
return img, width, height
return img, width, height, orig_w, orig_h
@torch.inference_mode()
def load_first_frame(self, img):
if isinstance(img, torch.Tensor):
img = img.to(self.device) # Ensure the tensor is on the correct device
self.condition_state = self._init_state(
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["num_frames"] = len(self.condition_state["images"])
self.condition_state["video_height"] = height
self.condition_state["video_width"] = width
self._get_image_feature(frame_idx=0, batch_size=1)
def add_conditioning_frame(self, img):
if isinstance(img, torch.Tensor):
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["num_frames"] = len(self.condition_state["images"])
self._get_image_feature(
@@ -235,14 +237,15 @@ class SAM2TensorPredictor(SAM2Base):
points = torch.cat([box_coords, points], dim=1)
labels = torch.cat([box_labels, labels], dim=1)
if normalize_coords:
video_H = self.condition_state["video_height"]
video_W = self.condition_state["video_width"]
points = points / torch.tensor([video_W, video_H]).to(points.device)
#video_H = self.condition_state["video_height"]
#video_W = self.condition_state["video_width"]
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
points = points * self.image_size
points = points.to(self.condition_state["device"])
labels = labels.to(self.condition_state["device"])
if not clear_old_points:
point_inputs = point_inputs_per_frame.get(frame_idx, None)
else:
@@ -342,14 +345,16 @@ class SAM2TensorPredictor(SAM2Base):
if labels.dim() == 1:
labels = labels.unsqueeze(0) # add batch dimension
if normalize_coords:
video_H = self.condition_state["video_height"]
video_W = self.condition_state["video_width"]
points = points / torch.tensor([video_W, video_H]).to(points.device)
#video_H = self.condition_state["video_height"]
#video_W = self.condition_state["video_width"]
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
points = points * self.image_size
points = points.to(self.condition_state["device"])
labels = labels.to(self.condition_state["device"])
if not clear_old_points:
point_inputs = point_inputs_per_frame.get(frame_idx, None)
else:
@@ -769,7 +774,7 @@ class SAM2TensorPredictor(SAM2Base):
if isinstance(img, torch.Tensor):
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"]
obj_ids = self.condition_state["obj_ids"]
+1 -1
View File
@@ -58,7 +58,7 @@ def get_connected_components(mask):
- counts: A tensor of shape (N, 1, H, W) containing the area of the connected
components for foreground pixels and 0 for background pixels.
"""
from sam2 import _C
from sam2_realtime import _C
return _C.get_connected_componnets(mask.to(torch.uint8).contiguous())
+7 -3
View File
@@ -4,6 +4,10 @@
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
import os
# Set the CUDA architecture list
os.environ["TORCH_CUDA_ARCH_LIST"] = "8.0 8.6+PTX 8.7 9.0 9.0a"
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
@@ -28,7 +32,7 @@ REQUIRED_PACKAGES = [
]
def get_extensions():
srcs = ["sam2/csrc/connected_components.cu"]
srcs = ["sam2_realtime/csrc/connected_components.cu"]
compile_args = {
"cxx": [],
"nvcc": [
@@ -38,7 +42,7 @@ def get_extensions():
"-D__CUDA_NO_HALF2_OPERATORS__",
],
}
ext_modules = [CUDAExtension("sam2._C", srcs, extra_compile_args=compile_args)]
ext_modules = [CUDAExtension("sam2_realtime._C", srcs, extra_compile_args=compile_args)]
return ext_modules
@@ -53,7 +57,7 @@ setup(
license=LICENSE,
packages=find_packages(),
install_requires=REQUIRED_PACKAGES,
python_requires=">=3.11.10",
python_requires=">=3.10.15",
ext_modules=get_extensions(),
cmdclass={"build_ext": BuildExtension.with_options(no_python_abi_suffix=True)},
)