fixed for bacth inputs

This commit is contained in:
Ram Deshmukh
2024-03-06 17:25:13 +05:30
parent fd89979d62
commit 3f2ee1d80a
+23 -9
View File
@@ -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