diff --git a/histogram_matching.py b/histogram_matching.py new file mode 100644 index 0000000..f7b4694 --- /dev/null +++ b/histogram_matching.py @@ -0,0 +1,87 @@ +# from MIT licensed https://github.com/nemodleo/pytorch-histogram-matching +import torch +import torch.nn as nn +import torch.nn.functional as F + +class Histogram_Matching(nn.Module): + def __init__(self, differentiable=False): + super(Histogram_Matching, self).__init__() + self.differentiable = differentiable + + def forward(self, dst, ref): + # B C + B, C, H, W = dst.size() + # assertion + assert dst.device == ref.device + # [B*C 256] + hist_dst = self.cal_hist(dst) + hist_ref = self.cal_hist(ref) + # [B*C 256] + tables = self.cal_trans_batch(hist_dst, hist_ref) + # [B C H W] + rst = dst.clone() + for b in range(B): + for c in range(C): + rst[b,c] = tables[b*c, (dst[b,c] * 255).long()] + # [B C H W] + rst /= 255. + return rst + + def cal_hist(self, img): + B, C, H, W = img.size() + # [B*C 256] + if self.differentiable: + hists = self.soft_histc_batch(img * 255, bins=256, min=0, max=256, sigma=3*25) + else: + hists = torch.stack([torch.histc(img[b,c] * 255, bins=256, min=0, max=255) for b in range(B) for c in range(C)]) + hists = hists.float() + hists = F.normalize(hists, p=1) + # BC 256 + bc, n = hists.size() + # [B*C 256 256] + triu = torch.ones(bc, n, n, device=hists.device).triu() + # [B*C 256] + hists = torch.bmm(hists[:,None,:], triu)[:,0,:] + return hists + + def soft_histc_batch(self, x, bins=256, min=0, max=256, sigma=3*25): + # B C H W + B, C, H, W = x.size() + # [B*C H*W] + x = x.view(B*C, -1) + # 1 + delta = float(max - min) / float(bins) + # [256] + centers = float(min) + delta * (torch.arange(bins, device=x.device, dtype=torch.bfloat16) + 0.5) + # [B*C 1 H*W] + x = torch.unsqueeze(x, 1) + # [1 256 1] + centers = centers[None,:,None] + # [B*C 256 H*W] + x = x - centers + # [B*C 256 H*W] + x = x.type(torch.bfloat16) + # [B*C 256 H*W] + x = torch.sigmoid(sigma * (x + delta/2)) - torch.sigmoid(sigma * (x - delta/2)) + # [B*C 256] + x = x.sum(dim=2) + # [B*C 256] + x = x.type(torch.float32) + # prevent oom + # torch.cuda.empty_cache() + return x + + def cal_trans_batch(self, hist_dst, hist_ref): + # [B*C 256 256] + hist_dst = hist_dst[:,None,:].repeat(1,256,1) + # [B*C 256 256] + hist_ref = hist_ref[:,:,None].repeat(1,1,256) + # [B*C 256 256] + table = hist_dst - hist_ref + # [B*C 256 256] + table = torch.where(table>=0, 1., 0.) + # [B*C 256] + table = torch.sum(table, dim=1) - 1 + # [B*C 256] + table = torch.clamp(table, min=0, max=255) + return table diff --git a/image.py b/image.py index aa0bc17..b50e622 100644 --- a/image.py +++ b/image.py @@ -878,6 +878,126 @@ class ExtractKeyframes: return (image[keyframes], ','.join(map(str, keyframes)),) +class ImageColorMatch: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "reference": ("IMAGE",), + "color_space": (["LAB", "YCbCr", "RGB", "LUV", "YUV", "XYZ"],), + "factor": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05, }), + "device": (["auto", "cpu", "gpu"],), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "execute" + CATEGORY = "essentials/image processing" + + def execute(self, image, reference, color_space, factor, device): + import kornia + + if "gpu" == device: + device = comfy.model_management.get_torch_device() + elif "auto" == device: + device = comfy.model_management.intermediate_device() + else: + device = 'cpu' + + image = image.permute([0, 3, 1, 2]).to(device) + reference = reference.permute([0, 3, 1, 2]).to(device) + + if "LAB" == color_space: + image = kornia.color.rgb_to_lab(image) + reference = kornia.color.rgb_to_lab(reference) + elif "YCbCr" == color_space: + image = kornia.color.rgb_to_ycbcr(image) + reference = kornia.color.rgb_to_ycbcr(reference) + elif "LUV" == color_space: + image = kornia.color.rgb_to_luv(image) + reference = kornia.color.rgb_to_luv(reference) + elif "YUV" == color_space: + image = kornia.color.rgb_to_yuv(image) + reference = kornia.color.rgb_to_yuv(reference) + elif "XYZ" == color_space: + image = kornia.color.rgb_to_xyz(image) + reference = kornia.color.rgb_to_xyz(reference) + + image_mean, image_std = self.compute_mean_std(image) + reference_mean, reference_std = self.compute_mean_std(reference) + out = ((image - image_mean) / (image_std + 1e-6)) * (reference_std + 1e-6) + reference_mean + out = factor * out + (1 - factor) * image + + if "LAB" == color_space: + out = kornia.color.lab_to_rgb(out) + elif "YCbCr" == color_space: + out = kornia.color.ycbcr_to_rgb(out) + elif "LUV" == color_space: + out = kornia.color.luv_to_rgb(out) + elif "YUV" == color_space: + out = kornia.color.yuv_to_rgb(out) + elif "XYZ" == color_space: + out = kornia.color.xyz_to_rgb(out) + + out = out.permute([0, 2, 3, 1]).clamp(0, 1).to(comfy.model_management.intermediate_device()) + + return (out,) + + def compute_mean_std(self, image): + mean = torch.mean(image, dim=(2, 3), keepdim=True) + std = torch.std(image, dim=(2, 3), keepdim=True) + return mean, std + +class ImageHistogramMatch: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "reference": ("IMAGE",), + "method": (["pytorch", "skimage"],), + "factor": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05, }), + "device": (["auto", "cpu", "gpu"],), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "execute" + CATEGORY = "essentials/image processing" + + def execute(self, image, reference, method, factor, device): + if "gpu" == device: + device = comfy.model_management.get_torch_device() + elif "auto" == device: + device = comfy.model_management.intermediate_device() + else: + device = 'cpu' + + if "pytorch" in method: + from .histogram_matching import Histogram_Matching + + image = image.permute([0, 3, 1, 2]).to(device) + reference = reference.permute([0, 3, 1, 2]).to(device)[0].unsqueeze(0) + image.requires_grad = True + reference.requires_grad = True + + out = [] + + for i in image: + i = i.unsqueeze(0) + hm = Histogram_Matching(differentiable=True) + out.append(hm(i, reference)) + out = torch.cat(out, dim=0) + out = factor * out + (1 - factor) * image + out = out.permute([0, 2, 3, 1]).clamp(0, 1) + else: + from skimage.exposure import match_histograms + + out = torch.from_numpy(match_histograms(image.cpu().numpy(), reference.cpu().numpy(), channel_axis=3)).to(device) + out = factor * out + (1 - factor) * image.to(device) + + return (out.to(comfy.model_management.intermediate_device()),) """ ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ @@ -907,7 +1027,7 @@ class ImageToDevice: else: device = 'cpu' - image = image.to(device) + image = image.clone().to(device) torch.cuda.empty_cache() return (image,) @@ -975,6 +1095,8 @@ IMAGE_CLASS_MAPPINGS = { "ImageDesaturate+": ImageDesaturate, "PixelOEPixelize+": PixelOEPixelize, "ImagePosterize+": ImagePosterize, + "ImageColorMatch+": ImageColorMatch, + "ImageHistogramMatch+": ImageHistogramMatch, # Utilities "GetImageSize+": GetImageSize, @@ -1011,6 +1133,8 @@ IMAGE_NAME_MAPPINGS = { "ImageDesaturate+": "🔧 Image Desaturate", "PixelOEPixelize+": "🔧 Pixelize", "ImagePosterize+": "🔧 Image Posterize", + "ImageColorMatch+": "🔧 Image Color Match", + "ImageHistogramMatch+": "🔧 Image Histogram Match", # Utilities "GetImageSize+": "🔧 Get Image Size",