diff --git a/README.MD b/README.MD
index b8dbe29..81345e8 100644
--- a/README.MD
+++ b/README.MD
@@ -150,6 +150,7 @@ Please try downgrading the ```protobuf``` dependency package to 3.20.3, or set e
**If the dependency package error after updating, please double clicking ```repair_dependency.bat``` (for Official ComfyUI Protable) or ```repair_dependency_aki.bat``` (for ComfyUI-aki-v1.x) in the plugin folder to reinstall the dependency packages.
+* [SAM2Ultra](#SAM2Ultra) and ObjectDetector nodes support image batch.
* [SAM2Ultra](#SAM2Ultra) and [SAM2VideoUltra](#SAM2VideoUltra) nodes add support for SAM2.1 model, including [kijai](https://github.com/kijai)'s FP16 model. Download model files from [BaiduNetdisk](https://pan.baidu.com/s/1xaQYBA6ktxvAxm310HXweQ?pwd=auki) or [huggingface.co/Kijai/sam2-safetensors](https://huggingface.co/Kijai/sam2-safetensors/tree/main) and copy to ```ComfyUI/models/sam2``` folder.
* Commit [JoyCaption2Split](#JoyCaption2Split) and [LoadJoyCaption2Model](#LoadJoyCaption2Model) nodes, Sharing the model across multiple JoyCaption2 nodes improves efficiency.
* [SegmentAnythingUltra](#SegmentAnythingUltra) and [SegmentAnythingUltraV2](#SegmentAnythingUltraV2) add the ```cache_model``` option, Easy to flexibly manage VRAM usage.
diff --git a/README_CN.MD b/README_CN.MD
index a54c8dd..48a23d8 100644
--- a/README_CN.MD
+++ b/README_CN.MD
@@ -127,6 +127,7 @@ If this call came from a _pb2.py file, your generated code is out of date and mu
## 更新说明
**如果本插件更新后出现依赖包错误,请双击运行插件目录下的```install_requirements.bat```(官方便携包),或 ```install_requirements_aki.bat```(秋叶整合包) 重新安装依赖包。
+* [SAM2Ultra](#SAM2Ultra) 及 ObjectDetector 节点支持图像批次。
* [SAM2Ultra](#SAM2Ultra) 及 [SAM2VideoUltra](#SAM2VideoUltra) 节点增加支持SAM2.1模型,包括[kijai](https://github.com/kijai)量化版fp16模型。请从请从[百度网盘](https://pan.baidu.com/s/1xaQYBA6ktxvAxm310HXweQ?pwd=auki) 或者 [huggingface.co/Kijai/sam2-safetensors](https://huggingface.co/Kijai/sam2-safetensors/tree/main)下载模型文件并复制到```ComfyUI/models/sam2```文件夹。
* 添加 [JoyCaption2Split](#JoyCaption2Split) 和 [LoadJoyCaption2Model](#LoadJoyCaption2Model) 节点,在多个JoyCaption2节点时共用模型提高效率。
* [SegmentAnythingUltra](#SegmentAnythingUltra) 和 [SegmentAnythingUltraV2](#SegmentAnythingUltraV2) 增加 ```cache_model``` 参数,便于灵活管理显存。
diff --git a/__init__.py b/__init__.py
index 240d335..6ead716 100644
--- a/__init__.py
+++ b/__init__.py
@@ -10,35 +10,38 @@ NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
python = sys.executable
-base_path = os.path.dirname(folder_paths.base_path)
-extentions_folder = os.path.join(base_path, "web", "extensions", "dzNodes")
-javascript_folder = os.path.join(os.path.dirname(os.path.realpath(__file__)), "js")
-outdate_file_list = ['comfy_shared.js', 'debug.js', 'mtb_widgets.js', 'parse-css.js', 'dz_widgets.js']
-if not os.path.exists(extentions_folder):
- print('# 😺dzNodes: Making the "web\extensions\dzNodes" folder')
- os.makedirs(extentions_folder, exist_ok=True)
-else:
- for i in outdate_file_list:
- outdate_file = os.path.join(extentions_folder, i)
- if os.path.exists(outdate_file):
- os.remove(outdate_file)
+try:
+ base_path = os.path.dirname(folder_paths.base_path)
+ extentions_folder = os.path.join(base_path, "web", "extensions", "dzNodes")
+ javascript_folder = os.path.join(os.path.dirname(os.path.realpath(__file__)), "js")
+ outdate_file_list = ['comfy_shared.js', 'debug.js', 'mtb_widgets.js', 'parse-css.js', 'dz_widgets.js']
+
+ if not os.path.exists(extentions_folder):
+ print('# 😺dzNodes: Making the "web\extensions\dzNodes" folder')
+ os.makedirs(extentions_folder, exist_ok=True)
+ else:
+ for i in outdate_file_list:
+ outdate_file = os.path.join(extentions_folder, i)
+ if os.path.exists(outdate_file):
+ os.remove(outdate_file)
-result = filecmp.dircmp(javascript_folder, extentions_folder)
+ result = filecmp.dircmp(javascript_folder, extentions_folder)
-if result.left_only or result.diff_files:
- print('# 😺dzNodes: Update to javascripts files detected')
- file_list = list(result.left_only)
- file_list.extend(x for x in result.diff_files if x not in file_list)
-
- for file in file_list:
- print(f'# 😺dzNodes:: Copying {file} to extensions folder')
- src_file = os.path.join(javascript_folder, file)
- dst_file = os.path.join(extentions_folder, file)
- if os.path.exists(dst_file):
- os.remove(dst_file)
- shutil.copy(src_file, dst_file)
+ if result.left_only or result.diff_files:
+ print('# 😺dzNodes: Update to javascripts files detected')
+ file_list = list(result.left_only)
+ file_list.extend(x for x in result.diff_files if x not in file_list)
+ for file in file_list:
+ print(f'# 😺dzNodes:: Copying {file} to extensions folder')
+ src_file = os.path.join(javascript_folder, file)
+ dst_file = os.path.join(extentions_folder, file)
+ if os.path.exists(dst_file):
+ os.remove(dst_file)
+ shutil.copy(src_file, dst_file)
+except Exception as e:
+ print(f'# 😺dzNodes: Error in update js files: {e}')
def get_ext_dir(subpath=None, mkdir=False):
dir = os.path.dirname(__file__)
diff --git a/py/object_detector.py b/py/object_detector.py
index f97ec80..f58de97 100644
--- a/py/object_detector.py
+++ b/py/object_detector.py
@@ -105,6 +105,7 @@ class LS_OBJECT_DETECTOR_FL2:
def object_detector_fl2(self, image, prompt, florence2_model, sort_method, bbox_select, select_index):
+ ret_bboxes = []
bboxes = []
ret_previews = []
max_new_tokens = 512
@@ -115,27 +116,30 @@ class LS_OBJECT_DETECTOR_FL2:
model = florence2_model['model']
processor = florence2_model['processor']
- img = tensor2pil(image[0]).convert("RGB")
- task = 'caption to phrase grounding'
- from .florence2_ultra import process_image
- results, _ = process_image(model, processor, img, task,
- max_new_tokens, num_beams, do_sample,
- fill_mask, prompt)
+ for img in image:
+ img = tensor2pil(img.unsqueeze(0)).convert("RGB")
+ task = 'caption to phrase grounding'
+ from .florence2_ultra import process_image
+ results, _ = process_image(model, processor, img, task,
+ max_new_tokens, num_beams, do_sample,
+ fill_mask, prompt)
- if isinstance(results, dict):
- results["width"] = img.width
- results["height"] = img.height
+ if isinstance(results, dict):
+ results["width"] = img.width
+ results["height"] = img.height
- bboxes = self.fbboxes_to_list(results)
- bboxes = sort_bboxes(bboxes, sort_method)
- bboxes = select_bboxes(bboxes, bbox_select, select_index)
- preview = draw_bounding_boxes(img, bboxes, color="random", line_width=-1)
- ret_previews.append(pil2tensor(preview))
- if len(bboxes) == 0:
- log(f"{self.NODE_NAME} no object found", message_type='warning')
- else:
- log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
- return (standardize_bbox(bboxes), torch.cat(ret_previews, dim=0))
+ bboxes = self.fbboxes_to_list(results)
+ bboxes = sort_bboxes(bboxes, sort_method)
+ bboxes = select_bboxes(bboxes, bbox_select, select_index)
+ preview = draw_bounding_boxes(img, bboxes, color="random", line_width=-1)
+ ret_previews.append(pil2tensor(preview))
+ ret_bboxes.append(standardize_bbox(bboxes))
+ if len(bboxes) == 0:
+ log(f"{self.NODE_NAME} no object found", message_type='warning')
+ else:
+ log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+
+ return (ret_bboxes, torch.cat(ret_previews, dim=0))
def fbboxes_to_list(self, F_BBOXES) -> list:
if isinstance(F_BBOXES, str):
@@ -210,31 +214,34 @@ class LS_OBJECT_DETECTOR_MASK:
def object_detector_mask(self, object_mask, sort_method, bbox_select, select_index):
+ ret_bboxes = []
+ ret_previews = []
bboxes = []
if object_mask.dim() == 2:
object_mask = torch.unsqueeze(object_mask, 0)
- cv_mask = tensor2cv2(object_mask[0])
- cv_mask = cv2.cvtColor(cv_mask, cv2.COLOR_BGR2GRAY)
- _, binary = cv2.threshold(cv_mask, 127, 255, cv2.THRESH_BINARY)
- # invert mask
- # binary = cv2.bitwise_not(binary)
- contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
- for contour in contours:
- x, y, w, h = cv2.boundingRect(contour)
- bboxes.append([x, y, x + w, y + h])
- bboxes = sort_bboxes(bboxes, sort_method)
- bboxes = select_bboxes(bboxes, bbox_select, select_index)
- ret_previews = []
- preview = draw_bounding_boxes(tensor2pil(object_mask[0]).convert("RGB"), bboxes, color="random", line_width=-1)
- ret_previews.append(pil2tensor(preview))
+ for msk in object_mask:
+ cv_mask = tensor2cv2(msk)
+ cv_mask = cv2.cvtColor(cv_mask, cv2.COLOR_BGR2GRAY)
+ _, binary = cv2.threshold(cv_mask, 127, 255, cv2.THRESH_BINARY)
+ # invert mask
+ # binary = cv2.bitwise_not(binary)
+ contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
+ for contour in contours:
+ x, y, w, h = cv2.boundingRect(contour)
+ bboxes.append([x, y, x + w, y + h])
+ bboxes = sort_bboxes(bboxes, sort_method)
+ bboxes = select_bboxes(bboxes, bbox_select, select_index)
+ preview = draw_bounding_boxes(tensor2pil(msk).convert("RGB"), bboxes, color="random", line_width=-1)
+ ret_previews.append(pil2tensor(preview))
- if len(bboxes) == 0:
- log(f"{self.NODE_NAME} no object found", message_type='warning')
- else:
- log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ if len(bboxes) == 0:
+ log(f"{self.NODE_NAME} no object found", message_type='warning')
+ else:
+ log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ ret_bboxes.append(standardize_bbox(bboxes))
- return (standardize_bbox(bboxes), torch.cat(ret_previews, dim=0))
+ return (ret_bboxes, torch.cat(ret_previews, dim=0))
class LS_OBJECT_DETECTOR_YOLO8:
@@ -271,31 +278,34 @@ class LS_OBJECT_DETECTOR_YOLO8:
model_path = os.path.join(folder_paths.models_dir, 'yolo')
yolo_model = YOLO(os.path.join(model_path, yolo_model))
+ ret_bboxes = []
bboxes = []
ret_previews = []
- img = torch.unsqueeze(image[0], 0)
- _image = tensor2pil(img)
- results = yolo_model(_image, retina_masks=True)
- for result in results:
- yolo_plot_image = cv2.cvtColor(result.plot(), cv2.COLOR_BGR2RGB)
+ for img in image:
+ img = torch.unsqueeze(img.unsqueeze(0), 0)
+ _image = tensor2pil(img)
+ results = yolo_model(_image, retina_masks=True)
+ for result in results:
+ yolo_plot_image = cv2.cvtColor(result.plot(), cv2.COLOR_BGR2RGB)
- # no mask, if have box, draw box
- if result.boxes is not None and len(result.boxes.xyxy) > 0:
- for box in result.boxes:
- x1, y1, x2, y2 = box.xyxy[0].cpu().numpy()
- bboxes.append([x1, y1, x2, y2])
- bboxes = sort_bboxes(bboxes, sort_method)
- bboxes = select_bboxes(bboxes, bbox_select, select_index)
- preview = draw_bounding_boxes(_image.convert("RGB"), bboxes, color="random", line_width=-1)
- ret_previews.append(pil2tensor(preview))
+ # no mask, if have box, draw box
+ if result.boxes is not None and len(result.boxes.xyxy) > 0:
+ for box in result.boxes:
+ x1, y1, x2, y2 = box.xyxy[0].cpu().numpy()
+ bboxes.append([x1, y1, x2, y2])
+ bboxes = sort_bboxes(bboxes, sort_method)
+ bboxes = select_bboxes(bboxes, bbox_select, select_index)
+ preview = draw_bounding_boxes(_image.convert("RGB"), bboxes, color="random", line_width=-1)
+ ret_previews.append(pil2tensor(preview))
- if len(bboxes) == 0:
- log(f"{self.NODE_NAME} no object found", message_type='warning')
- else:
- log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ if len(bboxes) == 0:
+ log(f"{self.NODE_NAME} no object found", message_type='warning')
+ else:
+ log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ ret_bboxes.append(standardize_bbox(bboxes))
- return (standardize_bbox(bboxes), torch.cat(ret_previews, dim=0),)
+ return (ret_bboxes, torch.cat(ret_previews, dim=0),)
class LS_OBJECT_DETECTOR_YOLOWORLD:
@@ -332,33 +342,44 @@ class LS_OBJECT_DETECTOR_YOLOWORLD:
def object_detector_yoloworld(self, image, yolo_world_model,
confidence_threshold, nms_iou_threshold, prompt,
sort_method, bbox_select, select_index):
+ ret_previews = []
+ ret_bboxes = []
+
import supervision as sv
model=self.load_yolo_world_model(yolo_world_model, prompt)
- infer_outputs = []
- img = (255 * image[0].cpu().numpy()).astype(np.uint8)
- results = model.infer(
- img, confidence=confidence_threshold)
- detections = sv.Detections.from_inference(results)
- detections = detections.with_nms(
- class_agnostic=False,
- threshold=nms_iou_threshold
- )
- infer_outputs.append(detections)
- bboxes = infer_outputs[0].xyxy.tolist()
- bboxes = [[int(value) for value in sublist] for sublist in bboxes]
- bboxes = sort_bboxes(bboxes, sort_method)
- bboxes = select_bboxes(bboxes, bbox_select, select_index)
- ret_previews = []
- preview = draw_bounding_boxes(tensor2pil(image[0]).convert('RGB'), bboxes, color="random", line_width=-1)
- ret_previews.append(pil2tensor(preview))
- if len(bboxes) == 0:
- log(f"{self.NODE_NAME} no object found", message_type='warning')
- else:
- log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ for i in image:
+ infer_outputs = []
+ # img = (255 * img.unsqueeze(0).cpu().numpy()).astype(np.uint8)
+ img = tensor2np(i)
+ results = model.infer(
+ img, confidence=confidence_threshold)
+ detections = sv.Detections.from_inference(results)
+ detections = detections.with_nms(
+ class_agnostic=False,
+ threshold=nms_iou_threshold
+ )
+ infer_outputs.append(detections)
- return (standardize_bbox(bboxes), torch.cat(ret_previews, dim=0))
+ if len(infer_outputs[0].xyxy) > 0:
+ bboxes = infer_outputs[0].xyxy.tolist()
+ bboxes = [[int(value) for value in sublist] for sublist in bboxes]
+ bboxes = sort_bboxes(bboxes, sort_method)
+ bboxes = select_bboxes(bboxes, bbox_select, select_index)
+ else:
+ bboxes = [[0, 0, i.shape[1], i.shape[0]]]
+
+ preview = draw_bounding_boxes(tensor2pil(i.unsqueeze(0)).convert('RGB'), bboxes, color="random", line_width=-1)
+ ret_previews.append(pil2tensor(preview))
+
+ if len(bboxes) == 0:
+ log(f"{self.NODE_NAME} no object found", message_type='warning')
+ else:
+ log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info')
+ ret_bboxes.append(standardize_bbox(bboxes))
+
+ return (ret_bboxes, torch.cat(ret_previews, dim=0))
def process_categories(self, categories: str) -> List[str]:
return [category.strip().lower() for category in categories.split(',')]
diff --git a/py/sam_2_ultra.py b/py/sam_2_ultra.py
index 1deff2b..66cf9da 100644
--- a/py/sam_2_ultra.py
+++ b/py/sam_2_ultra.py
@@ -285,9 +285,9 @@ class LS_SAM2_ULTRA:
log(f"{self.NODE_NAME}: Downloading SAM2 model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/sam2-safetensors",
- allow_patterns=[f"*{sam2_model}*"],
- local_dir=sam2_path,
- local_dir_use_symlinks=False)
+ allow_patterns=[f"*{sam2_model}*"],
+ local_dir=sam2_path,
+ local_dir_use_symlinks=False)
model_mapping = {
"2.0": {
@@ -311,81 +311,115 @@ class LS_SAM2_ULTRA:
None
)
log(f"{self.NODE_NAME}: Using model config: {model_cfg_path}")
-
model = load_model(model_path, model_cfg_path, segmentor, dtype, device)
offload_device = mm.unet_offload_device()
- # B, H, W, C = image.shape
indexs = extract_numbers(select_index)
- # Handle possible bboxes
- if len(bboxes) == 0:
- log(f"{self.NODE_NAME} skipped, because bboxes is empty.", message_type='error')
- return (image, None)
- else:
- boxes_np_batch = []
- for bbox_list in bboxes:
- boxes_np = []
- for bbox in bbox_list:
- boxes_np.append(bbox)
- boxes_np = np.array(boxes_np)
- boxes_np_batch.append(boxes_np)
- if bbox_select == "all":
- final_box = np.array(boxes_np_batch)
- elif bbox_select == "by_index":
- final_box = []
- try:
- for i in indexs:
- final_box.append(boxes_np_batch[i])
- except IndexError:
- log(f"{self.NODE_NAME} invalid bbox index {i}", message_type='warning')
- else:
- final_box = np.array(boxes_np_batch[0])
- # final_labels = None
-
- mask_list = []
try:
model.to(device)
except:
model.model.to(device)
-
autocast_condition = not mm.is_device_mps(device)
- with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
- image_np = (image.contiguous() * 255).byte().numpy()
- comfy_pbar = ProgressBar(len(image_np))
- tqdm_pbar = tqdm(total=len(image_np), desc="Processing Images")
- for i in range(len(image_np)):
- model.set_image(image_np[i])
- # if len(image_np) > 1:
- # input_box = final_box[i]
- input_box = final_box
+ for index in range(len(image)):
+ img = image[index].unsqueeze(0)
- out_masks, scores, logits = model.predict(
- point_coords=None,
- point_labels=None,
- box=input_box,
- multimask_output=True,
- mask_input=None,
- )
-
- if out_masks.ndim == 3:
- sorted_ind = np.argsort(scores)[::-1]
- out_masks = out_masks[sorted_ind][0] # choose only the best result for now
- # scores = scores[sorted_ind]
- # logits = logits[sorted_ind]
- mask_list.append(np.expand_dims(out_masks, axis=0))
+ # Handle possible bboxes
+ if len(bboxes) == 0:
+ log(f"{self.NODE_NAME} skipped, because bboxes is empty.", message_type='error')
+ return (image, None)
+ else:
+ boxes_np_batch = []
+ for bbox_list in bboxes[index]:
+ boxes_np = []
+ for bbox in bbox_list:
+ boxes_np.append(bbox)
+ boxes_np = np.array(boxes_np)
+ boxes_np_batch.append(boxes_np)
+ if bbox_select == "all":
+ final_box = np.array(boxes_np_batch)
+ elif bbox_select == "by_index":
+ final_box = []
+ try:
+ for i in indexs:
+ final_box.append(boxes_np_batch[i])
+ except IndexError:
+ log(f"{self.NODE_NAME} invalid bbox index {i}", message_type='warning')
else:
- _, _, H, W = out_masks.shape
- # Combine masks for all object IDs in the frame
- combined_mask = np.zeros((H, W), dtype=bool)
- for out_mask in out_masks:
- combined_mask = np.logical_or(combined_mask, out_mask)
- combined_mask = combined_mask.astype(np.uint8)
- mask_list.append(combined_mask)
- comfy_pbar.update(1)
- tqdm_pbar.update(1)
+ final_box = np.array(boxes_np_batch[0])
+
+ mask_list = []
+
+ with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
+
+ image_np = (img.contiguous() * 255).byte().numpy()
+ comfy_pbar = ProgressBar(len(image_np))
+ tqdm_pbar = tqdm(total=len(image_np), desc="Processing Images")
+ for i in range(len(image_np)):
+ model.set_image(image_np[i])
+ # if len(image_np) > 1:
+ # input_box = final_box[i]
+ input_box = final_box
+
+ out_masks, scores, logits = model.predict(
+ point_coords=None,
+ point_labels=None,
+ box=input_box,
+ multimask_output=True,
+ mask_input=None,
+ )
+
+ if out_masks.ndim == 3:
+ sorted_ind = np.argsort(scores)[::-1]
+ out_masks = out_masks[sorted_ind][0] # choose only the best result for now
+ # scores = scores[sorted_ind]
+ # logits = logits[sorted_ind]
+ mask_list.append(np.expand_dims(out_masks, axis=0))
+ else:
+ _, _, H, W = out_masks.shape
+ # Combine masks for all object IDs in the frame
+ combined_mask = np.zeros((H, W), dtype=bool)
+ for out_mask in out_masks:
+ combined_mask = np.logical_or(combined_mask, out_mask)
+ combined_mask = combined_mask.astype(np.uint8)
+ mask_list.append(combined_mask)
+ comfy_pbar.update(1)
+ tqdm_pbar.update(1)
+
+ out_list = []
+ for mask in mask_list:
+ mask_tensor = torch.from_numpy(mask)
+ mask_tensor = mask_tensor.permute(1, 2, 0)
+ mask_tensor = mask_tensor[:, :, 0]
+ out_list.append(mask_tensor)
+ mask_tensor = torch.stack(out_list, dim=0).cpu().float()
+ _mask = mask_tensor.squeeze()
+
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+ orig_image = tensor2pil(img)
+ detail_range = detail_erode + detail_dilate
+ if process_detail:
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(pil2tensor(orig_image), _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(pil2tensor(orig_image), _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device,
+ max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = tensor2pil(_mask)
+
+ ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
if cache_model:
try:
@@ -398,40 +432,6 @@ class LS_SAM2_ULTRA:
else:
del model
clear_memory()
-
- out_list = []
- for mask in mask_list:
- mask_tensor = torch.from_numpy(mask)
- mask_tensor = mask_tensor.permute(1, 2, 0)
- mask_tensor = mask_tensor[:, :, 0]
- out_list.append(mask_tensor)
- mask_tensor = torch.stack(out_list, dim=0).cpu().float()
- _mask = mask_tensor.squeeze()
-
- if detail_method == 'VITMatte(local)':
- local_files_only = True
- else:
- local_files_only = False
- orig_image = tensor2pil(image[0])
- detail_range = detail_erode + detail_dilate
- if process_detail:
- if detail_method == 'GuidedFilter':
- _mask = guided_filter_alpha(pil2tensor(orig_image), _mask, detail_range // 6 + 1)
- _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
- elif detail_method == 'PyMatting':
- _mask = tensor2pil(mask_edge_detail(pil2tensor(orig_image), _mask, detail_range // 8 + 1, black_point, white_point))
- else:
- _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
- _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device,
- max_megapixels=max_megapixels)
- _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
- else:
- _mask = tensor2pil(_mask)
-
- ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
- ret_images.append(pil2tensor(ret_image))
- ret_masks.append(image2mask(_mask))
-
log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0))
diff --git a/pyproject.toml b/pyproject.toml
index b0861e7..bdf6d4a 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,7 +1,7 @@
[project]
name = "comfyui_layerstyle"
description = "A set of nodes for ComfyUI it generate image like Adobe Photoshop's Layer Style. the Drop Shadow is first completed node, and follow-up work is in progress."
-version = "1.0.83"
+version = "1.0.85"
license = "MIT"
dependencies = ["numpy", "pillow", "torch", "matplotlib", "Scipy", "scikit_image", "scikit_learn", "opencv-contrib-python", "pymatting", "segment_anything", "timm", "addict", "yapf", "colour-science", "wget", "mediapipe", "loguru", "typer_config", "fastapi", "rich", "google-generativeai", "diffusers", "omegaconf", "tqdm", "transformers", "kornia", "image-reward", "ultralytics", "blend_modes", "blind-watermark", "qrcode", "pyzbar", "transparent-background", "huggingface_hub", "accelerate", "bitsandbytes", "torchscale", "wandb", "hydra-core", "psd-tools", "inference-cli[yolo-world]", "inference-gpu[yolo-world]", "onnxruntime", "peft", "iopath"]