diff --git a/__init__.py b/__init__.py index 743c3ec..47ef93f 100644 --- a/__init__.py +++ b/__init__.py @@ -9,89 +9,20 @@ from PIL import Image, ImageFilter import math import comfy.utils import node_helpers +from einops import rearrange +import folder_paths +import os +from time import time -def tensor2image(tensor): +def tensor2image_sm(tensor): tensor = tensor.cpu() image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy() image = Image.fromarray(image_np, mode='RGB') return image -def phi2narry(img): +def phi2narry_sm(img): img = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0) return img -def bbox_from_mask(mask_l: Image.Image) -> tuple[int, int, int, int]: - """Return tight (x1, y1, x2, y2) bounding box of non-zero pixels.""" - arr = np.array(mask_l, dtype=np.uint8) - ys, xs = np.where(arr > 0) - if xs.size == 0: - raise ValueError("Mask is empty — nothing to refine.") - w, h = mask_l.size - return (max(0, int(xs.min())), - max(0, int(ys.min())), - min(w, int(xs.max()) + 1), - min(h, int(ys.max()) + 1)) - -def focus_crop( - image: Image.Image, - mask_l: Image.Image, - bbox: tuple[int, int, int, int], - margin: int = 64, - ) -> tuple[Image.Image, Image.Image, tuple[int, int, int, int]]: - """ - Crop around *bbox* so the diffusion model works on a ~1024² region. - - Returns (cropped_image, cropped_mask, crop_box). - """ - iw, ih = image.size - s = math.sqrt(1024 * 1024 / float(iw * ih)) - x1, y1, x2, y2 = bbox - - cx1 = max(0, int(math.floor(max(0.0, x1 * s - margin) / s))) - cy1 = max(0, int(math.floor(max(0.0, y1 * s - margin) / s))) - cx2 = min(iw, int(math.ceil(min(iw * s, x2 * s + margin) / s))) - cy2 = min(ih, int(math.ceil(min(ih * s, y2 * s + margin) / s))) - - crop_box = (cx1, cy1, cx2, cy2) - return image.crop(crop_box), mask_l.crop(crop_box), crop_box - -def paste_back( - original: Image.Image, - generated: Image.Image, - mask_l: Image.Image, - crop_box: tuple[int, int, int, int] | None = None, - mask_grow: int = 3, - blend_blur: int = 5, - ) -> Image.Image: - """Blend *generated* back into *original* through a smoothed mask.""" - m = mask_l.convert("L") - if mask_grow > 0: - m = m.filter(ImageFilter.MaxFilter(size=2 * mask_grow + 1)) - if blend_blur > 0: - m = m.filter(ImageFilter.GaussianBlur(radius=float(blend_blur))) - - target = original.crop(crop_box) if crop_box else original - dst = np.asarray(target.convert("RGB")).astype(np.float32) - src = np.asarray( - generated.convert("RGB").resize(target.size, Image.BICUBIC) - ).astype(np.float32) - alpha = np.asarray(m.resize(target.size, Image.BILINEAR)).astype(np.float32) / 255.0 - - # blended = src * alpha[:, :, None] + dst * (1.0 - alpha[:, :, None]) - blended = src * alpha[:, :, None] + dst * (1.0 - alpha[:, :, None]) - composited = Image.fromarray( - np.clip(blended, 0, 255).astype(np.uint8), mode="RGB" - ) - if crop_box: - result = original.copy() - result.paste(composited, (crop_box[0], crop_box[1])) - return result - return composited - -def binarise_mask_to_rgb(mask_l: Image.Image) -> Image.Image: - """Convert an L-mode mask to a clean binary RGB image (spatial condition for the model).""" - arr = np.where(np.array(mask_l, dtype=np.uint8) > 0, 255, 0).astype(np.uint8) - return Image.fromarray(arr, mode="L").convert("RGB") - class RefineAnything_Pasteback(io.ComfyNode): @classmethod def define_schema(cls): @@ -104,13 +35,94 @@ class RefineAnything_Pasteback(io.ComfyNode): io.Conditioning.Input("cond" ), io.Int.Input("mask_grow",default=3, min=0, max=4097, step=1, display_mode=io.NumberDisplay.number), io.Int.Input("blend_blur",default=5, min=0, max=4096, step=1, display_mode=io.NumberDisplay.number), + io.Boolean.Input("adain",default=True), + io.Boolean.Input("wavelet",default=False), + io.Boolean.Input("save_rgba",default=True), ], - outputs=[io.Image.Output(display_name="image"),], + outputs=[ + io.Image.Output(display_name="image"), + io.Image.Output(display_name="transparent_bg"), + ] + , ) @classmethod - def execute(cls,generated_image,cond,mask_grow,blend_blur ) -> io.NodeOutput: - result=paste_back(cond["origin_image"],tensor2image(generated_image),cond["model_mask"],crop_box=cond["crop_box"],mask_grow=mask_grow,blend_blur=blend_blur) - return io.NodeOutput(phi2narry(result)) + def execute(cls,generated_image,cond,mask_grow,blend_blur,adain,wavelet,save_rgba ) -> io.NodeOutput: + def paste_back( + original: Image.Image, + generated: Image.Image, + mask_l: Image.Image, + crop_box: tuple[int, int, int, int] | None = None, + mask_grow: int = 3, + blend_blur: int = 5, + wavelet: bool = False, + adain: bool = False, + save_rgba: bool = True, + + ) -> Image.Image: + """Blend *generated* back into *original* through a smoothed mask.""" + m = mask_l.convert("L") + if mask_grow > 0: + m = m.filter(ImageFilter.MaxFilter(size=2 * mask_grow + 1)) + if blend_blur > 0: + m = m.filter(ImageFilter.GaussianBlur(radius=float(blend_blur))) + + target = original.crop(crop_box) if crop_box else original + dst = np.asarray(target.convert("RGB")).astype(np.float32) + + alpha = np.asarray(m.resize(target.size, Image.BILINEAR)).astype(np.float32) / 255.0 + # blended = src * alpha[:, :, None] + dst * (1.0 - alpha[:, :, None]) + if adain: + from .align_color import adain_color_fix + src =adain_color_fix(generated.convert("RGB").resize(target.size, Image.BICUBIC) , target) + src=np.asarray(src).astype(np.float32) + else: + src = np.asarray( + generated.convert("RGB").resize(target.size, Image.BICUBIC) + ).astype(np.float32) + blended = src * alpha[:, :, None] + dst * (1.0 - alpha[:, :, None]) + composited = Image.fromarray( + np.clip(blended, 0, 255).astype(np.uint8), mode="RGB" + ) + transparent_bg = Image.new('RGBA', original.size, (0, 0, 0, 0)) + prefix = f"composited_{int(time())}" + if wavelet: + from .align_color import wavelet_reconstruction + x1 = phi2narry_sm(composited).permute(0, 3, 1, 2) #--> torch.Size([1, 3, 1024, 1024]) + x1 = rearrange(x1[-1], "c h w -> h w c").to("cpu") + x1 = wavelet_reconstruction(x1.permute(2, 0, 1), phi2narry_sm(target).permute(0, 3, 1, 2).squeeze(0).to("cpu")) + x1 = x1.clamp(0, 1) + img=x1.unsqueeze(0).permute(0, 2, 3, 1) #torch.Size([1, 673, 818, 3]) + if crop_box: + original_tensor = phi2narry_sm(original) + x1, y1, x2, y2 = crop_box + comp_h, comp_w = y2 - y1, x2 - x1 + composited_resized = torch.nn.functional.interpolate( + img.permute(0, 3, 1, 2), + size=(comp_h, comp_w), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) # [H_crop, W_crop, C] + result_tensor = original_tensor.clone() + + result_rgba= tensor2image_sm(composited_resized[0]) + transparent_bg.paste(result_rgba,(crop_box[0], crop_box[1])) + if save_rgba: + transparent_bg.save(os.path.join(folder_paths.get_output_directory(), f"{prefix}.png")) + result_tensor[0, y1:y2, x1:x2, :] = composited_resized[0] + return result_tensor,transparent_bg + return img,transparent_bg + if crop_box: + result = original.copy() + result.paste(composited, (crop_box[0], crop_box[1])) + + transparent_bg.paste(composited,(crop_box[0], crop_box[1])) + if save_rgba: + transparent_bg.save(os.path.join(folder_paths.get_output_directory(), f"{prefix}.png")) + return result,transparent_bg + return composited ,transparent_bg + result,transparent_bg=paste_back(cond["origin_image"],tensor2image_sm(generated_image),cond["model_mask"],crop_box=cond["crop_box"],mask_grow=mask_grow,blend_blur=blend_blur,wavelet=wavelet, adain=adain,save_rgba=save_rgba) + + return io.NodeOutput( result if wavelet else phi2narry_sm(result),phi2narry_sm(transparent_bg)) class RefineAnything_PreImg(io.ComfyNode): @classmethod @@ -133,20 +145,58 @@ class RefineAnything_PreImg(io.ComfyNode): @classmethod def execute(cls,origin_image,mask_image,do_focus_crop, ) -> io.NodeOutput: - origin_image = tensor2image(origin_image) - mask_l = tensor2image(mask_image).convert("L") + def binarise_mask_to_rgb(mask_l: Image.Image) -> Image.Image: + """Convert an L-mode mask to a clean binary RGB image (spatial condition for the model).""" + arr = np.where(np.array(mask_l, dtype=np.uint8) > 0, 255, 0).astype(np.uint8) + return Image.fromarray(arr, mode="L").convert("RGB") + def bbox_from_mask(mask_l: Image.Image) -> tuple[int, int, int, int]: + """Return tight (x1, y1, x2, y2) bounding box of non-zero pixels.""" + arr = np.array(mask_l, dtype=np.uint8) + ys, xs = np.where(arr > 0) + if xs.size == 0: + raise ValueError("Mask is empty — nothing to refine.") + w, h = mask_l.size + return (max(0, int(xs.min())), + max(0, int(ys.min())), + min(w, int(xs.max()) + 1), + min(h, int(ys.max()) + 1)) + + origin_image = tensor2image_sm(origin_image) + mask_l = tensor2image_sm(mask_image).convert("L") if mask_l.size != origin_image.size: mask_l = mask_l.resize(origin_image.size, Image.NEAREST) bbox = bbox_from_mask(mask_l) model_image, model_mask = origin_image, mask_l crop_box=None + def focus_crop( + image: Image.Image, + mask_l: Image.Image, + bbox: tuple[int, int, int, int], + margin: int = 64, + ) -> tuple[Image.Image, Image.Image, tuple[int, int, int, int]]: + """ + Crop around *bbox* so the diffusion model works on a ~1024² region. + + Returns (cropped_image, cropped_mask, crop_box). + """ + iw, ih = image.size + s = math.sqrt(1024 * 1024 / float(iw * ih)) + x1, y1, x2, y2 = bbox + + cx1 = max(0, int(math.floor(max(0.0, x1 * s - margin) / s))) + cy1 = max(0, int(math.floor(max(0.0, y1 * s - margin) / s))) + cx2 = min(iw, int(math.ceil(min(iw * s, x2 * s + margin) / s))) + cy2 = min(ih, int(math.ceil(min(ih * s, y2 * s + margin) / s))) + + crop_box = (cx1, cy1, cx2, cy2) + return image.crop(crop_box), mask_l.crop(crop_box), crop_box if do_focus_crop: model_image, model_mask, crop_box = focus_crop( origin_image, mask_l, bbox, margin=64, ) cond={"crop_box":crop_box,"model_mask":model_mask,"origin_image":origin_image} - return io.NodeOutput(phi2narry(model_image),phi2narry(binarise_mask_to_rgb(model_mask)),cond) + return io.NodeOutput(phi2narry_sm(model_image),phi2narry_sm(binarise_mask_to_rgb(model_mask)),cond) class TextEncodeQwenImageEditPlus_NoAppend(io.ComfyNode): diff --git a/align_color.py b/align_color.py new file mode 100644 index 0000000..b519294 --- /dev/null +++ b/align_color.py @@ -0,0 +1,119 @@ +''' +# -------------------------------------------------------------------------------- +# Color fixed script from Li Yi (https://github.com/pkuliyi2015/sd-webui-stablesr/blob/master/srmodule/colorfix.py) +# -------------------------------------------------------------------------------- +''' + +import torch +from PIL import Image +from torch import Tensor +from torch.nn import functional as F +from torchvision.transforms import ToTensor, ToPILImage + + +def adain_color_fix(target: Image, source: Image): + # Convert images to tensors + to_tensor = ToTensor() + target_tensor = to_tensor(target).unsqueeze(0) + source_tensor = to_tensor(source).unsqueeze(0) + + # Apply adaptive instance normalization + result_tensor = adaptive_instance_normalization(target_tensor, source_tensor) + + # Convert tensor back to image + to_image = ToPILImage() + result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0)) + + return result_image + +def wavelet_color_fix(target: Image, source: Image): + # Convert images to tensors + to_tensor = ToTensor() + target_tensor = to_tensor(target).unsqueeze(0) + source_tensor = to_tensor(source).unsqueeze(0) + + # Apply wavelet reconstruction + result_tensor = wavelet_reconstruction(target_tensor, source_tensor) + + # Convert tensor back to image + to_image = ToPILImage() + result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0)) + + return result_image + +def calc_mean_std(feat: Tensor, eps=1e-5): + """Calculate mean and std for adaptive_instance_normalization. + Args: + feat (Tensor): 4D tensor. + eps (float): A small value added to the variance to avoid + divide-by-zero. Default: 1e-5. + """ + size = feat.size() + assert len(size) == 4, 'The input feature should be 4D tensor.' + b, c = size[:2] + feat_var = feat.reshape(b, c, -1).var(dim=2) + eps + feat_std = feat_var.sqrt().reshape(b, c, 1, 1) + feat_mean = feat.reshape(b, c, -1).mean(dim=2).reshape(b, c, 1, 1) + return feat_mean, feat_std + +def adaptive_instance_normalization(content_feat:Tensor, style_feat:Tensor): + """Adaptive instance normalization. + Adjust the reference features to have the similar color and illuminations + as those in the degradate features. + Args: + content_feat (Tensor): The reference feature. + style_feat (Tensor): The degradate features. + """ + size = content_feat.size() + style_mean, style_std = calc_mean_std(style_feat) + content_mean, content_std = calc_mean_std(content_feat) + normalized_feat = (content_feat - content_mean.expand(size)) / content_std.expand(size) + return normalized_feat * style_std.expand(size) + style_mean.expand(size) + +def wavelet_blur(image: Tensor, radius: int): + """ + Apply wavelet blur to the input tensor. + """ + # input shape: (1, 3, H, W) + # convolution kernel + kernel_vals = [ + [0.0625, 0.125, 0.0625], + [0.125, 0.25, 0.125], + [0.0625, 0.125, 0.0625], + ] + kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device) + # add channel dimensions to the kernel to make it a 4D tensor + kernel = kernel[None, None] + # repeat the kernel across all input channels + kernel = kernel.repeat(3, 1, 1, 1) + image = F.pad(image, (radius, radius, radius, radius), mode='replicate') + # apply convolution + output = F.conv2d(image, kernel, groups=3, dilation=radius) + return output + +def wavelet_decomposition(image: Tensor, levels=5): + """ + Apply wavelet decomposition to the input tensor. + This function only returns the low frequency & the high frequency. + """ + high_freq = torch.zeros_like(image) + for i in range(levels): + radius = 2 ** i + low_freq = wavelet_blur(image, radius) + high_freq += (image - low_freq) + image = low_freq + + return high_freq, low_freq + +def wavelet_reconstruction(content_feat:Tensor, style_feat:Tensor): + """ + Apply wavelet decomposition, so that the content will have the same color as the style. + """ + # calculate the wavelet decomposition of the content feature + content_high_freq, content_low_freq = wavelet_decomposition(content_feat) + del content_low_freq + # calculate the wavelet decomposition of the style feature + style_high_freq, style_low_freq = wavelet_decomposition(style_feat) + del style_high_freq + # reconstruct the content feature with the style's high frequency + return content_high_freq + style_low_freq \ No newline at end of file