Author SHA1 Message Date
Peter Schroedl 4a3b0beac1 change node input to json 2024-12-03 09:46:45 -08:00
Peter Schroedl 28e19ce5d2 add node to calc center of BBOX data 2024-12-03 09:40:10 -08:00
4 changed files with 52 additions and 29 deletions
+34 -2
View File
@@ -261,12 +261,44 @@ class Sam2RealtimeSegmentation:
return (stacked_frames, stacked_masks)
class BoundingBoxToCenter:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"bbox_data": ("JSON",),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("center_coordinates",)
FUNCTION = "convert_bbox_to_center"
CATEGORY = "SAM2-Realtime"
def convert_bbox_to_center(self, bbox_data):
try:
bbox_list = ast.literal_eval(bbox_data)
tlx, tly, brx, bry = bbox_list[0][0]
center_x = int((tlx + brx) / 2)
center_y = int((tly + bry) / 2)
center_coords = f"[[{center_x}, {center_y}]]"
return (center_coords,)
except (ValueError, SyntaxError, IndexError) as e:
print(f"Error processing bounding box data: {e}")
return ("[[0, 0]]",)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel,
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation,
"BoundingBoxToCenter": BoundingBoxToCenter
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation",
"BoundingBoxToCenter": "BoundingBox To Center"
}
+1 -1
View File
@@ -4,4 +4,4 @@ tqdm>=4.66.1
hydra-core>=1.3.2
iopath>=0.1.10
pillow>=9.4.0
git+https://github.com/pschroedl/ComfyUI-SAM2-Realtime.git@main#egg=sam2_realtime
-e .
+16 -21
View File
@@ -14,6 +14,7 @@ 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."""
@@ -54,9 +55,7 @@ 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
@@ -69,29 +68,28 @@ class SAM2TensorPredictor(SAM2Base):
img /= img_std
height, width = img.shape[1:] # CHW format
return img, width, height, orig_w, orig_h
return img, width, height
@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, orig_w, orig_h = self.prepare_data(img, image_size=self.image_size)
self._orig_hw = (orig_w, orig_h)
img, width, height = self.prepare_data(img, image_size=self.image_size)
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(
@@ -237,15 +235,14 @@ 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"]
orig_w, orig_h = self._orig_hw
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
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)
# 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:
@@ -345,16 +342,14 @@ 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"]
orig_w, orig_h = self._orig_hw
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
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)
# 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:
@@ -774,7 +769,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 -5
View File
@@ -4,10 +4,6 @@
# 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
@@ -57,7 +53,7 @@ setup(
license=LICENSE,
packages=find_packages(),
install_requires=REQUIRED_PACKAGES,
python_requires=">=3.10.15",
python_requires=">=3.11.10",
ext_modules=get_extensions(),
cmdclass={"build_ext": BuildExtension.with_options(no_python_abi_suffix=True)},
)