From 84b6473ba0bbdafd9c57875a389282d239811391 Mon Sep 17 00:00:00 2001 From: chflame Date: Wed, 20 Mar 2024 21:50:48 +0800 Subject: [PATCH] fix bug of LayerColor nodes lost alpha --- py/color_adapter.py | 8 ++++---- py/color_correct_HSV.py | 4 ++++ py/color_correct_LAB.py | 4 ++++ py/color_correct_LUTapply.py | 4 ++++ py/color_correct_RGB.py | 4 ++++ py/color_correct_YUV.py | 4 ++++ py/color_correct_brightness&contrast.py | 21 ++++++++++++--------- py/color_correct_exposure.py | 8 +++++++- py/color_correct_gamma.py | 4 ++++ 9 files changed, 47 insertions(+), 14 deletions(-) diff --git a/py/color_adapter.py b/py/color_adapter.py index fe58a61..c02fb08 100644 --- a/py/color_adapter.py +++ b/py/color_adapter.py @@ -28,8 +28,6 @@ class ColorAdapter: def color_adapter(self, image, color_ref_image, opacity): ret_images = [] - # if color_ref_image.shape[0] > 0: - # color_ref_image = torch.unsqueeze(color_ref_image[0], 0) l_images = [] r_images = [] @@ -41,10 +39,12 @@ class ColorAdapter: _image = l_images[i] _ref = r_images[i] if len(ret_images) > i else r_images[-1] - _canvas = tensor2pil(_image).convert('RGB') + __image = tensor2pil(_image) + _canvas = __image.convert('RGB') ret_image = color_adapter(_canvas, tensor2pil(_ref).convert('RGB')) ret_image = chop_image(_canvas, ret_image, blend_mode='normal', opacity=opacity) - + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_HSV.py b/py/color_correct_HSV.py index 42f3394..76f04f1 100644 --- a/py/color_correct_HSV.py +++ b/py/color_correct_HSV.py @@ -33,6 +33,7 @@ class ColorCorrectHSV: for i in image: i = torch.unsqueeze(i,0) + __image = tensor2pil(i) _h, _s, _v = tensor2pil(i).convert('HSV').split() if H != 0 : _h = image_hue_offset(_h, H) @@ -42,6 +43,9 @@ class ColorCorrectHSV: _v = image_gray_offset(_v, V) ret_image = image_channel_merge((_h, _s, _v), 'HSV') + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_LAB.py b/py/color_correct_LAB.py index d87ff31..1c726c1 100644 --- a/py/color_correct_LAB.py +++ b/py/color_correct_LAB.py @@ -33,6 +33,7 @@ class ColorCorrectLAB: for i in image: i = torch.unsqueeze(i, 0) + __image = tensor2pil(i) _l, _a, _b = tensor2pil(i).convert('LAB').split() if L != 0 : _l = image_gray_offset(_l, L) @@ -42,6 +43,9 @@ class ColorCorrectLAB: _b = image_gray_offset(_b, B) ret_image = image_channel_merge((_l, _a, _b), 'LAB') + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_LUTapply.py b/py/color_correct_LUTapply.py index 02828e5..930817b 100644 --- a/py/color_correct_LUTapply.py +++ b/py/color_correct_LUTapply.py @@ -31,8 +31,12 @@ class ColorCorrectLUTapply: for i in image: i = torch.unsqueeze(i, 0) _image = tensor2pil(i) + lut_file = LUT_DICT[LUT] ret_image = apply_lut(_image, lut_file, log=(color_space == 'log')) + + if _image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, _image.split()[-1]) ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_RGB.py b/py/color_correct_RGB.py index 9f2368a..d97bf67 100644 --- a/py/color_correct_RGB.py +++ b/py/color_correct_RGB.py @@ -33,6 +33,7 @@ class ColorCorrectRGB: for i in image: i = torch.unsqueeze(i,0) + __image = tensor2pil(i) _r, _g, _b = tensor2pil(i).convert('RGB').split() if R != 0 : _r = image_gray_offset(_r, R) @@ -42,6 +43,9 @@ class ColorCorrectRGB: _b = image_gray_offset(_b, B) ret_image = image_channel_merge((_r, _g, _b), 'RGB') + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_YUV.py b/py/color_correct_YUV.py index 63c0a16..3f666f9 100644 --- a/py/color_correct_YUV.py +++ b/py/color_correct_YUV.py @@ -33,6 +33,7 @@ class ColorCorrectYUV: for i in image: i = torch.unsqueeze(i, 0) + __image = tensor2pil(i) _y, _u, _v = tensor2pil(i).convert('YCbCr').split() if Y != 0 : _y = image_gray_offset(_y, Y) @@ -42,6 +43,9 @@ class ColorCorrectYUV: _v = image_gray_offset(_v, V) ret_image = image_channel_merge((_y, _u, _v), 'YCbCr') + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') diff --git a/py/color_correct_brightness&contrast.py b/py/color_correct_brightness&contrast.py index 7c63ed5..f6da08a 100644 --- a/py/color_correct_brightness&contrast.py +++ b/py/color_correct_brightness&contrast.py @@ -33,18 +33,21 @@ class ColorCorrectBrightnessAndContrast: for i in image: i = torch.unsqueeze(i,0) - - _image = tensor2pil(i).convert('RGB') + __image = tensor2pil(i) + ret_image = __image.convert('RGB') if brightness != 1: - brightness_image = ImageEnhance.Brightness(_image) - _image = brightness_image.enhance(factor=brightness) + brightness_image = ImageEnhance.Brightness(ret_image) + ret_image = brightness_image.enhance(factor=brightness) if contrast != 1: - contrast_image = ImageEnhance.Contrast(_image) - _image = contrast_image.enhance(factor=contrast) + contrast_image = ImageEnhance.Contrast(ret_image) + ret_image = contrast_image.enhance(factor=contrast) if saturation != 1: - color_image = ImageEnhance.Color(_image) - _image = color_image.enhance(factor=saturation) - ret_images.append(pil2tensor(_image)) + color_image = ImageEnhance.Color(ret_image) + ret_image = color_image.enhance(factor=saturation) + + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') return (torch.cat(ret_images, dim=0),) diff --git a/py/color_correct_exposure.py b/py/color_correct_exposure.py index ba7bc59..08daa4b 100644 --- a/py/color_correct_exposure.py +++ b/py/color_correct_exposure.py @@ -31,6 +31,7 @@ class ColorCorrectExposure: for i in image: i = torch.unsqueeze(i, 0) + __image = tensor2pil(i) t = i.detach().clone().cpu().numpy().astype(np.float32) more = t[:, :, :, :3] > 0 t[:, :, :, :3][more] *= pow(2, exposure / 32) @@ -38,7 +39,12 @@ class ColorCorrectExposure: bp = -exposure / 250 scale = 1 / (1 - bp) t = np.clip((t - bp) * scale, 0.0, 1.0) - ret_images.append(torch.from_numpy(t)) + ret_image = tensor2pil(torch.from_numpy(t)) + + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') return (torch.cat(ret_images, dim=0),) diff --git a/py/color_correct_gamma.py b/py/color_correct_gamma.py index c2c70d5..eb1dc4b 100644 --- a/py/color_correct_gamma.py +++ b/py/color_correct_gamma.py @@ -31,8 +31,12 @@ class ColorCorrectGamma: for i in image: i = torch.unsqueeze(i, 0) + __image = tensor2pil(i) ret_image = gamma_trans(tensor2pil(i), gamma) + if __image.mode == 'RGBA': + ret_image = RGB2RGBA(ret_image, __image.split()[-1]) + ret_images.append(pil2tensor(ret_image)) log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')