Support SAM2 models.

https://github.com/ltdrdata/ComfyUI-Impact-Pack/issues/684
This commit is contained in:
Dr.Lt.Data
2025-07-07 00:34:10 +09:00
parent 705698faf2
commit d444e0acad
9 changed files with 321 additions and 61 deletions
+6 -2
View File
@@ -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
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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):
+67
View File
@@ -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
+33 -13
View File
@@ -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, )
+48
View File
@@ -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
View File
@@ -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
View File
@@ -6,4 +6,5 @@ opencv-python-headless
scipy>=1.11.4
numpy
dill
matplotlib
matplotlib
git+https://github.com/facebookresearch/sam2