diff --git a/nodes.py b/nodes.py index f8469fb..99c83a4 100644 --- a/nodes.py +++ b/nodes.py @@ -128,31 +128,32 @@ class Florence2toCoordinates: CATEGORY = "SAM2" def segment(self, data, index, batch=False): - print(data) try: coordinates = coordinates.replace("'", '"') coordinates = json.loads(coordinates) except: coordinates = data - print("Type of data:", type(data)) - print("Data:", data) + if len(data)==0: return (json.dumps([{'x': 0, 'y': 0}]),) center_points = [] + def get_bboxes(item): + return item["bboxes"] if isinstance(item, dict) else item + if index.strip(): # Check if index is not empty indexes = [int(i) for i in index.split(",")] else: # If index is empty, use all indices from data[0] - indexes = list(range(len(data[0]))) + indexes = list(range(len(get_bboxes(data[0])))) print("Indexes:", indexes) bboxes = [] if batch: for idx in indexes: - if 0 <= idx < len(data[0]): + if 0 <= idx < len(get_bboxes(data[0])): for i in range(len(data)): - bbox = data[i][idx] if i < len(data) else data[0].get("bboxes", [])[idx] + bbox = get_bboxes(data[i])[idx] min_x, min_y, max_x, max_y = bbox center_x = int((min_x + max_x) / 2) center_y = int((min_y + max_y) / 2) @@ -160,8 +161,8 @@ class Florence2toCoordinates: bboxes.append(bbox) else: for idx in indexes: - if 0 <= idx < len(data[0]): - bbox = data[0].get("bboxes", [])[idx] if isinstance(data[0], dict) else data[0][idx] + if 0 <= idx < len(get_bboxes(data[0])): + bbox = get_bboxes(data[0])[idx] min_x, min_y, max_x, max_y = bbox center_x = int((min_x + max_x) / 2) center_y = int((min_y + max_y) / 2) diff --git a/pyproject.toml b/pyproject.toml index 1a9050b..1abdecc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [project] name = "comfyui-segment-anything-2" description = "Nodes to use [a/segment-anything-2](https://github.com/facebookresearch/segment-anything-2) for image or video segmentation." -version = "1.0.0" +version = "1.0.1" license = {file = "LICENSE"} -dependencies = ["pyyaml", "numpy<=1.26.4"] +dependencies = [] [project.urls] Repository = "https://github.com/kijai/ComfyUI-segment-anything-2" diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index e55b533..0000000 --- a/requirements.txt +++ /dev/null @@ -1,2 +0,0 @@ -pyyaml -iopath diff --git a/sam2/modeling/backbones/hieradet.py b/sam2/modeling/backbones/hieradet.py index 217a054..9003309 100644 --- a/sam2/modeling/backbones/hieradet.py +++ b/sam2/modeling/backbones/hieradet.py @@ -10,7 +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 iopath.common.file_io import g_pathmgr from ....sam2.modeling.backbones.utils import ( PatchEmbed, @@ -264,10 +264,10 @@ 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)) + # 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 diff --git a/sam2/utils/misc.py b/sam2/utils/misc.py index abb888a..1e49097 100644 --- a/sam2/utils/misc.py +++ b/sam2/utils/misc.py @@ -19,12 +19,12 @@ def get_sdpa_settings(): old_gpu = torch.cuda.get_device_properties(0).major < 7 # only use Flash Attention on Ampere (8.0) or newer GPUs use_flash_attn = torch.cuda.get_device_properties(0).major >= 8 and platform.system() == 'Linux' - if not use_flash_attn: - warnings.warn( - "Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.", - category=UserWarning, - stacklevel=2, - ) + # if not use_flash_attn: + # warnings.warn( + # "Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.", + # category=UserWarning, + # stacklevel=2, + # ) # keep math kernel for PyTorch versions before 2.2 (Flash Attention v2 is only # available on PyTorch 2.2+, while Flash Attention v1 cannot handle all cases) pytorch_version = tuple(int(v) for v in torch.__version__.split(".")[:2])