From 6e0dfa07a139f49bb5e8d20799e1e24598208820 Mon Sep 17 00:00:00 2001 From: "curious.jp" <48515264+curiousjp@users.noreply.github.com> Date: Tue, 28 May 2024 23:09:33 +1000 Subject: [PATCH] Rewrote flattenAgainstOriginal to use RGBA --- __init__.py | 44 +++++++++++++++++++++++++++++++++++++------- 1 file changed, 37 insertions(+), 7 deletions(-) diff --git a/__init__.py b/__init__.py index 1840379..e8f15ef 100644 --- a/__init__.py +++ b/__init__.py @@ -59,16 +59,46 @@ class FlattenAgainstOriginal: CATEGORY = "image" def flattenAgainstOriginal(self, base_image, candidates): + print('inbound shape', base_image.shape) + target = base_image.clone() + for b_idx, b in enumerate(target): + if b.shape[-1] == 4: + b_rgb = b[..., :3] + b_alp = b[..., -1:] + elif b.shape[-1] == 3: + b_rgb = b + b_alp = torch.ones( + b.shape[:-1] + (1,), + dtype = b.dtype, + device = b.device + ) + else: + raise ValueError('final dimension of base images must be 3 or 4') - if len(base_image) > 1: - raise Exception('ERROR: FlattenAgainstOriginal does not allow image batches for the base_image.') + for c_idx, c in enumerate(candidates): + c_rgb = c[..., :3] + c_alp = c[..., -1:] - cand_masks = (candidates != base_image) - flattened = base_image.clone() - for i, candidate in enumerate(candidates): - flattened = torch.where(cand_masks[i], candidate, flattened) + new_a = c_alp + (b_alp * (1 - c_alp)) + mask_area = (new_a > 0).squeeze(-1) - return (flattened,) + new_rgb = torch.zeros_like(b_rgb) + + new_rgb[mask_area] = ( + c_rgb[mask_area] * c_alp[mask_area] + + b_rgb[mask_area] * b_alp[mask_area] * (1 - c_alp[mask_area]) + ) / new_a[mask_area] + + b_rgb = new_rgb + b_alp = new_a + + if target.shape[-1] == 4: + target[b_idx] = torch.cat((b_rgb, b_alp), dim = -1) + else: + target[b_idx] = b_rgb + + print('return shape', target.shape) + return (target,) class CombinatorialDetailer: