From 8ffc4efcd6a12649ba5f524602dd24a74ea8145b Mon Sep 17 00:00:00 2001 From: "Dr.Lt.Data" Date: Tue, 9 Apr 2024 21:16:29 +0900 Subject: [PATCH] fix: better compatibility with other nodes. https://github.com/ltdrdata/ComfyUI-Impact-Pack/issues/547 --- modules/impact/config.py | 2 +- modules/impact/core.py | 7 ++++--- modules/impact/impact_pack.py | 7 +++++-- 3 files changed, 10 insertions(+), 6 deletions(-) diff --git a/modules/impact/config.py b/modules/impact/config.py index 71701fc..4c4a238 100644 --- a/modules/impact/config.py +++ b/modules/impact/config.py @@ -2,7 +2,7 @@ import configparser import os -version_code = [4, 87, 4] +version_code = [4, 87, 5] version = f"V{version_code[0]}.{version_code[1]}" + (f'.{version_code[2]}' if len(version_code) > 2 else '') dependency_version = 20 diff --git a/modules/impact/core.py b/modules/impact/core.py index b961d54..2d007b4 100644 --- a/modules/impact/core.py +++ b/modules/impact/core.py @@ -515,9 +515,10 @@ class ESAMWrapper: return [detected_masks.squeeze(0)] -def make_sam_mask(sam_obj, segs, image, detection_hint, dilation, +def make_sam_mask(sam, segs, image, detection_hint, dilation, threshold, bbox_expansion, mask_hint_threshold, mask_hint_use_negative): + sam_obj = sam.sam_wrapper sam_obj.prepare_device() try: @@ -775,9 +776,9 @@ def every_three_pick_last(stacked_masks): return selected_masks -def make_sam_mask_segmented(sam_obj, segs, image, detection_hint, dilation, +def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, threshold, bbox_expansion, mask_hint_threshold, mask_hint_use_negative): - + sam_obj = sam.sam_wrapper sam_obj.prepare_device() try: diff --git a/modules/impact/impact_pack.py b/modules/impact/impact_pack.py index b98583c..781ca00 100644 --- a/modules/impact/impact_pack.py +++ b/modules/impact/impact_pack.py @@ -111,8 +111,10 @@ class SAMLoader: esam = esam_loader.load_esam_model('CUDA')[0] sam_obj = core.ESAMWrapper(esam, device_mode) + esam.sam_wrapper = sam_obj + print(f"Loads EfficientSAM model: (device:{device_mode})") - return (sam_obj, ) + return (esam, ) modelname = folder_paths.get_full_path("sams", model_name) @@ -136,9 +138,10 @@ class SAMLoader: is_auto_mode = device_mode == "AUTO" sam_obj = core.SAMWrapper(sam, is_auto_mode=is_auto_mode, safe_to_gpu=safe_to) + sam.sam_wrapper = sam_obj print(f"Loads SAM model: {modelname} (device:{device_mode})") - return (sam_obj, ) + return (sam, ) class ONNXDetectorForEach: