From e10daee9edea458fc709f60e725970a25567fca4 Mon Sep 17 00:00:00 2001 From: Acly Date: Sun, 24 Nov 2024 22:51:42 +0100 Subject: [PATCH] Nodes for stacking and weighting reference images with flux redux model --- __init__.py | 4 ++ nodes.py | 130 +++++++++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 133 insertions(+), 1 deletion(-) diff --git a/__init__.py b/__init__.py index 436c77a..a96c7f2 100644 --- a/__init__.py +++ b/__init__.py @@ -6,6 +6,8 @@ NODE_CLASS_MAPPINGS = { "ETN_SendImageWebSocket": nodes.SendImageWebSocket, "ETN_CropImage": nodes.CropImage, "ETN_ApplyMaskToImage": nodes.ApplyMaskToImage, + "ETN_ReferenceImage": nodes.ReferenceImage, + "ETN_ApplyReferenceImages": nodes.ApplyReferenceImages, "ETN_TileLayout": tile.TileLayout, "ETN_ExtractImageTile": tile.ExtractImageTile, "ETN_ExtractMaskTile": tile.ExtractMaskTile, @@ -32,6 +34,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ETN_SendImageWebSocket": "Send Image (WebSocket)", "ETN_CropImage": "Crop Image", "ETN_ApplyMaskToImage": "Apply Mask to Image", + "ETN_ReferenceImage": "Reference Image", + "ETN_ApplyReferenceImages": "Apply Reference Images", "ETN_TileLayout": "Create Tile Layout", "ETN_ExtractImageTile": "Extract Image Tile", "ETN_ExtractMaskTile": "Extract Mask Tile", diff --git a/nodes.py b/nodes.py index 68da235..4247510 100644 --- a/nodes.py +++ b/nodes.py @@ -1,11 +1,17 @@ from __future__ import annotations +from copy import copy +from typing import NamedTuple from PIL import Image import numpy as np import base64 import torch +import torch.nn.functional as F from io import BytesIO from server import PromptServer, BinaryEventTypes +from comfy.clip_vision import ClipVisionModel +from comfy.sd import StyleModel + class LoadImageBase64: @classmethod @@ -50,7 +56,8 @@ class LoadMaskBase64: if img.dim() == 3: # RGB(A) input, use red channel img = img[:, :, 0] return (img.unsqueeze(0),) - + + class SendImageWebSocket: @classmethod def INPUT_TYPES(s): @@ -84,6 +91,7 @@ class SendImageWebSocket: return {"ui": {"images": results}} + class CropImage: """Deprecated, ComfyUI has an ImageCrop node now which does the same.""" @@ -169,3 +177,123 @@ class ApplyMaskToImage: out[i, 3, :, :] = alpha return (to_bhwc(out),) + + +class _ReferenceImageData(NamedTuple): + image: torch.Tensor + weight: float + range: tuple[float, float] + + +class ReferenceImage: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}), + "range_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}), + "range_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), + }, + "optional": { + "reference_images": ("REFERENCE_IMAGE",), + }, + } + + CATEGORY = "external_tooling" + RETURN_TYPES = ("REFERENCE_IMAGE",) + RETURN_NAMES = ("reference_images",) + FUNCTION = "append" + + def append( + self, + image: torch.Tensor, + weight: float, + range_start: float, + range_end: float, + reference_images: list[_ReferenceImageData] | None = None, + ): + imgs = copy(reference_images) if reference_images is not None else [] + imgs.append(_ReferenceImageData(image, weight, (range_start, range_end))) + return (imgs,) + + +class ApplyReferenceImages: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "conditioning": ("CONDITIONING",), + "clip_vision": ("CLIP_VISION",), + "style_model": ("STYLE_MODEL",), + "references": ("REFERENCE_IMAGE",), + } + } + + CATEGORY = "external_tooling" + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "apply" + + def apply( + self, + conditioning: list[list], + clip_vision: ClipVisionModel, + style_model: StyleModel, + references: list[_ReferenceImageData], + ): + delimiters = {0.0, 1.0} + delimiters |= set(r.range[0] for r in references) + delimiters |= set(r.range[1] for r in references) + delimiters = sorted(delimiters) + ranges = [(delimiters[i], delimiters[i + 1]) for i in range(len(delimiters) - 1)] + + embeds = [_encode_image(r.image, clip_vision, style_model, r.weight) for r in references] + base = conditioning[0][0] + result = [] + for start, end in ranges: + e = [ + embeds[i] + for i, r in enumerate(references) + if r.range[0] <= start and r.range[1] >= end + ] + options = conditioning[0][1].copy() + options["start_percent"] = start + options["end_percent"] = end + result.append((torch.cat([base] + e, dim=1), options)) + + return (result,) + + +def _encode_image( + image: torch.Tensor, clip_vision: ClipVisionModel, style_model: StyleModel, weight: float +): + e = clip_vision.encode_image(image) + e = style_model.get_cond(e).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0) + e = _downsample_image_cond(e, weight) + return e + + +def _downsample_image_cond(cond: torch.Tensor, weight: float): + match weight: + case x if x >= 1.0: + return cond + case x if x <= 0.0: + return torch.zeros_like(cond) + case x if x >= 0.6: + factor = 2 + case x if x >= 0.3: + factor = 3 + case _: + factor = 4 + + # Downsample the clip vision embedding to make it smaller, resulting in less impact + # compared to other conditioning. + # See https://github.com/kaibioinfo/ComfyUI_AdvancedRefluxControl + (b, t, h) = cond.shape + m = int(np.sqrt(t)) + cond = F.interpolate( + cond.view(b, m, m, h).transpose(1, -1), + size=(m // factor, m // factor), + mode="area", + ) + return cond.transpose(1, -1).reshape(b, -1, h)