Support SAM2 models.
https://github.com/ltdrdata/ComfyUI-Impact-Pack/issues/684
This commit is contained in:
@@ -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.
|
||||
|
||||
+6
-4
@@ -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)",
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
+157
-39
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, )
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
+1
-1
@@ -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"]
|
||||
|
||||
|
||||
+2
-1
@@ -6,4 +6,5 @@ opencv-python-headless
|
||||
scipy>=1.11.4
|
||||
numpy
|
||||
dill
|
||||
matplotlib
|
||||
matplotlib
|
||||
git+https://github.com/facebookresearch/sam2
|
||||
Reference in New Issue
Block a user