diff --git a/README.md b/README.md index e4c7ca6..9251878 100644 --- a/README.md +++ b/README.md @@ -1,2 +1,361 @@ -# ComfyUI-Loop-image -A comfyui node that uses the loop function to process images and masks +# ComfyUI Loop Image + +[English](#english) | [简体中文](#简体中文) + + +# ComfyUI Loop Image + +## Introduction +ComfyUI Loop Image is a node package specifically designed for image loop processing. It provides two main processing modes: Batch Image Processing and Single Image Processing, along with supporting image segmentation and merging functions. + +## Differences between Batch and Single Processing + +### Batch Image Processing +- Suitable for scenarios requiring simultaneous processing of multiple different regions +- Uses Mask Segmentation node to divide images into multiple parts +- Processes one segmented region per iteration +- Merges results through Mask Merge after all regions are processed + +### Single Image Processing +- Suitable for scenarios requiring multiple processing passes on the same image +- Uses the result of the previous iteration as input for the next +- Enables progressive image modification +- Ideal for iterative optimization tasks + +## Node Documentation + +### 1. Batch Processing Nodes + +#### Mask Segmentation🐰 +- **Functionality** + - Automatically segments a mask containing multiple independent regions into separate mask sequences + - Each segmented mask corresponds to an independent region in the original image + - Segmentation based on connected component analysis + +- **Segmentation Rules** + - Independent regions are identified as separate parts + - Regions with holes are properly processed, maintaining hole structure + +- **Sequence Rules** + - Masks are arranged from left to right, then top to bottom + - Sorting based on leftmost pixel position, then topmost pixel position + - This order determines subsequent processing sequence + - Example: In a mask with three regions, leftmost region is iteration 0, middle is 1, rightmost is 2 + +#### Batch Image Loop Open🐰 +- **Input/Output Details** + - Inputs: + - segmented_images: Image sequence from Mask Segmentation + - segmented_masks: Mask sequence from Mask Segmentation + - Outputs: + - current_image: Currently processed image portion + - current_mask: Current iteration mask + - max_iterations: Total iteration count (equals number of segmented regions) + - iteration_count: Current iteration number (starts from 0) + +- **Usage Notes** + - current_image and current_mask can be used directly for subsequent processing + - iteration_count can connect to Loop Index Switch for different processing parameters + - max_iterations used for loop control, usually doesn't need manual handling + +#### Batch Image Loop Close🐰 +- **Input/Output Details** + - Inputs: + - flow_control: Control signal from Loop Open + - current_image: Currently processed image + - current_mask: Current processed mask + - max_iterations: Total iteration count from Loop Open + - Outputs: + - result_images: All processed image sequences + - result_masks: All processed mask sequences + +#### Mask Merge🐰 +- **Functionality** + - Merges multiple processed image regions back into the original image + - Uses masks to ensure each processed region is correctly placed + - Maintains original content in unprocessed areas + +- **Usage Tips** + - original_image: Use original input image + - processed_images: Connect to result_images output from Loop Close + - masks: Connect to result_masks output from Loop Close + +This batch processing system allows you to apply different processing methods to different regions of an image, particularly suitable for scenarios requiring differentiated processing of various image parts. + +### 2. Single Image Processing Nodes + +#### Single Image Loop Open🐰 +- **Functionality** + - Performs multiple iterations of processing on a single image + - Uses the result of each iteration as input for the next + - Suitable for progressive enhancement or multiple optimization scenarios + +- **Input Parameters** + - **Required Inputs**: + - image: Original image to process + - max_iterations: Maximum iteration count (1-100) + - **Optional Inputs**: + - mask: Optional processing area mask + +- **Output Parameters** + - current_image: Current iteration image (original image for first iteration, previous result for subsequent iterations) + - current_mask: Current mask (if provided) + - max_iterations: Set maximum iterations + - iteration_count: Current iteration number (starts from 0) + +#### Single Image Loop Close🐰 +- **Input Parameters** + - **Required Inputs**: + - flow_control: Control signal from Loop Open + - current_image: Currently processed image + - max_iterations: Maximum iterations from Loop Open + - **Optional Inputs**: + - current_mask: Processed mask (if using mask) + +- **Output Parameters** + - final_image: Final image after all iterations + - final_mask: Final mask (if using mask) + +#### Single Image Processing Features and Applications +1. **Progressive Processing** + - Each iteration builds on previous results + - Enables cumulative effects + - Suitable for scenarios requiring fine-tuning + +2. **Use Case Examples** + - Progressive image enhancement + - Iterative style transfer + - Multiple denoising passes + - Gradual detail optimization + +### 3. Special Function Node +- **Loop Index Switch🐰** + - Function: Select different inputs based on current iteration count + - Usage: + 1. Right-click node and select "Add Loop Input" + 2. Enter desired iteration number (0-99) + 3. Connect corresponding inputs + 4. Use "Remove Loop Input" to delete unwanted inputs + - Note: Only inputs corresponding to current iteration are computed, others are skipped for efficiency + +## Usage Recommendations +1. Use batch processing for scenarios requiring different processing in different image regions +2. Use single image processing for scenarios requiring multiple optimization iterations +3. Utilize Loop Index Switch to implement different parameters for different iterations +4. Control iteration count to avoid over-processing + +## Example Workflows +TODO + +## Acknowledgments +This project references the following excellent open source projects: +- [ComfyUI-Easy-Use](https://github.com/yolain/ComfyUI-Easy-Use/) - Provided excellent node design ideas and implementation references +- [execution-inversion-demo-comfyui](https://github.com/BadCafeCode/execution-inversion-demo-comfyui) - Provided core implementation ideas for loop control +- [cozy_ex_dynamic](https://github.com/cozy-comfyui/cozy_ex_dynamic) - Provided implementation reference for dynamic input nodes + +Special thanks to the authors of these projects for their contributions to the ComfyUI community! + +## About +For more ComfyUI tutorials and updates, visit: +- Bilibili: [CyberEve](https://space.bilibili.com/16993154) +- Content includes: + - ComfyUI node development tutorials + - Workflow usage tutorials + - Latest feature updates + - AI drawing tips + +If you find this project helpful, please follow the author's Bilibili account for more resources! + +--- + + + + +## 简介 +ComfyUI Loop Image是一个专门用于处理图像循环操作的节点包。它提供了两种主要的循环处理模式:批量图像处理(Batch)和单图像重复处理(Single),以及配套的图像分割与合并功能。 + + +## 批量处理与单图处理的区别 + +### 批量图像处理 +- 适用于需要同时处理多个不同区域的场景 +- 通过Mask Segmentation节点将图像分割成多个部分 +- 每次循环处理一个分割区域 +- 所有区域处理完成后通过Mask Merge合并结果 + +### 单图像处理 +- 适用于需要对同一图像进行多次处理的场景 +- 每次循环使用上一次的处理结果作为输入 +- 可以实现渐进式的图像修改 +- 适合迭代优化类的任务 + + +## 节点说明 + + +### 1. 批量处理节点详解 + + +#### Mask Segmentation🐰 (遮罩分割) +- **功能说明** + - 将一个包含多个独立区域的遮罩图自动分割成独立的遮罩序列 + - 每个分割后的遮罩对应原图中的一个独立区域 + - 分割基于连通区域分析,即相互不连接的区域会被分为不同部分 + +- **分割规则** + - 相互独立的区域会被识别为不同的部分 + - 包含孔洞的区域会被正确处理,保持孔洞结构 + +- **顺序规则** + - 分割后的遮罩按照从左到右排列,若左右位置相等,再按照从上到下的顺序 + - 排序依据是每个区域最左边的像素点的位置,再按照最上边的像素点的位置 + - 这个顺序决定了后续循环处理的顺序 + - 例如:如果遮罩中有三个区域,最左边的区域将是第0次迭代,中间的是第1次,最右边的是第2次 + + +#### Batch Image Loop Open🐰 (批量循环开始) +- **输入输出详解** + - 输入: + - segmented_images: 来自Mask Segmentation的图像序列 + - segmented_masks: 来自Mask Segmentation的遮罩序列 + - 输出: + - current_image: 当前迭代处理的图像部分 + - current_mask: 当前迭代的遮罩 + - max_iterations: 总迭代次数(等于分割区域的数量) + - iteration_count: 当前迭代次数(从0开始) + +- **使用说明** + - current_image和current_mask可以直接用于后续处理 + - iteration_count可以连接到Loop Index Switch来选择不同的处理参数 + - max_iterations用于循环控制,一般不需要手动使用 + + +#### Batch Image Loop Close🐰 (批量循环结束) +- **输入输出详解** + - 输入: + - flow_control: 来自Loop Open的控制信号 + - current_image: 处理后的当前图像 + - current_mask: 处理后的当前遮罩 + - max_iterations: 来自Loop Open的总迭代次数 + - 输出: + - result_images: 所有处理完成的图像序列 + - result_masks: 所有处理完成的遮罩序列 + + +#### Mask Merge🐰 (遮罩合并) +- **功能说明** + - 将循环处理后的多个图像区域合并回原始图像 + - 使用遮罩确保每个处理过的区域正确放回原位 + - 保持未处理区域的原始内容不变 + +- **使用技巧** + - original_image: 使用原始输入图像 + - processed_images: 连接Loop Close的result_images输出 + - masks: 连接Loop Close的result_masks输出 + +这样的批量处理系统允许你对图像的不同区域应用不同的处理方法,特别适合需要对图像不同部分进行差异化处理的场景。 + + +### 2. 单图处理节点详解 + +#### Single Image Loop Open🐰 (单图循环开始) +- **功能说明** + - 对同一张图像进行多次迭代处理 + - 每次迭代都使用上一次的处理结果作为输入 + - 适合需要渐进式改善或多次优化的场景 + +- **输入参数详解** + - **必需输入**: + - image: 需要处理的原始图像 + - max_iterations: 最大迭代次数(1-100) + - **可选输入**: + - mask: 可选的处理区域遮罩 + +- **输出参数详解** + - current_image: 当前迭代的图像(第一次是原始图像,之后是上一次处理的结果) + - current_mask: 当前使用的遮罩(如果提供了遮罩) + - max_iterations: 设定的最大迭代次数 + - iteration_count: 当前迭代次数(从0开始) + + +#### Single Image Loop Close🐰 (单图循环结束) +- **输入参数详解** + - **必需输入**: + - flow_control: 来自Loop Open的控制信号 + - current_image: 当前迭代处理后的图像 + - max_iterations: 来自Loop Open的最大迭代次数 + - **可选输入**: + - current_mask: 处理后的遮罩(如果使用了遮罩) + +- **输出参数详解** + - final_image: 所有迭代完成后的最终图像 + - final_mask: 最终的遮罩(如果使用了遮罩) + + +#### 单图处理的特点和应用场景 +1. **渐进式处理** + - 每次迭代都基于上一次的结果 + - 可以实现累积效果 + - 适合需要多次微调的场景 + +2. **使用场景示例** + - 图像渐进式增强 + - 迭代式风格转换 + - 多次降噪处理 + - 逐步细节优化 + + +### 与Loop Index Switch的配合使用 +- 可以使用Loop Index Switch根据iteration_count选择不同的处理参数 + +这种单图循环处理方式特别适合需要精细调整或渐进式改善的场景,通过多次迭代可以达到更理想的处理效果。配合Loop Index Switch,还可以实现更复杂的参数控制策略。 + + +### 3. 特殊功能节点 +- **Loop Index Switch🐰** + - 功能:根据当前循环次数选择不同的输入 + - 使用方法: + 1. 右键点击节点选择"Add Loop Input" + 2. 输入想要添加的循环序号(0-99) + 3. 连接对应的输入 + 4. 可以通过"Remove Loop Input"删除不需要的输入 + - 注意:只有当前迭代次数对应的输入会被计算,其他输入会被跳过,提高效率 + + +## 使用建议 +1. 批量处理适合需要在图像不同区域应用不同处理的场景 +2. 单图处理适合需要多次迭代优化的场景 +3. 合理使用Loop Index Switch节点可以实现在不同迭代次数使用不同参数 +4. 注意控制循环次数,避免过度处理 + + +## 示例工作流 +TODO + + +## 致谢 + +本项目在开发过程中参考和借鉴了以下优秀的开源项目: + +- [ComfyUI-Easy-Use](https://github.com/yolain/ComfyUI-Easy-Use/) - 提供了优秀的节点设计思路和实现参考 +- [execution-inversion-demo-comfyui](https://github.com/BadCafeCode/execution-inversion-demo-comfyui) - 提供了循环控制的核心实现思路 +- [cozy_ex_dynamic](https://github.com/cozy-comfyui/cozy_ex_dynamic) - 提供了动态输入节点的实现参考 + +特别感谢这些项目的作者们为ComfyUI社区做出的贡献! + + +## 关于作者 + +欢迎访问作者的B站主页,获取更多ComfyUI教程和更新: +- B站:[CyberEve](https://space.bilibili.com/16993154) +- 内容包括: + - ComfyUI节点开发教程 + - 工作流使用教程 + - 最新功能更新介绍 + - AI绘画技巧分享 + +如果您觉得这个项目对您有帮助,欢迎关注作者B站账号获取更多资源! + +--- + +*Note: 本项目遵循开源协议,欢迎提出建议和改进意见。* \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..9949b67 --- /dev/null +++ b/__init__.py @@ -0,0 +1,13 @@ +from .flow_control import CyberEve_Loop_CLASS_MAPPINGS, CyberEve_Loop_DISPLAY_NAME_MAPPINGS +from .mask_split import Mask_CLASS_MAPPINGS, Mask_DISPLAY_NAME_MAPPINGS + +WEB_DIRECTORY = "./web" +NODE_CLASS_MAPPINGS = {} +NODE_CLASS_MAPPINGS.update(CyberEve_Loop_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(Mask_CLASS_MAPPINGS) + +NODE_DISPLAY_NAME_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS.update(CyberEve_Loop_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(Mask_DISPLAY_NAME_MAPPINGS) + + diff --git a/flow_control.py b/flow_control.py new file mode 100644 index 0000000..1985ac7 --- /dev/null +++ b/flow_control.py @@ -0,0 +1,485 @@ +from comfy_execution.graph_utils import GraphBuilder, is_link +from .tools import VariantSupport +import torch +from nodes import NODE_CLASS_MAPPINGS as ALL_NODE_CLASS_MAPPINGS + +@VariantSupport() +class BatchImageLoopOpen: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + inputs = { + "required": { + "segmented_images": ("IMAGE", {"forceInput": True}), + "segmented_masks": ("MASK", {"forceInput": True}), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + "iteration_count": ("INT", {"default": 0}), + } + } + return inputs + + RETURN_TYPES = tuple(["FLOW_CONTROL", "IMAGE", "MASK", "INT", "INT"]) + RETURN_NAMES = tuple(["FLOW_CONTROL", "current_image", "current_mask", "max_iterations", "iteration_count"]) + FUNCTION = "while_loop_open" + CATEGORY = "CyberEveLoop🐰" + + def while_loop_open(self, segmented_images, segmented_masks, unique_id=None, iteration_count=0): + print(f"while_loop_open Processing iteration {iteration_count}") + + # 确保输入是张量 + if isinstance(segmented_images, list): + segmented_images = torch.cat(segmented_images, dim=0) + if isinstance(segmented_masks, list): + segmented_masks = torch.cat(segmented_masks, dim=0) + + max_iterations = segmented_images.shape[0] + if max_iterations == 0: + raise ValueError("No images provided in segmented_images") + + # 获取当前迭代的图片和蒙版 + current_image = segmented_images[iteration_count:iteration_count+1] + current_mask = segmented_masks[iteration_count:iteration_count+1] + + return tuple(["stub", current_image, current_mask, max_iterations, iteration_count]) + +@VariantSupport() +class BatchImageLoopClose: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + inputs = { + "required": { + "flow_control": ("FLOW_CONTROL", {"rawLink": True}), + "current_image": ("IMAGE",), + "current_mask": ("MASK",), + "max_iterations": ("INT", {"forceInput": True}), + }, + "hidden": { + "dynprompt": "DYNPROMPT", + "unique_id": "UNIQUE_ID", + "result_images": ("IMAGE",), + "result_masks": ("MASK",), + "iteration_count": ("INT", {"default": 0}), + } + } + return inputs + + RETURN_TYPES = tuple(["IMAGE", "MASK"]) + RETURN_NAMES = tuple(["result_images", "result_masks"]) + FUNCTION = "while_loop_close" + CATEGORY = "CyberEveLoop🐰" + + def explore_dependencies(self, node_id, dynprompt, upstream, parent_ids): + node_info = dynprompt.get_node(node_id) + if "inputs" not in node_info: + return + + for k, v in node_info["inputs"].items(): + if is_link(v): + parent_id = v[0] + display_id = dynprompt.get_display_node_id(parent_id) + display_node = dynprompt.get_node(display_id) + class_type = display_node["class_type"] + # 排除循环结束节点 + if class_type not in ['BatchImageLoopClose']: + parent_ids.append(display_id) + if parent_id not in upstream: + upstream[parent_id] = [] + self.explore_dependencies(parent_id, dynprompt, upstream, parent_ids) + upstream[parent_id].append(node_id) + + def explore_output_nodes(self, dynprompt, upstream, output_nodes, parent_ids): + """探索并添加输出节点的连接""" + for parent_id in upstream: + display_id = dynprompt.get_display_node_id(parent_id) + for output_id in output_nodes: + id = output_nodes[output_id][0] + if id in parent_ids and display_id == id and output_id not in upstream[parent_id]: + if '.' in parent_id: + arr = parent_id.split('.') + arr[len(arr)-1] = output_id + upstream[parent_id].append('.'.join(arr)) + else: + upstream[parent_id].append(output_id) + + def collect_contained(self, node_id, upstream, contained): + if node_id not in upstream: + return + for child_id in upstream[node_id]: + if child_id not in contained: + contained[child_id] = True + self.collect_contained(child_id, upstream, contained) + + def while_loop_close(self, flow_control, current_image, current_mask, max_iterations, + iteration_count=0, result_images=None, result_masks=None, + dynprompt=None, unique_id=None,): + print(f"Iteration {iteration_count} of {max_iterations}") + + # 维度处理 + if len(current_image.shape) == 3: + current_image = current_image.unsqueeze(0) + if len(current_mask.shape) == 2: + current_mask = current_mask.unsqueeze(0) + + # 结果初始化 + if result_images is None: + result_images = torch.zeros((max_iterations,) + current_image.shape[1:], + dtype=current_image.dtype, + device=current_image.device) + result_masks = torch.zeros((max_iterations,) + current_mask.shape[1:], + dtype=current_mask.dtype, + device=current_mask.device) + + # 存储当前结果 + result_images[iteration_count:iteration_count+1] = current_image + result_masks[iteration_count:iteration_count+1] = current_mask + + # 检查是否继续循环 + if iteration_count >= max_iterations - 1: + print(f"Loop finished with {iteration_count + 1} iterations") + return (result_images, result_masks) + + # 准备下一次循环 + this_node = dynprompt.get_node(unique_id) + upstream = {} + parent_ids = [] + self.explore_dependencies(unique_id, dynprompt, upstream, parent_ids) + parent_ids = list(set(parent_ids)) # 去重 + + # 获取并处理输出节点 + prompts = dynprompt.get_original_prompt() + output_nodes = {} + for id in prompts: + node = prompts[id] + if "inputs" not in node: + continue + class_type = node["class_type"] + if class_type in ALL_NODE_CLASS_MAPPINGS: + class_def = ALL_NODE_CLASS_MAPPINGS[class_type] + if hasattr(class_def, 'OUTPUT_NODE') and class_def.OUTPUT_NODE == True: + for k, v in node['inputs'].items(): + if is_link(v): + output_nodes[id] = v + + # 创建新图 + graph = GraphBuilder() + self.explore_output_nodes(dynprompt, upstream, output_nodes, parent_ids) + + contained = {} + open_node = flow_control[0] + self.collect_contained(open_node, upstream, contained) + contained[unique_id] = True + contained[open_node] = True + + # 创建节点 + for node_id in contained: + original_node = dynprompt.get_node(node_id) + node = graph.node(original_node["class_type"], + "Recurse" if node_id == unique_id else node_id) + node.set_override_display_id(node_id) + + # 设置连接 + for node_id in contained: + original_node = dynprompt.get_node(node_id) + node = graph.lookup_node("Recurse" if node_id == unique_id else node_id) + for k, v in original_node["inputs"].items(): + if is_link(v) and v[0] in contained: + parent = graph.lookup_node(v[0]) + node.set_input(k, parent.out(v[1])) + else: + node.set_input(k, v) + + # 设置节点参数 + my_clone = graph.lookup_node("Recurse") + my_clone.set_input("iteration_count", iteration_count + 1) + my_clone.set_input("result_images", result_images) + my_clone.set_input("result_masks", result_masks) + + new_open = graph.lookup_node(open_node) + new_open.set_input("iteration_count", iteration_count + 1) + + print(f"Continuing to iteration {iteration_count + 1}") + + return { + "result": tuple([my_clone.out(0), my_clone.out(1)]), + "expand": graph.finalize(), + } + + +@VariantSupport() +class SingleImageLoopOpen: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + inputs = { + "required": { + "image": ("IMAGE",), + "max_iterations": ("INT", {"default": 5, "min": 1, "max": 100}), + }, + "optional": { + "mask": ("MASK",), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + "iteration_count": ("INT", {"default": 0}), + "previous_image": ("IMAGE",), + "previous_mask": ("MASK",), + } + } + return inputs + + RETURN_TYPES = tuple(["FLOW_CONTROL", "IMAGE", "MASK", "INT", "INT"]) + RETURN_NAMES = tuple(["FLOW_CONTROL", "current_image", "current_mask", "max_iterations", "iteration_count"]) + FUNCTION = "loop_open" + CATEGORY = "CyberEveLoop🐰" + + def loop_open(self, image, max_iterations, mask=None, unique_id=None, + iteration_count=0, previous_image=None, previous_mask=None): + print(f"SingleImageLoopOpen Processing iteration {iteration_count}") + + # 确保维度正确 + if len(image.shape) == 3: + image = image.unsqueeze(0) + if mask is not None and len(mask.shape) == 2: + mask = mask.unsqueeze(0) + + # 使用上一次循环的结果(如果有) + current_image = previous_image if previous_image is not None and iteration_count > 0 else image + current_mask = previous_mask if previous_mask is not None and iteration_count > 0 else mask + + return tuple(["stub", current_image, current_mask, max_iterations, iteration_count]) + +@VariantSupport() +class SingleImageLoopClose: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + inputs = { + "required": { + "flow_control": ("FLOW_CONTROL", {"rawLink": True}), + "current_image": ("IMAGE",), + "max_iterations": ("INT", {"forceInput": True}), + }, + "optional": { + "current_mask": ("MASK",), + }, + "hidden": { + "dynprompt": "DYNPROMPT", + "unique_id": "UNIQUE_ID", + "iteration_count": ("INT", {"default": 0}), + } + } + return inputs + + RETURN_TYPES = tuple(["IMAGE", "MASK"]) + RETURN_NAMES = tuple(["final_image", "final_mask"]) + FUNCTION = "loop_close" + CATEGORY = "CyberEveLoop🐰" + + def explore_dependencies(self, node_id, dynprompt, upstream, parent_ids): + node_info = dynprompt.get_node(node_id) + if "inputs" not in node_info: + return + + for k, v in node_info["inputs"].items(): + if is_link(v): + parent_id = v[0] + display_id = dynprompt.get_display_node_id(parent_id) + display_node = dynprompt.get_node(display_id) + class_type = display_node["class_type"] + if class_type not in ['SingleImageLoopClose']: + parent_ids.append(display_id) + if parent_id not in upstream: + upstream[parent_id] = [] + self.explore_dependencies(parent_id, dynprompt, upstream, parent_ids) + upstream[parent_id].append(node_id) + + def explore_output_nodes(self, dynprompt, upstream, output_nodes, parent_ids): + for parent_id in upstream: + display_id = dynprompt.get_display_node_id(parent_id) + for output_id in output_nodes: + id = output_nodes[output_id][0] + if id in parent_ids and display_id == id and output_id not in upstream[parent_id]: + if '.' in parent_id: + arr = parent_id.split('.') + arr[len(arr)-1] = output_id + upstream[parent_id].append('.'.join(arr)) + else: + upstream[parent_id].append(output_id) + + def collect_contained(self, node_id, upstream, contained): + if node_id not in upstream: + return + for child_id in upstream[node_id]: + if child_id not in contained: + contained[child_id] = True + self.collect_contained(child_id, upstream, contained) + + def loop_close(self, flow_control, current_image, max_iterations, current_mask=None, + iteration_count=0, dynprompt=None, unique_id=None): + print(f"Iteration {iteration_count} of {max_iterations}") + + # 维度处理 + if len(current_image.shape) == 3: + current_image = current_image.unsqueeze(0) + if current_mask is not None and len(current_mask.shape) == 2: + current_mask = current_mask.unsqueeze(0) + + # 检查是否继续循环 + if iteration_count >= max_iterations - 1: + print(f"Loop finished with {iteration_count + 1} iterations") + return (current_image, current_mask if current_mask is not None else torch.zeros_like(current_image[:,:,:,0])) + + # 准备下一次循环 + this_node = dynprompt.get_node(unique_id) + upstream = {} + parent_ids = [] + self.explore_dependencies(unique_id, dynprompt, upstream, parent_ids) + parent_ids = list(set(parent_ids)) + + # 获取并处理输出节点 + prompts = dynprompt.get_original_prompt() + output_nodes = {} + for id in prompts: + node = prompts[id] + if "inputs" not in node: + continue + class_type = node["class_type"] + if class_type in ALL_NODE_CLASS_MAPPINGS: + class_def = ALL_NODE_CLASS_MAPPINGS[class_type] + if hasattr(class_def, 'OUTPUT_NODE') and class_def.OUTPUT_NODE == True: + for k, v in node['inputs'].items(): + if is_link(v): + output_nodes[id] = v + + # 创建新图 + graph = GraphBuilder() + self.explore_output_nodes(dynprompt, upstream, output_nodes, parent_ids) + + contained = {} + open_node = flow_control[0] + self.collect_contained(open_node, upstream, contained) + contained[unique_id] = True + contained[open_node] = True + + # 创建节点 + for node_id in contained: + original_node = dynprompt.get_node(node_id) + node = graph.node(original_node["class_type"], + "Recurse" if node_id == unique_id else node_id) + node.set_override_display_id(node_id) + + # 设置连接 + for node_id in contained: + original_node = dynprompt.get_node(node_id) + node = graph.lookup_node("Recurse" if node_id == unique_id else node_id) + for k, v in original_node["inputs"].items(): + if is_link(v) and v[0] in contained: + parent = graph.lookup_node(v[0]) + node.set_input(k, parent.out(v[1])) + else: + node.set_input(k, v) + + # 设置节点参数 + my_clone = graph.lookup_node("Recurse") + my_clone.set_input("iteration_count", iteration_count + 1) + + new_open = graph.lookup_node(open_node) + new_open.set_input("iteration_count", iteration_count + 1) + new_open.set_input("previous_image", current_image) + if current_mask is not None: + new_open.set_input("previous_mask", current_mask) + + print(f"Continuing to iteration {iteration_count + 1}") + + return { + "result": tuple([my_clone.out(0), my_clone.out(1)]), + "expand": graph.finalize(), + } + + + +@VariantSupport() +class LoopIndexSwitch: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + """ + 预定义100个隐藏的lazy输入 + """ + optional_inputs = { + "default_value": ("*", {"lazy": True}), # 默认值也设为lazy + } + # 添加100个隐藏的lazy输入 + hidden_inputs = {} + for i in range(100): + hidden_inputs[f"while_{i}"] = ("*", {"lazy": True}) + + return { + "required": { + "iteration_count": ("INT", {"forceInput": True}), # 当前迭代次数 + }, + "optional": optional_inputs, + "hidden": hidden_inputs, + } + + RETURN_TYPES = ("*",) + FUNCTION = "index_switch" + CATEGORY = "CyberEveLoop🐰" + + def check_lazy_status(self, iteration_count, **kwargs): + """ + 检查当前迭代需要的输入和默认值 + """ + needed = [] + current_key = f"while_{iteration_count}" + + # 检查当前迭代的输入 + if current_key in kwargs : + needed.append(current_key) + else: + needed.append("default_value") + + print(f"Index switch needed: {needed}") + return needed + + + + def index_switch(self, iteration_count, **kwargs): + """ + 根据当前迭代次数选择对应的输入值 + """ + current_key = f"while_{iteration_count}" + + if current_key in kwargs and kwargs[current_key] is not None: + return (kwargs[current_key],) + return (kwargs.get("default_value"),) + + +CyberEve_Loop_CLASS_MAPPINGS = { + "CyberEve_BatchImageLoopOpen": BatchImageLoopOpen, + "CyberEve_BatchImageLoopClose": BatchImageLoopClose, + "CyberEve_LoopIndexSwitch": LoopIndexSwitch, + "CyberEve_SingleImageLoopOpen": SingleImageLoopOpen, + "CyberEve_SingleImageLoopClose": SingleImageLoopClose, +} + +CyberEve_Loop_DISPLAY_NAME_MAPPINGS = { + "CyberEve_BatchImageLoopOpen": "Batch Image Loop Open🐰", + "CyberEve_BatchImageLoopClose": "Batch Image Loop Close🐰", + "CyberEve_LoopIndexSwitch": "Loop Index Switch🐰", + "CyberEve_SingleImageLoopOpen": "Single Image Loop Open🐰", + "CyberEve_SingleImageLoopClose": "Single Image Loop Close🐰", +} \ No newline at end of file diff --git a/mask_split.py b/mask_split.py new file mode 100644 index 0000000..0a75f6a --- /dev/null +++ b/mask_split.py @@ -0,0 +1,246 @@ +import torch +import torch.nn.functional as F +import cv2 +import numpy as np + + +class MaskSplit: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "mask": ("MASK",), + + }, + } + + RETURN_TYPES = ("IMAGE","MASK") + RETURN_NAMES = ("segmented_images","segmented_masks") + FUNCTION = "segment_mask" + + CATEGORY = "CyberEveLoop🐰" + + def find_top_left_point(self, mask_np): + """找到mask中最左上角的点""" + # 找到所有非零点 + y_coords, x_coords = np.nonzero(mask_np) + if len(x_coords) == 0: + return float('inf'), float('inf') + + # 找到最小x值 + min_x = np.min(x_coords) + # 在最小x值的点中找到最小y值 + min_y = np.min(y_coords[x_coords == min_x]) + + return min_x, min_y + + def segment_mask(self, mask, image): + """使用OpenCV快速分割蒙版并处理图像""" + # 保存原始设备信息 + device = mask.device if isinstance(mask, torch.Tensor) else torch.device('cpu') + + # 确保mask是正确的形状并转换为numpy数组 + if isinstance(mask, torch.Tensor): + if len(mask.shape) == 2: + mask = mask.unsqueeze(0) + mask_np = (mask[0] * 255).cpu().numpy().astype(np.uint8) + else: + mask_np = (mask * 255).astype(np.uint8) + + # 使用OpenCV找到轮廓 + contours, hierarchy = cv2.findContours( + mask_np, + cv2.RETR_TREE, + cv2.CHAIN_APPROX_SIMPLE + ) + + mask_info = [] # 用于排序的信息列表 + + if hierarchy is not None and len(contours) > 0: + hierarchy = hierarchy[0] + contour_masks = {} + + # 创建每个轮廓的mask + for i, contour in enumerate(contours): + mask = np.zeros_like(mask_np) + cv2.drawContours(mask, [contour], -1, 255, -1) + contour_masks[i] = mask + + # 处理每个轮廓 + processed_indices = set() + + for i, (contour, h) in enumerate(zip(contours, hierarchy)): + if i in processed_indices: + continue + + current_mask = contour_masks[i].copy() + child_idx = h[2] + + if child_idx != -1: + while child_idx != -1: + current_mask = cv2.subtract(current_mask, contour_masks[child_idx]) + processed_indices.add(child_idx) + child_idx = hierarchy[child_idx][0] + + # 找到最左上角的点 + min_x, min_y = self.find_top_left_point(current_mask) + + # 转换为tensor + mask_tensor = torch.from_numpy(current_mask).float() / 255.0 + mask_tensor = mask_tensor.unsqueeze(0) + mask_tensor = mask_tensor.to(device) + + # 保存mask和排序信息 + mask_info.append((mask_tensor, min_x, min_y)) + processed_indices.add(i) + + # 如果没有找到任何轮廓,使用原始mask + if not mask_info: + if isinstance(mask, torch.Tensor): + mask_info.append((mask, 0, 0)) + else: + mask_tensor = torch.from_numpy(mask).float() + if len(mask_tensor.shape) == 2: + mask_tensor = mask_tensor.unsqueeze(0) + mask_tensor = mask_tensor.to(device) + mask_info.append((mask_tensor, 0, 0)) + + # 根据最左上角点排序 + mask_info.sort(key=lambda x: (x[1], x[2])) + + # 确保image是正确的形状 + if len(image.shape) == 3: + image = image.unsqueeze(0) + + # 处理masks和images + result_masks = None + result_images = None + + for mask_tensor, _, _ in mask_info: + # 处理masks + if result_masks is None: + result_masks = mask_tensor + else: + result_masks = torch.cat([result_masks, mask_tensor], dim=0) + + # 处理images + if result_images is None: + result_images = image.clone() + else: + result_images = torch.cat([result_images, image.clone()], dim=0) + + return (result_images, result_masks) + + + +class MaskMerge: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "original_image": ("IMAGE",), + }, + "optional": { + "processed_images": ("IMAGE", {"forceInput": True}), + "masks": ("MASK", {"forceInput": True}), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("merged_image",) + # INPUT_IS_LIST = True + FUNCTION = "merge_masked_images" + + CATEGORY = "CyberEveLoop🐰" + + def resize_tensor(self, x, size, mode='bilinear'): + """调整tensor尺寸的辅助函数""" + # 确保输入是4D tensor [B,C,H,W] + orig_dim = x.dim() + if orig_dim == 3: + x = x.unsqueeze(0) + + # 如果是图像 [B,H,W,C],需要转换为 [B,C,H,W] + if x.shape[-1] in [1, 3, 4]: + x = x.permute(0, 3, 1, 2) + + # 执行调整 + x = F.interpolate(x, size=size, mode=mode, align_corners=False if mode in ['bilinear', 'bicubic'] else None) + + # 转换回原始格式 + if x.shape[1] in [1, 3, 4]: + x = x.permute(0, 2, 3, 1) + + # 如果原始输入是3D,去掉batch维度 + if orig_dim == 3: + x = x.squeeze(0) + + return x + + def merge_masked_images(self, original_image, processed_images=None, masks=None): + """合并处理后的图像""" + # 确保输入有效 + if processed_images is None or masks is None: + return (original_image,) + + # 确保原始图像维度正确 + if len(original_image.shape) == 3: + original_image = original_image.unsqueeze(0) + + # 创建结果图像的副本 + result = original_image.clone() + + # 确保processed_images和masks都是张量 + if isinstance(processed_images, list): + processed_images = torch.cat(processed_images, dim=0) + if isinstance(masks, list): + masks = torch.cat(masks, dim=0) + + # 获取目标尺寸 + target_height = original_image.shape[1] + target_width = original_image.shape[2] + + # 调整处理图像的尺寸(如果需要) + if processed_images.shape[1:3] != (target_height, target_width): + processed_images = self.resize_tensor( + processed_images, + (target_height, target_width), + mode='bilinear' + ) + + # 调整蒙版尺寸(如果需要) + if masks.shape[1:3] != (target_height, target_width): + masks = self.resize_tensor( + masks, + (target_height, target_width), + mode='bilinear' + ) + + # 扩展蒙版维度以匹配图像通道 + masks = masks.unsqueeze(-1).expand(-1, -1, -1, 3) + + # 批量处理所有图片 + for i in range(processed_images.shape[0]): + current_image = processed_images[i:i+1] + current_mask = masks[i:i+1] + result = current_mask * current_image + (1 - current_mask) * result + + return (result,) + +Mask_CLASS_MAPPINGS = { + "CyberEve_MaskSegmentation": MaskSplit, + "CyberEve_MaskMerge": MaskMerge, +} + +Mask_DISPLAY_NAME_MAPPINGS = { + "CyberEve_MaskSegmentation": "Mask Segmentation🐰", + "CyberEve_MaskMerge": "Mask Merge🐰", +} + diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..6f1e523 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +opencv-python +numpy \ No newline at end of file diff --git a/tools.py b/tools.py new file mode 100644 index 0000000..409bc71 --- /dev/null +++ b/tools.py @@ -0,0 +1,52 @@ +def MakeSmartType(t): + if isinstance(t, str): + return SmartType(t) + return t + +class SmartType(str): + def __ne__(self, other): + if self == "*" or other == "*": + return False + selfset = set(self.split(',')) + otherset = set(other.split(',')) + return not selfset.issubset(otherset) + +def VariantSupport(): + def decorator(cls): + if hasattr(cls, "INPUT_TYPES"): + old_input_types = getattr(cls, "INPUT_TYPES") + def new_input_types(*args, **kwargs): + types = old_input_types(*args, **kwargs) + for category in ["required", "optional"]: + if category not in types: + continue + for key, value in types[category].items(): + if isinstance(value, tuple): + types[category][key] = (MakeSmartType(value[0]),) + value[1:] + return types + setattr(cls, "INPUT_TYPES", new_input_types) + if hasattr(cls, "RETURN_TYPES"): + old_return_types = cls.RETURN_TYPES + setattr(cls, "RETURN_TYPES", tuple(MakeSmartType(x) for x in old_return_types)) + if hasattr(cls, "VALIDATE_INPUTS"): + # Reflection is used to determine what the function signature is, so we can't just change the function signature + raise NotImplementedError("VariantSupport does not support VALIDATE_INPUTS yet") + else: + def validate_inputs(input_types): + inputs = cls.INPUT_TYPES() + for key, value in input_types.items(): + if isinstance(value, SmartType): + continue + if "required" in inputs and key in inputs["required"]: + expected_type = inputs["required"][key][0] + elif "optional" in inputs and key in inputs["optional"]: + expected_type = inputs["optional"][key][0] + else: + expected_type = None + if expected_type is not None and MakeSmartType(value) != expected_type: + return f"Invalid type of {key}: {value} (expected {expected_type})" + return True + setattr(cls, "VALIDATE_INPUTS", validate_inputs) + return cls + return decorator + diff --git a/web/node/dynamicnode.js b/web/node/dynamicnode.js new file mode 100644 index 0000000..7694886 --- /dev/null +++ b/web/node/dynamicnode.js @@ -0,0 +1,108 @@ +import { app } from "../../../scripts/app.js" + +const _ID = "CyberEve_LoopIndexSwitch"; +const _PREFIX = "while_"; +const _TYPE = "*"; +const MAX_INPUTS = 100; + +app.registerExtension({ + name: 'CyberEveLoop.LoopIndexSwitch', + async beforeRegisterNodeDef(nodeType, nodeData, app) { + // 添加调试信息 + console.log("Registering extension for:", nodeData.name); + console.log("Looking for:", _ID); + if (nodeData.name !== _ID) { + return; + } + + // 添加右键菜单选项 + const onGetExtraMenuOptions = nodeType.prototype.getExtraMenuOptions; + nodeType.prototype.getExtraMenuOptions = function(_, options) { + if (onGetExtraMenuOptions) { + onGetExtraMenuOptions.apply(this, arguments); + } + + // 添加输入槽 + options.push({ + content: "Add Loop Input", + callback: () => { + // 弹出对话框让用户输入循环次数 + const number = prompt("Enter loop iteration number (0-99):", "0"); + if (number === null || isNaN(number)) { + return; + } + + const num = parseInt(number); + if (num < 0 || num >= MAX_INPUTS) { + alert(`Please enter a number between 0 and ${MAX_INPUTS-1}`); + return; + } + + const slotName = `${_PREFIX}${num}`; + + // 检查是否已存在该输入槽 + if (this.inputs.find(input => input.name === slotName)) { + alert(`Input for iteration ${num} already exists!`); + return; + } + + // 添加新的输入槽 + this.addInput(slotName, _TYPE); + this.graph.setDirtyCanvas(true); + } + }); + + // 删除输入槽子菜单 + const valueInputs = this.inputs.filter(input => + input.name.startsWith(_PREFIX) + ); + + if (valueInputs.length > 0) { + const removeOptions = valueInputs.map(input => ({ + content: `Remove iteration ${input.name.substring(_PREFIX.length)}`, + callback: () => { + const index = this.inputs.findIndex(i => i.name === input.name); + if (index !== -1) { + this.removeInput(index); + this.graph.setDirtyCanvas(true); + } + } + })); + + options.push({ + content: "Remove Loop Input", + submenu: { + options: removeOptions + } + }); + } + }; + + // 序列化节点时保存输入槽信息 + nodeType.prototype.serialize = function() { + const data = LGraphNode.prototype.serialize.apply(this); + data.inputSlots = this.inputs.filter(input => + input.name.startsWith(_PREFIX) + ).map(input => ({ + name: input.name, + type: input.type + })); + return data; + }; + + // 反序列化时恢复输入槽 + nodeType.prototype.configure = function(data) { + LGraphNode.prototype.configure.apply(this, arguments); + if (data.inputSlots) { + // 移除所有循环相关的输入 + this.inputs = this.inputs.filter(input => + !input.name.startsWith(_PREFIX) + ); + // 恢复保存的输入槽 + data.inputSlots.forEach(slot => { + this.addInput(slot.name, slot.type); + }); + } + }; + } +});