diff --git a/flow_control.py b/flow_control.py index 1985ac7..88dbc46 100644 --- a/flow_control.py +++ b/flow_control.py @@ -27,24 +27,54 @@ class BatchImageLoopOpen: FUNCTION = "while_loop_open" CATEGORY = "CyberEveLoop🐰" + def standardize_input(self, images, masks): + """ + 标准化输入格式 + images: 确保是4D tensor [B,H,W,C] + masks: 确保是3D tensor [B,H,W] + """ + # 处理images + if isinstance(images, list): + images = torch.cat(images, dim=0) + if len(images.shape) == 3: # [H,W,C] -> [1,H,W,C] + images = images.unsqueeze(0) + assert len(images.shape) == 4, f"Images must be 4D [B,H,W,C], got shape {images.shape}" + + # 处理masks + if isinstance(masks, list): + masks = torch.cat(masks, dim=0) + if len(masks.shape) == 2: # [H,W] -> [1,H,W] + masks = masks.unsqueeze(0) + assert len(masks.shape) == 3, f"Masks must be 3D [B,H,W], got shape {masks.shape}" + + # 确保batch维度相同 + assert images.shape[0] == masks.shape[0], \ + f"Batch size mismatch: images {images.shape[0]} vs masks {masks.shape[0]}" + + return images, masks + + 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) + # 标准化输入 + segmented_images, segmented_masks = self.standardize_input(segmented_images, segmented_masks) + # 获取最大迭代次数 max_iterations = segmented_images.shape[0] if max_iterations == 0: raise ValueError("No images provided in segmented_images") + # 验证迭代计数(修改这里) + if iteration_count >= max_iterations: + raise ValueError(f"Iteration count {iteration_count} exceeds max iterations {max_iterations}") + # 获取当前迭代的图片和蒙版 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: @@ -116,32 +146,78 @@ class BatchImageLoopClose: contained[child_id] = True self.collect_contained(child_id, upstream, contained) + + def standardize_input(self, image, mask): + """ + 标准化输入格式 + image: 确保是4D tensor [B,H,W,C] + mask: 确保是3D tensor [B,H,W] + """ + # 处理image + if len(image.shape) == 3: # [H,W,C] -> [1,H,W,C] + image = image.unsqueeze(0) + assert len(image.shape) == 4, f"Image must be 4D [B,H,W,C], got shape {image.shape}" + + # 处理mask + if len(mask.shape) == 2: # [H,W] -> [1,H,W] + mask = mask.unsqueeze(0) + assert len(mask.shape) == 3, f"Mask must be 3D [B,H,W], got shape {mask.shape}" + + return image, mask + + + def initialize_results(self, max_iterations, current_image, current_mask): + """ + 初始化结果张量,确保与MaskSplit输出格式一致 + """ + # 确保维度正确 + assert len(current_image.shape) == 4, "Current image must be 4D [B,H,W,C]" + assert len(current_mask.shape) == 3, "Current mask must be 3D [B,H,W]" + + # 创建结果张量,确保格式一致 + result_images = torch.zeros( + (max_iterations, current_image.shape[1], current_image.shape[2], current_image.shape[3]), + dtype=current_image.dtype, + device=current_image.device + ) # 明确指定 [B,H,W,C] + + result_masks = torch.zeros( + (max_iterations, current_mask.shape[1], current_mask.shape[2]), + dtype=current_mask.dtype, + device=current_mask.device + ) # 明确指定 [B,H,W] + + return result_images, result_masks + + 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) + # 标准化输入,确保格式一致 + current_image, current_mask = self.standardize_input(current_image, current_mask) - # 结果初始化 - 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) + # 验证迭代计数(修改这里) + if iteration_count >= max_iterations: + raise ValueError(f"Iteration count {iteration_count} exceeds max iterations {max_iterations}") + + # 结果初始化或验证 + if result_images is None or result_masks is None: + result_images, result_masks = self.initialize_results(max_iterations, current_image, current_mask) + else: + # 验证现有结果的维度和格式 + assert result_images.shape[0] == max_iterations and len(result_images.shape) == 4, \ + f"Result images must be 4D [B,H,W,C] with batch size {max_iterations}" + assert result_masks.shape[0] == max_iterations and len(result_masks.shape) == 3, \ + f"Result masks must be 3D [B,H,W] with batch size {max_iterations}" # 存储当前结果 result_images[iteration_count:iteration_count+1] = current_image result_masks[iteration_count:iteration_count+1] = current_mask - - # 检查是否继续循环 - if iteration_count >= max_iterations - 1: + + # 检查是否继续循环(修改这里) + if iteration_count == max_iterations - 1: print(f"Loop finished with {iteration_count + 1} iterations") return (result_images, result_masks) diff --git a/mask_split.py b/mask_split.py index 0a75f6a..bca15b7 100644 --- a/mask_split.py +++ b/mask_split.py @@ -155,11 +155,40 @@ class MaskMerge: RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("merged_image",) - # INPUT_IS_LIST = True FUNCTION = "merge_masked_images" - CATEGORY = "CyberEveLoop🐰" + def standardize_input(self, image, processed_images=None, masks=None): + """ + 标准化输入格式 + - image: [H,W,C] -> [1,H,W,C] + - processed_images: [...] -> [B,H,W,C] + - masks: [...] -> [B,H,W] + """ + # 处理原始图像 + if len(image.shape) == 3: + image = image.unsqueeze(0) + assert len(image.shape) == 4, f"Original image must be 4D [B,H,W,C], got shape {image.shape}" + + # 处理processed_images + if processed_images is not None: + if isinstance(processed_images, list): + processed_images = torch.cat(processed_images, dim=0) + if len(processed_images.shape) == 3: + processed_images = processed_images.unsqueeze(0) + assert len(processed_images.shape) == 4, \ + f"Processed images must be 4D [B,H,W,C], got shape {processed_images.shape}" + + # 处理masks + if masks is not None: + if isinstance(masks, list): + masks = torch.cat(masks, dim=0) + if len(masks.shape) == 2: + masks = masks.unsqueeze(0) + assert len(masks.shape) == 3, f"Masks must be 3D [B,H,W], got shape {masks.shape}" + + return image, processed_images, masks + def resize_tensor(self, x, size, mode='bilinear'): """调整tensor尺寸的辅助函数""" # 确保输入是4D tensor [B,C,H,W] @@ -190,19 +219,14 @@ class MaskMerge: if processed_images is None or masks is None: return (original_image,) - # 确保原始图像维度正确 - if len(original_image.shape) == 3: - original_image = original_image.unsqueeze(0) + # 标准化输入 + original_image, processed_images, masks = self.standardize_input( + original_image, processed_images, masks + ) # 创建结果图像的副本 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] @@ -232,7 +256,9 @@ class MaskMerge: current_mask = masks[i:i+1] result = current_mask * current_image + (1 - current_mask) * result + assert len(result.shape) == 4, "Output must be 4D [B,H,W,C]" return (result,) + Mask_CLASS_MAPPINGS = { "CyberEve_MaskSegmentation": MaskSplit, diff --git a/pyproject.toml b/pyproject.toml index c7c2a02..1ce77e0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "loop-image" description = "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." -version = "1.0.3" +version = "1.0.4" license = {file = "LICENSE"} dependencies = ["opencv-python", "numpy"]