Nodes for stacking and weighting reference images with flux redux model
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user