From 3f2ee1d80aa673dc6a5ad5d37e86ca2fd1ff3271 Mon Sep 17 00:00:00 2001 From: Ram Deshmukh Date: Wed, 6 Mar 2024 17:25:13 +0530 Subject: [PATCH] fixed for bacth inputs --- __init__.py | 32 +++++++++++++++++++++++--------- 1 file changed, 23 insertions(+), 9 deletions(-) diff --git a/__init__.py b/__init__.py index c5ecf32..08c3ed4 100644 --- a/__init__.py +++ b/__init__.py @@ -2391,7 +2391,7 @@ class TRI3DCompositeImageSplitter: """ return { "required": { - "image": ("IMAGE",) + "images": ("IMAGE",) }, } @@ -2401,16 +2401,30 @@ class TRI3DCompositeImageSplitter: CATEGORY = "TRI3D" - def main(self, image): - image = from_torch_image(image) - h,w,_ = image.shape - image1 = image[:,:w//2,:] - image2 = image[:,w//2:,:] + def main(self, images): + + def tensor_to_cv2_img(tensor, remove_alpha=False): + i = 255. * tensor.cpu().numpy() # This will give us (H, W, C) + img = np.clip(i, 0, 255).astype(np.uint8) + return img - image1 = torch.from_numpy(image1.astype(np.float32)/255.0)[None,] - image2 = torch.from_numpy(image2.astype(np.float32)/255.0)[None,] + def cv2_img_to_tensor(img): + img = img.astype(np.float32) / 255.0 + img = torch.from_numpy(img)[ + None, + ] + return img + images1 = [] + images2 = [] + for image in images: + image = tensor_to_cv2_img(image) + h,w,_ = image.shape + image1 = image[:,:w//2,:] + image2 = image[:,w//2:,:] + images1.append(cv2_img_to_tensor(image1).squeeze(0)) + images2.append(cv2_img_to_tensor(image2).squeeze(0)) - return (image1, image2) + return (torch.stack(images1), torch.stack(images2)) # A dictionary that contains all nodes you want to export with their names # NOTE: names should be globally unique