Trying to have some tolerance to None as argument

This commit is contained in:
Salvador E. Tropea
2025-11-12 07:28:31 -03:00
parent a34689b27c
commit f59f7b7d9f
2 changed files with 29 additions and 3 deletions
+19 -2
View File
@@ -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):
+10 -1
View File
@@ -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