From 27c5368bcc83941ce3290ff1aeb126b15a6d8922 Mon Sep 17 00:00:00 2001 From: "Dr.Lt.Data" Date: Mon, 7 Jul 2025 00:34:10 +0900 Subject: [PATCH] Support SAM2 models. https://github.com/ltdrdata/ComfyUI-Impact-Pack/issues/684 --- README.md | 8 +- __init__.py | 10 +- modules/impact/config.py | 2 +- modules/impact/core.py | 196 +++++++++++++++++++++++++++------- modules/impact/detectors.py | 67 ++++++++++++ modules/impact/impact_pack.py | 46 +++++--- modules/impact/utils.py | 48 +++++++++ pyproject.toml | 2 +- requirements.txt | 3 +- 9 files changed, 321 insertions(+), 61 deletions(-) diff --git a/README.md b/README.md index 1c5b1a9..7544def 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,7 @@ This node pack helps to conveniently enhance images through Detector, Detailer, NOTE: The UltralyticsDetectorProvider node is not part of the ComfyUI-Impact-Pack. To use the UltralyticsDetectorProvider node, please install the ComfyUI-Impact-Subpack separately. ## NOTICE +* V8.18: Support [facebookresearch/sam2](https://github.com/facebookresearch/sam2) models * V8.0: The `Impact Subpack` is no longer installed automatically. To use `UltralyticsDetectorProvider` nodes, please install the `Impact Subpack` separately. * V7.6: Automatic installation is no longer supported. Please install using ComfyUI-Manager, or manually install requirements.txt and run install.py to complete the installation. * V7.0: Supports Switch based on Execution Model Inversion. @@ -59,7 +60,7 @@ NOTE: The UltralyticsDetectorProvider node is not part of the ComfyUI-Impact-Pac ## Custom Nodes ### [Detector nodes](https://github.com/ltdrdata/ComfyUI-extension-tutorials/blob/Main/ComfyUI-Impact-Pack/tutorial/detectors.md) - * `SAMLoader` - Loads the SAM model. + * `SAMLoader (Impact)` - Loads the SAM model. * `ONNXDetectorProvider` - Loads the ONNX model to provide BBOX_DETECTOR. * `CLIPSegDetectorProvider` - Wrapper for CLIPSeg to provide BBOX_DETECTOR. * You need to install the ComfyUI-CLIPSeg node extension. @@ -70,6 +71,9 @@ NOTE: The UltralyticsDetectorProvider node is not part of the ComfyUI-Impact-Pac * As a result, it outputs the `combined_mask`, which is a unified mask, and `batch_masks`, which are multiple masks grouped together in batch form. * While `batch_masks` may not be completely separated, it provides functionality to perform some level of segmentation. * `Simple Detector (SEGS)` - Operating primarily with `BBOX_DETECTOR`, and with the additional provision of `SAM_MODEL` or `SEGM_DETECTOR`, this node internally generates improved SEGS through mask operations on both *bbox* and *silhouette*. It serves as a convenient tool to simplify a somewhat intricate workflow. + * `Simple Detector for Video (SEGS)` – Performs detection on videos composed of image frames. Instead of using a single mask, it performs detection individually on each image frame and generates a SEGS object with a batch of masks. + * `SAM2 Video Detector (SEGS)` – Similar to `Simple Detector for Video (SEGS)`, but utilizes SAM2’s video tracking technology to generate a SEGS object with a batch of masks. + * To use this node, you must select a SAM2 model in the SAMLoader. ### ControlNet, IPAdapter * `ControlNetApply (SEGS)` - To apply ControlNet in SEGS, you need to use the Preprocessor Provider node from the Inspire Pack to utilize this node. @@ -101,7 +105,7 @@ NOTE: The UltralyticsDetectorProvider node is not part of the ComfyUI-Impact-Pac * `DetailerDebug (SEGS)` - Refines the image based on SEGS. Additionally, it provides the ability to monitor the cropped image and the refined image of the cropped image. * To prevent regeneration caused by the seed that does not change every time when using 'external_seed', please disable the 'seed random generate' option in the 'Detailer...' node. * `MASK to SEGS` - Generates SEGS based on the mask. - * `MASK to SEGS For AnimateDiff` - Generates SEGS based on the mask for AnimateDiff. + * `MASK to SEGS For Video` - Generates SEGS based on the mask for Video. (Renamed from `MASK to SEGS For AnimateDiff`) * When using a single mask, convert it to SEGS to apply it to the entire frame. * When using a batch mask, the contour fill feature is disabled. * `MediaPipe FaceMesh to SEGS` - Separate each landmark from the mediapipe facemesh image to create labeled SEGS. diff --git a/__init__.py b/__init__.py index 7e31379..b233c2e 100644 --- a/__init__.py +++ b/__init__.py @@ -159,6 +159,7 @@ NODE_CLASS_MAPPINGS = { "SegmDetectorSEGS": SegmDetectorForEach, "ONNXDetectorSEGS": BboxDetectorForEach, "ImpactSimpleDetectorSEGS_for_AD": SimpleDetectorForAnimateDiff, + "ImpactSAM2VideoDetectorSEGS": SAM2VideoDetectorSEGS, "ImpactSimpleDetectorSEGS": SimpleDetectorForEach, "ImpactSimpleDetectorSEGSPipe": SimpleDetectorForEachPipe, "ImpactControlNetApplySEGS": ControlNetApplySEGS, @@ -304,7 +305,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "BboxDetectorSEGS": "BBOX Detector (SEGS)", "SegmDetectorSEGS": "SEGM Detector (SEGS)", "ONNXDetectorSEGS": "ONNX Detector (SEGS/legacy) - use BBOXDetector", - "ImpactSimpleDetectorSEGS_for_AD": "Simple Detector for AnimateDiff (SEGS)", + "ImpactSimpleDetectorSEGS_for_AD": "Simple Detector for Video (SEGS)", + "ImpactSAM2VideoDetectorSEGS": "SAM2 Video Detector (SEGS)", "ImpactSimpleDetectorSEGS": "Simple Detector (SEGS)", "ImpactSimpleDetectorSEGSPipe": "Simple Detector (SEGS/pipe)", "ImpactControlNetApplySEGS": "ControlNetApply (SEGS) - DEPRECATED", @@ -316,7 +318,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SegsToCombinedMask": "SEGS to MASK (combined)", "MediaPipeFaceMeshToSEGS": "MediaPipe FaceMesh to SEGS", "MaskToSEGS": "MASK to SEGS", - "MaskToSEGS_for_AnimateDiff": "MASK to SEGS for AnimateDiff", + "MaskToSEGS_for_AnimateDiff": "MASK to SEGS for Video", "BitwiseAndMaskForEach": "Pixelwise(SEGS & SEGS)", "SubtractMaskForEach": "Pixelwise(SEGS - SEGS)", "ImpactSegsAndMask": "Pixelwise(SEGS & MASK)", @@ -331,8 +333,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "DetailerForEachPipe": "Detailer (SEGS/pipe)", "DetailerForEachDebug": "DetailerDebug (SEGS)", "DetailerForEachDebugPipe": "DetailerDebug (SEGS/pipe)", - "SEGSDetailerForAnimateDiff": "SEGSDetailer For AnimateDiff (SEGS/pipe)", - "DetailerForEachPipeForAnimateDiff": "Detailer For AnimateDiff (SEGS/pipe)", + "SEGSDetailerForAnimateDiff": "SEGSDetailer For Video (SEGS/pipe)", + "DetailerForEachPipeForAnimateDiff": "Detailer For Video (SEGS/pipe)", "SEGSUpscaler": "Upscaler (SEGS)", "SEGSUpscalerPipe": "Upscaler (SEGS/pipe)", diff --git a/modules/impact/config.py b/modules/impact/config.py index 4134998..509650b 100644 --- a/modules/impact/config.py +++ b/modules/impact/config.py @@ -1,7 +1,7 @@ import configparser import os -version_code = [8, 17, 1] +version_code = [8, 18] version = f"V{version_code[0]}.{version_code[1]}" + (f'.{version_code[2]}' if len(version_code) > 2 else '') dependency_version = 24 diff --git a/modules/impact/core.py b/modules/impact/core.py index 140f9ee..0f6e17a 100644 --- a/modules/impact/core.py +++ b/modules/impact/core.py @@ -1,17 +1,15 @@ -import copy import os import warnings -import numpy import torch +from sam2.sam2_image_predictor import SAM2ImagePredictor from segment_anything import SamPredictor from comfy_extras.nodes_custom_sampler import Noise_RandomNoise from impact.utils import * from collections import namedtuple import numpy as np -from skimage.measure import label -from PIL import ImageOps +from PIL import ImageOps, Image import nodes import comfy_extras.nodes_upscale_model as model_upscale @@ -26,6 +24,9 @@ from impact import utils from impact import impact_sampling from concurrent.futures import ThreadPoolExecutor import inspect +from collections import OrderedDict +from sam2.build_sam import build_sam2, build_sam2_video_predictor +import torch.nn.functional as F try: @@ -83,7 +84,7 @@ def set_previewbridge_image(node_id, file, item): def erosion_mask(mask, grow_mask_by): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) w = mask.shape[1] h = mask.shape[0] @@ -139,7 +140,7 @@ def mix_noise(from_noise, to_noise, strength, variation_method): class REGIONAL_PROMPT: def __init__(self, mask, sampler, variation_seed=0, variation_strength=0.0, variation_method='linear'): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) self.mask = mask self.sampler = sampler @@ -199,7 +200,7 @@ def create_segmasks(results): def gen_detection_hints_from_mask_area(x, y, mask, threshold, use_negative): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) points = [] plabs = [] @@ -318,7 +319,7 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max print(f"Detailer: segment upscale for ({bbox_w, bbox_h}) | crop region {w, h} x {upscale} -> {new_w, new_h}") # upscale - upscaled_image = tensor_resize(image, new_w, new_h) + upscaled_image = utils.tensor_resize(image, new_w, new_h) if detailer_hook is not None: upscaled_image = detailer_hook.post_upscale(upscaled_image, noise_mask) @@ -339,7 +340,7 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max print(f"[Impact Pack] ComfyUI is an outdated version.") positive, negative, latent_image = imc_encode(positive, negative, upscaled_image, vae, noise_mask) else: - latent_image = to_latent_image(upscaled_image, vae, vae_tiled_encode=vae_tiled_encode) + latent_image = utils.to_latent_image(upscaled_image, vae, vae_tiled_encode=vae_tiled_encode) if noise_mask is not None: latent_image['noise_mask'] = noise_mask @@ -398,7 +399,7 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max refined_image = detailer_hook.post_decode(refined_image) # downscale - refined_image = tensor_resize(refined_image, w, h) + refined_image = utils.tensor_resize(refined_image, w, h) # prevent mixing of device refined_image = refined_image.cpu() @@ -493,10 +494,10 @@ def enhance_detail_for_animatediff(image_frames, model, clip, vae, guide_size, g image = torch.from_numpy(image).unsqueeze(0) # upscale - upscaled_image = tensor_resize(image, new_w, new_h) + upscaled_image = utils.tensor_resize(image, new_w, new_h) # ksampler - samples = to_latent_image(upscaled_image, vae)['samples'] + samples = utils.to_latent_image(upscaled_image, vae)['samples'] if latent_frames is None: latent_frames = samples @@ -621,6 +622,118 @@ class SAMWrapper: return sam_predict(predictor, points, plabs, bbox, threshold) +class SAM2Wrapper: + def __init__(self, config, modelname, is_auto_mode, safe_to_gpu=None, device_mode="AUTO"): + self.config = config + self.modelname = modelname + self.image_predictor = None + self.video_predictor = None + self.device_mode = device_mode + self.safe_to_gpu = safe_to_gpu if safe_to_gpu is not None else SafeToGPU_stub() + self.is_auto_mode = is_auto_mode + + def prepare_device(self): + pass + + def prepare_image_device(self): + if self.is_auto_mode: + device = comfy.model_management.get_torch_device() + self.safe_to_gpu.to_device(self.image_predictor.model, device=device) + + def prepare_video_device(self): + if self.is_auto_mode: + device = comfy.model_management.get_torch_device() + self.safe_to_gpu.to_device(self.video_predictor, device=device) + + def release_device(self): + if self.is_auto_mode: + if self.image_predictor: + self.image_predictor.model.to(device="cpu") + if self.video_predictor: + self.video_predictor.to(device="cpu") + + def predict(self, image, points, plabs, bbox, threshold): + if self.image_predictor is None: + self.image_predictor = SAM2ImagePredictor(build_sam2(self.config, self.modelname)) + + self.prepare_image_device() + + self.image_predictor.set_image(image) + + return sam_predict(self.image_predictor, points, plabs, bbox, threshold) + + def predict_video_segs(self, image_frames, segs): + if self.video_predictor is None: + self.video_predictor = build_sam2_video_predictor(self.config, self.modelname) + + self.prepare_video_device() + + orig_video_height = image_frames.shape[1] + orig_video_width = image_frames.shape[2] + + image_frames, padding = utils.resize_with_padding(image_frames, self.video_predictor.image_size, self.video_predictor.image_size) + image_frames = image_frames.permute(0, 3, 1, 2) + + inference_state = {} + inference_state["images"] = image_frames + inference_state["num_frames"] = len(image_frames) + inference_state["video_height"] = self.video_predictor.image_size + inference_state["video_width"] = self.video_predictor.image_size + inference_state["offload_video_to_cpu"] = True + inference_state["offload_state_to_cpu"] = self.device_mode == "CPU" + inference_state["device"] = self.video_predictor.device + + if inference_state["offload_state_to_cpu"]: + inference_state["storage_device"] = torch.device("cpu") + else: + inference_state["storage_device"] = self.video_predictor.device + + inference_state["point_inputs_per_obj"] = {} + inference_state["mask_inputs_per_obj"] = {} + inference_state["cached_features"] = {} + inference_state["constants"] = {} + + inference_state["obj_id_to_idx"] = OrderedDict() + inference_state["obj_idx_to_id"] = OrderedDict() + inference_state["obj_ids"] = [] + + inference_state["output_dict_per_obj"] = {} + inference_state["temp_output_dict_per_obj"] = {} + inference_state["frames_tracked_per_obj"] = {} + self.video_predictor._get_image_feature(inference_state, frame_idx=0, batch_size=1) + + temp_masks = {} + for i in range(0, len(segs[1])): + bbox = segs[1][i].bbox + + adjusted_bbox = utils.adjust_bbox_after_resize( + bbox, + (orig_video_height, orig_video_width), + (self.video_predictor.image_size, self.video_predictor.image_size), + padding + ) + + print(f"bbox={bbox} / adjusted_bbox={adjusted_bbox}") + + points = [utils.center_of_bbox(adjusted_bbox)] + plabs = [1] + self.video_predictor.add_new_points_or_box(inference_state=inference_state, frame_idx=0, obj_id=i, points=points, labels=plabs, box=adjusted_bbox) + temp_masks[i] = [] + + for frame_idx, object_ids, masks in self.video_predictor.propagate_in_video(inference_state): + for i in object_ids: + m = masks[i] + m = m.permute(1, 2, 0) + temp_masks[i].append(m) + + result = {} + for k, v in temp_masks.items(): + m = torch.stack(v, dim=0) + m = utils.remove_padding(m, padding) + result[k] = utils.resize_with_padding(m, orig_video_width, orig_video_height)[0] + + return result + class ESAMWrapper: def __init__(self, model, device): self.model = model @@ -646,10 +759,15 @@ class ESAMWrapper: def make_sam_mask(sam, segs, image, detection_hint, dilation, threshold, bbox_expansion, mask_hint_threshold, mask_hint_use_negative): - if not hasattr(sam, 'sam_wrapper'): + if not hasattr(sam, 'sam_wrapper') and not isinstance(sam, SAM2Wrapper): raise Exception("[Impact Pack] Invalid SAMLoader is connected. Make sure 'SAMLoader (Impact)'.\nKnown issue: The ComfyUI-YOLO node overrides the SAMLoader (Impact), making it unusable. You need to uninstall ComfyUI-YOLO.\n\n\n") - sam_obj = sam.sam_wrapper + + if isinstance(sam, SAM2Wrapper): + sam_obj = sam + else: + sam_obj = sam.sam_wrapper + sam_obj.prepare_device() try: @@ -667,7 +785,7 @@ def make_sam_mask(sam, segs, image, detection_hint, dilation, for i in range(len(segs)): bbox = segs[i].bbox - center = center_of_bbox(segs[i].bbox) + center = utils.center_of_bbox(segs[i].bbox) points.append(center) # small point is background, big point is foreground @@ -682,7 +800,7 @@ def make_sam_mask(sam, segs, image, detection_hint, dilation, else: for i in range(len(segs)): bbox = segs[i].bbox - center = center_of_bbox(bbox) + center = utils.center_of_bbox(bbox) x1 = max(bbox[0] - bbox_expansion, 0) y1 = max(bbox[1] - bbox_expansion, 0) @@ -728,7 +846,7 @@ def make_sam_mask(sam, segs, image, detection_hint, dilation, plabs = [1, 1, 1, 1] elif detection_hint == "mask-point-bbox": - center = center_of_bbox(segs[i].bbox) + center = utils.center_of_bbox(segs[i].bbox) points.append(center) plabs = [1] @@ -749,14 +867,14 @@ def make_sam_mask(sam, segs, image, detection_hint, dilation, total_masks += detected_masks # merge every collected masks - mask = combine_masks2(total_masks) + mask = utils.combine_masks2(total_masks) finally: sam_obj.release_device() if mask is not None: mask = mask.float() - mask = dilate_mask(mask.cpu().numpy(), dilation) + mask = utils.dilate_mask(mask.cpu().numpy(), dilation) mask = torch.from_numpy(mask) else: size = image.shape[0], image.shape[1] @@ -807,7 +925,7 @@ def generate_detection_hints(image, seg, center, detection_hint, dilated_bbox, m plabs = [1, 1, 1, 1] elif detection_hint == "mask-point-bbox": - center = center_of_bbox(seg.bbox) + center = utils.center_of_bbox(seg.bbox) points.append(center) plabs = [1] @@ -897,7 +1015,7 @@ def segs_scale_match(segs, target_shape): cropped_mask = cropped_mask.squeeze(0).squeeze(0).numpy() if cropped_image is not None: - cropped_image = tensor_resize(cropped_image if isinstance(cropped_image, torch.Tensor) else torch.from_numpy(cropped_image), new_w, new_h) + cropped_image = utils.tensor_resize(cropped_image if isinstance(cropped_image, torch.Tensor) else torch.from_numpy(cropped_image), new_w, new_h) cropped_image = cropped_image.numpy() new_seg = SEG(cropped_image, cropped_mask, seg.confidence, crop_region, bbox, seg.label, seg.control_net_wrapper) @@ -937,7 +1055,7 @@ def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, for i in range(len(segs)): bbox = segs[i].bbox - center = center_of_bbox(bbox) + center = utils.center_of_bbox(bbox) points.append(center) # small point is background, big point is foreground @@ -952,7 +1070,7 @@ def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, else: for i in range(len(segs)): bbox = segs[i].bbox - center = center_of_bbox(bbox) + center = utils.center_of_bbox(bbox) x1 = max(bbox[0] - bbox_expansion, 0) y1 = max(bbox[1] - bbox_expansion, 0) x2 = min(bbox[2] + bbox_expansion, image.shape[1]) @@ -969,7 +1087,7 @@ def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, total_masks += detected_masks # merge every collected masks - mask = combine_masks2(total_masks) + mask = utils.combine_masks2(total_masks) finally: sam_obj.release_device() @@ -978,7 +1096,7 @@ def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, if mask is not None: mask = mask.float() - mask = dilate_mask(mask.cpu().numpy(), dilation) + mask = utils.dilate_mask(mask.cpu().numpy(), dilation) mask = torch.from_numpy(mask) mask = mask.to(device=mask_working_device) else: @@ -995,7 +1113,7 @@ def make_sam_mask_segmented(sam, segs, image, detection_hint, dilation, def segs_bitwise_and_mask(segs, mask): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) if mask is None: print("[SegsBitwiseAndMask] Cannot operate: MASK is empty.") @@ -1021,7 +1139,7 @@ def segs_bitwise_and_mask(segs, mask): def segs_bitwise_subtract_mask(segs, mask): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) if mask is None: print("[SegsBitwiseSubtractMask] Cannot operate: MASK is empty.") @@ -1077,7 +1195,7 @@ def dilate_segs(segs, factor): new_segs = [] for seg in segs[1]: - new_mask = dilate_mask(seg.cropped_mask, factor) + new_mask = utils.dilate_mask(seg.cropped_mask, factor) new_seg = SEG(seg.cropped_image, new_mask, seg.confidence, seg.crop_region, seg.bbox, seg.label, seg.control_net_wrapper) new_segs.append(new_seg) @@ -1109,7 +1227,7 @@ class ONNXDetector: x1, y1, x2, y2 = item_bbox if x2 - x1 > drop_size and y2 - y1 > drop_size: # minimum dimension must be (2,2) to avoid squeeze issue - crop_region = make_crop_region(w, h, item_bbox, crop_factor) + crop_region = utils.make_crop_region(w, h, item_bbox, crop_factor) if detailer_hook is not None: crop_region = item_bbox.post_crop_region(w, h, item_bbox, crop_region) @@ -1119,7 +1237,7 @@ class ONNXDetector: # prepare cropped mask cropped_mask = np.zeros((crop_y2 - crop_y1, crop_x2 - crop_x1)) cropped_mask[y1 - crop_y1:y2 - crop_y1, x1 - crop_x1:x2 - crop_x1] = 1 - cropped_mask = dilate_mask(cropped_mask, dilation) + cropped_mask = utils.dilate_mask(cropped_mask, dilation) # make items. just convert the integer label to a string item = SEG(None, cropped_mask, scores[i], crop_region, item_bbox, str(labels[i]), None) @@ -1194,7 +1312,7 @@ def mask_to_segs(mask, combined, crop_factor, bbox_fill, drop_size=1, label='A', np.max(indices[1]), np.max(indices[0]), ) - crop_region = make_crop_region( + crop_region = utils.make_crop_region( mask_i.shape[1], mask_i.shape[0], bbox, crop_factor ) x1, y1, x2, y2 = crop_region @@ -1228,7 +1346,7 @@ def mask_to_segs(mask, combined, crop_factor, bbox_fill, drop_size=1, label='A', x, y, w, h = cv2.boundingRect(contour) bbox = x, y, x + w, y + h - crop_region = make_crop_region( + crop_region = utils.make_crop_region( mask_i.shape[1], mask_i.shape[0], bbox, crop_factor, crop_min_size ) @@ -1302,7 +1420,7 @@ def mediapipe_facemesh_to_segs(image, crop_factor, bbox_fill, crop_min_size, dro tensor = torch.from_numpy(convex_segment) mask_tensor = torch.any(tensor != 0, dim=-1).float() mask_tensor = mask_tensor.squeeze(0) - mask_tensor = torch.from_numpy(dilate_mask(mask_tensor.numpy(), dilation)) + mask_tensor = torch.from_numpy(utils.dilate_mask(mask_tensor.numpy(), dilation)) mask_list.append(mask_tensor.unsqueeze(0)) return mask_list @@ -1537,7 +1655,7 @@ class TwoSamplersForMaskUpscaler: hook_full_opt=None, tile_size=512): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) mask = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])) @@ -1555,7 +1673,7 @@ class TwoSamplersForMaskUpscaler: def upscale(self, step_info, samples, upscale_factor, save_temp_prefix=None): scale_method, sample_schedule, use_tiled_vae, base_sampler, mask_sampler, mask, vae = self.params - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) self.prepare_hook(step_info) @@ -1585,7 +1703,7 @@ class TwoSamplersForMaskUpscaler: def upscale_shape(self, step_info, samples, w, h, save_temp_prefix=None): scale_method, sample_schedule, use_tiled_vae, base_sampler, mask_sampler, mask, vae = self.params - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) self.prepare_hook(step_info) @@ -1641,7 +1759,7 @@ class TwoSamplersForMaskUpscaler: return cur_step % 2 == 0 or cur_step >= total_step - 1 def do_samples(self, step_info, base_sampler, mask_sampler, sample_schedule, mask, upscaled_latent): - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) if self.is_full_sample_time(step_info, sample_schedule): print(f"step_info={step_info} / full time") @@ -2085,7 +2203,7 @@ class BBoxDetectorBasedOnCLIPSeg: def detect(self, image, bbox_threshold, bbox_dilation, bbox_crop_factor, drop_size=1, detailer_hook=None): mask = self.detect_combined(image, bbox_threshold, bbox_dilation) - mask = make_2d_mask(mask) + mask = utils.make_2d_mask(mask) segs = mask_to_segs(mask, False, bbox_crop_factor, True, drop_size, detailer_hook=detailer_hook) @@ -2115,7 +2233,7 @@ class BBoxDetectorBasedOnCLIPSeg: prompt = self.aux if self.prompt == '' and self.aux is not None else self.prompt mask, _, _ = CLIPSeg().segment_image(image, prompt, self.blur, threshold, dilation_factor) - mask = to_binary_mask(mask) + mask = utils.to_binary_mask(mask) return mask def setAux(self, x): diff --git a/modules/impact/detectors.py b/modules/impact/detectors.py index d2804de..4fc49d9 100644 --- a/modules/impact/detectors.py +++ b/modules/impact/detectors.py @@ -1,3 +1,5 @@ +import logging + import impact.core as core from nodes import MAX_RESOLUTION import impact.segs_nodes as segs_nodes @@ -298,6 +300,71 @@ class SimpleDetectorForEachPipe: sam_mask_hint_threshold, post_dilation=post_dilation, sam_model_opt=sam_model_opt, segm_detector_opt=segm_detector_opt, detailer_hook=detailer_hook) +class SAM2VideoDetectorSEGS: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image_frames": ("IMAGE", ), + + "bbox_detector": ("BBOX_DETECTOR", ), + "sam2_model": ("SAM_MODEL", ), + + "bbox_threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "sam2_threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + + "crop_factor": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 100, "step": 0.1}), + "drop_size": ("INT", {"min": 1, "max": MAX_RESOLUTION, "step": 1, "default": 10}), + } + } + + RETURN_TYPES = ("SEGS", ) + FUNCTION = "doit" + + CATEGORY = "ImpactPack/Detector" + + @staticmethod + def doit(bbox_detector, sam2_model, image_frames, bbox_threshold, sam2_threshold, crop_factor, drop_size): + if not isinstance(sam2_model, core.SAM2Wrapper): + logging.error("[Impact Pack] To use the SAM2VideoDetectorSEGS node, a SAM2 model must be provided as input to `sam2_model`.") + raise Exception("To use the SAM2VideoDetectorSEGS node, a SAM2 model must be provided as input to `sam2_model`.") + + segs = bbox_detector.detect(image_frames[0].unsqueeze(0), bbox_threshold, 0, 0, drop_size) + segs_masks = sam2_model.predict_video_segs(image_frames, segs) + + def get_whole_merged_mask(all_masks): + merged_mask = (all_masks[0] * 255).to(torch.uint8) + for mask in all_masks[1:]: + merged_mask |= (mask * 255).to(torch.uint8) + + merged_mask = (merged_mask / 255.0).to(torch.float32) + merged_mask = utils.to_binary_mask(merged_mask, 0.1)[0] + return merged_mask + + test_mask1 = None + test_mask2 = None + new_segs = [] + for k, v in segs_masks.items(): + v = v.squeeze(3) + m = get_whole_merged_mask(v) + test_mask2 = v + seg = segs_nodes.MaskToSEGS.doit(m, False, crop_factor, False, drop_size, contour_fill=True)[0][1] + + if len(seg) == 0: + continue + + seg = seg[0] + + x1, y1, x2, y2 = seg.crop_region + masks = [] + for mask in v: + masks.append(mask[y1:y2, x1:x2]) + cropped_mask = torch.stack(masks) + cropped_mask = (cropped_mask >= (sam2_threshold*100-50)).to(torch.uint8).cpu() + new_seg = SEG(seg.cropped_image, cropped_mask, seg.confidence, seg.crop_region, seg.bbox, seg.label, seg.control_net_wrapper) + new_segs.append(new_seg) + + return ((segs[0], new_segs), ) + class SimpleDetectorForAnimateDiff: @classmethod diff --git a/modules/impact/impact_pack.py b/modules/impact/impact_pack.py index e70d15e..4b3e0ce 100644 --- a/modules/impact/impact_pack.py +++ b/modules/impact/impact_pack.py @@ -92,10 +92,21 @@ class CLIPSegDetectorProvider: print("[ERROR] CLIPSegToBboxDetector: CLIPSeg custom node isn't installed. You must install biegert/ComfyUI-CLIPSeg extension to use this node.") +sam2_config_table = { + 'sam2.1_hiera_base_plus.pt': 'configs/sam2.1/sam2.1_hiera_b+.yaml', + 'sam2.1_hiera_large.pt': 'configs/sam2.1/sam2.1_hiera_l.yaml', + 'sam2.1_hiera_small.pt': 'configs/sam2.1/sam2.1_hiera_s.yaml', + 'sam2.1_hiera_tiny.pt': 'configs/sam2.1/sam2.1_hiera_t.yaml', + 'sam2_hiera_tiny.pt': 'configs/sam2/sam2_hiera_t.yaml', + 'sam2_hiera_small.pt': 'configs/sam2/sam2_hiera_s.yaml', + 'sam2_hiera_base_plus.pt': 'configs/sam2/sam2_hiera_b+.yaml', + 'sam2_hiera_large.pt': 'configs/sam2/sam2_hiera_l.yaml' +} + class SAMLoader: @classmethod def INPUT_TYPES(cls): - models = [x for x in folder_paths.get_filename_list("sams") if 'hq' not in x] + models = [x for x in folder_paths.get_filename_list("sams") if 'hq' not in x and (x.endswith('.pt') or x.endswith('.pth') or x.endswith('.safetensors'))] if 'ESAM_ModelLoader_Zho' in nodes.NODE_CLASS_MAPPINGS: models.append('ESAM') @@ -136,17 +147,22 @@ class SAMLoader: print(f"Loads EfficientSAM model: (device:{device_mode})") return (esam, ) - - modelname = folder_paths.get_full_path("sams", model_name) - - if 'vit_h' in model_name: - model_kind = 'vit_h' - elif 'vit_l' in model_name: - model_kind = 'vit_l' + elif model_name in sam2_config_table: + model_kind = 'sam2' + config = sam2_config_table[model_name] + modelname = folder_paths.get_full_path("sams", model_name) else: - model_kind = 'vit_b' + modelname = folder_paths.get_full_path("sams", model_name) + + if 'vit_h' in model_name: + model_kind = 'vit_h' + elif 'vit_l' in model_name: + model_kind = 'vit_l' + else: + model_kind = 'vit_b' + + sam = sam_model_registry[model_kind](checkpoint=modelname) - sam = sam_model_registry[model_kind](checkpoint=modelname) size = os.path.getsize(modelname) safe_to = core.SafeToGPU(size) @@ -158,10 +174,14 @@ 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 + if model_kind == 'sam2': + sam = core.SAM2Wrapper(config=config, modelname=modelname, is_auto_mode=is_auto_mode, safe_to_gpu=safe_to, device_mode=device_mode) + print(f"Loads SAM2 model: {modelname} (device:{device_mode})") + else: + 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})") - print(f"Loads SAM model: {modelname} (device:{device_mode})") return (sam, ) diff --git a/modules/impact/utils.py b/modules/impact/utils.py index 867f09c..dc25098 100644 --- a/modules/impact/utils.py +++ b/modules/impact/utils.py @@ -67,6 +67,54 @@ def tensor_convert_rgb(image, prefer_copy=True): raise ValueError(f"illegal conversion (channels: {n_channel} -> 3)") +def resize_with_padding(image, target_w: int, target_h: int): + _tensor_check_image(image) + b, h, w, c = image.shape + image = image.permute(0, 3, 1, 2) # B, C, H, W + + scale = min(target_w / w, target_h / h) + new_w, new_h = int(w * scale), int(h * scale) + + image = F.interpolate(image, size=(new_h, new_w), mode="bilinear", align_corners=False) + + pad_left = (target_w - new_w) // 2 + pad_right = target_w - new_w - pad_left + pad_top = (target_h - new_h) // 2 + pad_bottom = target_h - new_h - pad_top + + image = F.pad(image, (pad_left, pad_right, pad_top, pad_bottom), mode='constant', value=0) + + image = image.permute(0, 2, 3, 1) # B, H, W, C + return image, (pad_top, pad_bottom, pad_left, pad_right) + + +def remove_padding(image, padding): + pad_top, pad_bottom, pad_left, pad_right = padding + return image[:, pad_top:image.shape[1] - pad_bottom, pad_left:image.shape[2] - pad_right, :] + + +def adjust_bbox_after_resize(bbox, original_size, target_size, padding): + """ + bbox: (x1, y1, x2, y2) in original image + original_size: (original_h, original_w) + target_size: (target_h, target_w) + padding: (pad_top, pad_bottom, pad_left, pad_right) + """ + orig_h, orig_w = original_size + target_h, target_w = target_size + pad_top, pad_bottom, pad_left, pad_right = padding + + scale = min(target_w / orig_w, target_h / orig_h) + + # Apply scale + x1 = int(bbox[0] * scale + pad_left) + y1 = int(bbox[1] * scale + pad_top) + x2 = int(bbox[2] * scale + pad_left) + y2 = int(bbox[3] * scale + pad_top) + + return x1, y1, x2, y2 + + def general_tensor_resize(image, w: int, h: int): _tensor_check_image(image) image = image.permute(0, 3, 1, 2) diff --git a/pyproject.toml b/pyproject.toml index 1fdd611..0cc5aee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-impact-pack" description = "This node pack offers various detector nodes and detailer nodes that allow you to configure a workflow that automatically enhances facial details. And provide iterative upscaler." -version = "8.17.1" +version = "8.18" license = { file = "LICENSE.txt" } dependencies = ["segment-anything", "scikit-image", "piexif", "transformers", "opencv-python-headless", "GitPython", "scipy>=1.11.4"] diff --git a/requirements.txt b/requirements.txt index c7d9e0e..246a751 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,4 +6,5 @@ opencv-python-headless scipy>=1.11.4 numpy dill -matplotlib \ No newline at end of file +matplotlib +git+https://github.com/facebookresearch/sam2 \ No newline at end of file