diff --git a/src/nodes/helpers.py b/src/nodes/helpers.py index 9a99de9..93a8f67 100644 --- a/src/nodes/helpers.py +++ b/src/nodes/helpers.py @@ -18,6 +18,17 @@ except Exception: logger = main_logger +def empty_image(b=0, h=64, w=64, c=1): + if b == 0 and c == 1: + dims = (h, w) + elif c == 1: + dims = (b, h, w) + else: + dims = (b, h, w, c) + + return torch.zeros(dims, dtype=torch.float32, device="cpu") + + def upscale(image, width, height, upscale_method): # return F.interpolate(image, size=(height, width), mode=upscale_method) return common_upscale(image, width, height, upscale_method, crop="disabled") @@ -85,7 +96,7 @@ class CustomLoadImage(object): mask = np.array(i.convert('RGBA').getchannel('A')).astype(np.float32) / 255.0 mask = 1. - torch.from_numpy(mask) else: - mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") + mask = empty_image() output_images.append(image) output_masks.append(mask.unsqueeze(0)) @@ -116,11 +127,13 @@ class CustomLoadMask(object): if c == 'A': mask = 1. - mask else: - mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") + mask = empty_image() return (mask.unsqueeze(0),) def get_image_preview_info(file_name, where="input"): + if file_name is None: + return {} # This information is for the preview, as we are an output node and we return images # they will be displayed in our node. Quite simple. if os.path.isabs(file_name): @@ -135,6 +148,8 @@ def get_image_preview_info(file_name, where="input"): def load_one_image(file_name, disp_name, embed_transparency): + if file_name is None: + return (empty_image(b=1, c=3), empty_image(b=1), file_name) if not os.path.isabs(file_name): file_name = os.path.join(get_input_directory(), file_name) if not os.path.exists(file_name): @@ -177,6 +192,8 @@ def load_one_image(file_name, disp_name, embed_transparency): def load_one_mask(file_name, disp_name, channel='red'): + if file_name is None: + return (empty_image(b=1), file_name) if not os.path.isabs(file_name): file_name = os.path.join(get_input_directory(), file_name) if not os.path.exists(file_name): diff --git a/src/nodes/nodes_img.py b/src/nodes/nodes_img.py index 7260803..49ed993 100644 --- a/src/nodes/nodes_img.py +++ b/src/nodes/nodes_img.py @@ -544,6 +544,9 @@ class ImageDataset: # As we progress the number changes and the node is evaluated again # When no files are left we catch the exception and return 0, so the node will be actually evaluated # But this time will raise the exception indicating the process finished. + logger.debug(f"ImageDataset.IS_CHANGED {source} {pattern} {destination}") + if source is None or destination is None: + return float("NaN") try: images, _, _ = cls.generate_lists(source, pattern, destination, dest_ext, reference, sort_method, MAX_FILES, skip_first_images, select_every_nth, random_seed, show_info=False) @@ -648,7 +651,8 @@ class ImageDataset: break if not len(images): - raise ValueError("Finished processing images") + # raise ValueError("Finished processing images") + return ([None], [None], [None]) if show_info: cur_len = len(images) total = n_files @@ -825,6 +829,8 @@ class SaliencyEvaluationMetrics: logger.debug("Prediction mask already [0, 1]") for i in range(gt.shape[0]): + if img_name[index_name] is None: + return ([None], img_name, None, None, None, None, None) # Get the next name imgp = Path(img_name[index_name]) index_name += 1 @@ -961,6 +967,9 @@ class ConsolidateMetrics: def execute(self, metrics, img_name, destination): # --- 1. Input Validation and Flattening --- + if metrics[0] is None or img_name[0] is None or destination[0] is None: + return () + # The inputs are just lists no real need to do much flat_metrics = metrics flat_names = img_name