Remove useless dependencies, Florence2 fixes
This commit is contained in:
@@ -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)
|
||||
|
||||
+2
-2
@@ -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"
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
pyyaml
|
||||
iopath
|
||||
@@ -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
|
||||
|
||||
+6
-6
@@ -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])
|
||||
|
||||
Reference in New Issue
Block a user