This commit is contained in:
smthemex
2026-04-17 21:45:19 +08:00
committed by GitHub
parent e6c68da44c
commit 6fc657d09e
2 changed files with 251 additions and 82 deletions
+132 -82
View File
@@ -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
View File
@@ -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