Allow ApplyMaskToImage handle multiple masks (#11)

* Allow ApplyMaskToImage handle multiple masks
This commit is contained in:
Chenlei Hu
2024-03-04 09:52:38 +01:00
committed by GitHub
parent e27580efcd
commit bcb591c7b9
+22 -3
View File
@@ -131,11 +131,30 @@ class ApplyMaskToImage:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "apply_mask"
def apply_mask(self, image, mask):
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
# Move the channel to the second dimension for processing
out = image.movedim(-1, 1)
if out.shape[1] == 3: # RGB
# Check if the images are RGB, and if so, add an alpha channel initialized to 1
if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
# Ensure masks are unsqueezed to match the alpha channel dimension if needed
if mask.ndim == 2:
mask = mask.unsqueeze(0) # Add a batch dimension to masks
# For single mask, expand it to match size of image batch size.
if mask.shape[0] == 1:
mask = mask.repeat(out.shape[0], 1, 1)
assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
assert out.ndim == 4, f"Image should have shsape [B, C, H, W]. {out.shape}"
assert out.shape[-2:] == mask.shape[-2:], f"{out.shape[-2:]} != {mask.shape[-2:]}"
assert out.shape[0] == mask.shape[0], f"{out.shape[0]} != {mask.shape[0]}"
# Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]):
out[i, 3, :, :] = mask
out[i, 3, :, :] = mask[i]
# Move the channel back to its original dimension
out = out.movedim(1, -1)
return (out,)