init
This commit is contained in:
+132
-82
@@ -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):
|
||||
|
||||
+119
@@ -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
|
||||
Reference in New Issue
Block a user