SAM2Ultra and ObjectDetector nodes support image batch

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