2 Commits
Author SHA1 Message Date
Michael Poutre 600b304eee feat(Nodes): Add initial XY implementation 2023-09-16 00:38:08 -07:00
Michael Poutre 38192431a9 refactor(nodes): Deprecate StackImages 2023-09-16 00:38:05 -07:00
4 changed files with 269 additions and 27 deletions
+8 -1
View File
@@ -1,7 +1,8 @@
from custom_nodes.Comfy_KepListStuff.nodes.deprecated import StackImages
from custom_nodes.Comfy_KepListStuff.nodes.images import (
ImageLabelOverlay,
StackImages,
EmptyImages,
XYImage,
)
from custom_nodes.Comfy_KepListStuff.nodes.list_utils import (
ListLengthNode,
@@ -14,6 +15,7 @@ from custom_nodes.Comfy_KepListStuff.nodes.range_nodes import (
IntNumStepsRangeNode,
FloatNumStepsRangeNode,
)
from custom_nodes.Comfy_KepListStuff.nodes.xy import UnzippedProductAny
NODE_CLASS_MAPPINGS = {
"Range(Step) - Int": IntRangeNode,
@@ -26,5 +28,10 @@ NODE_CLASS_MAPPINGS = {
"Empty Images": EmptyImages,
"Join Image Lists": JoinImageLists,
"Join Float Lists": JoinFloatLists,
"XYAny": UnzippedProductAny,
"XYImage": XYImage
}
NODE_DISPLAY_NAME_MAPPINGS = {
"Stack Images": "Stack Images(Deprecated)",
}
+164
View File
@@ -0,0 +1,164 @@
from typing import Dict, Any, List, Optional, Tuple
from PIL import ImageFont, Image, ImageDraw
from torch import Tensor
import matplotlib.font_manager as fm
from custom_nodes.Comfy_KepListStuff.utils import tensor2pil, pil2tensor
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
# Our any instance wants to be a wildcard string
ANY = AnyType("*")
class StackImages:
def __init__(self) -> None:
pass
@classmethod
def INPUT_TYPES(s) -> Dict[str, Dict[str, Any]]:
return {
"required": {
"images": ("IMAGE",),
"splits": ("INT", {"forceInput": True, "min": 1}),
"stack_mode": (["horizontal", "vertical"], {"default": "horizontal"}),
"batch_stack_mode": (["horizontal", "vertical"], {"default": "horizontal"}),
},
"optional": {
"horizontal_labels": (ANY,{}),
"vertical_labels": (ANY,{}),
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Image",)
INPUT_IS_LIST = (True,)
OUTPUT_IS_LIST = (False,)
OUTPUT_NODE = True
FUNCTION = "stack_images"
CATEGORY = "List Stuff"
def stack_images(
self,
images: List[Tensor],
splits: List[int],
stack_mode: List[str],
batch_stack_mode: List[str],
horizontal_labels: Optional[List[str]] = None,
vertical_labels: Optional[List[str]] = None,
) -> Tuple[Tensor]:
if len(stack_mode) != 1:
raise Exception("Only single stack mode supported.")
if len(batch_stack_mode) != 1:
raise Exception("Only single batch stack mode supported.")
stack_direction = stack_mode[0]
batch_stack_direction = batch_stack_mode[0]
if len(splits) == 1:
splits = splits * (int(len(images) / splits[0]))
if sum(splits) != len(images):
splits.append(len(images) - sum(splits))
else:
if sum(splits) != len(images):
raise Exception("Sum of splits must equal number of images.")
batches = images
batch_size = len(batches[0])
image_h, image_w, _ = batches[0][0].size()
if batch_stack_direction == "horizontal":
batch_h = image_h
# stack horizontally
batch_w = image_w * batch_size
else:
# stack vertically
batch_h = image_h * batch_size
batch_w = image_w
if stack_direction == "horizontal":
full_w = batch_w * len(splits)
full_h = batch_h * max(splits)
else:
full_w = batch_w * max(splits)
full_h = batch_h * len(splits)
y_label_offset = 0
has_horizontal_labels = False
if horizontal_labels is not None:
horizontal_labels = [str(lbl) for lbl in horizontal_labels]
if stack_direction == "horizontal":
if len(horizontal_labels) != len(splits):
raise Exception("Number of horizontal labels must match number of splits.")
else:
if len(horizontal_labels) != max(splits):
raise Exception("Number of horizontal labels must match maximum split size.")
full_h += 60
y_label_offset = 60
has_horizontal_labels = True
x_label_offset = 0
has_vertical_labels = False
if vertical_labels is not None:
vertical_labels = [str(lbl) for lbl in vertical_labels]
if stack_direction == "horizontal":
if len(vertical_labels) != max(splits):
raise Exception("Number of vertical labels must match maximum split size.")
else:
if len(vertical_labels) != len(splits):
raise Exception("Number of vertical labels must match number of splits.")
full_w += 60
x_label_offset = 60
has_vertical_labels = True
full_image = Image.new("RGB", (full_w, full_h))
batch_idx = 0
if has_horizontal_labels:
assert horizontal_labels is not None
font = ImageFont.truetype(fm.findfont(fm.FontProperties()), 60)
for label_idx, label in enumerate(horizontal_labels):
x_offset = (batch_w * label_idx) + x_label_offset
draw = ImageDraw.Draw(full_image)
draw.rectangle((x_offset, 0, x_offset + batch_w, 60), fill="#ffffff")
draw.text((x_offset + (batch_w / 2), 0), label, fill="red", font=font)
if has_vertical_labels:
assert vertical_labels is not None
font = ImageFont.truetype(fm.findfont(fm.FontProperties()), 60)
for label_idx, label in enumerate(vertical_labels):
y_offset = (batch_h * label_idx) + y_label_offset
draw = ImageDraw.Draw(full_image)
draw.rectangle((0, y_offset, 60, y_offset + batch_h), fill="#ffffff")
draw.text((0, y_offset + (batch_h / 2)), label, fill="red", font=font)
for split_idx, split in enumerate(splits):
for idx_in_split in range(split):
batch_img = Image.new("RGB", (batch_w, batch_h))
batch = batches[batch_idx + idx_in_split]
if batch_stack_direction == "horizontal":
for img_idx, img in enumerate(batch):
x_offset = image_w * img_idx
batch_img.paste(tensor2pil(img), (x_offset, 0))
else:
for img_idx, img in enumerate(batch):
y_offset = image_h * img_idx
batch_img.paste(tensor2pil(img), (0, y_offset))
if stack_direction == "horizontal":
x_offset = batch_w * split_idx + x_label_offset
y_offset = batch_h * idx_in_split + y_label_offset
else:
x_offset = batch_w * idx_in_split + x_label_offset
y_offset = batch_h * split_idx + y_label_offset
full_image.paste(batch_img, (x_offset, y_offset))
batch_idx += split
return (pil2tensor(full_image),)
+30 -26
View File
@@ -99,7 +99,7 @@ class AnyType(str):
# Our any instance wants to be a wildcard string
any = AnyType("*")
class StackImages:
class XYImage:
def __init__(self) -> None:
pass
@@ -109,12 +109,12 @@ class StackImages:
"required": {
"images": ("IMAGE",),
"splits": ("INT", {"forceInput": True, "min": 1}),
"stack_mode": (["horizontal", "vertical"], {"default": "horizontal"}),
"flip_axis": (["False", "True"], {"default": "False"}),
"batch_stack_mode": (["horizontal", "vertical"], {"default": "horizontal"}),
},
"optional": {
"horizontal_labels": (any,{}),
"vertical_labels": (any,{}),
"x_labels": (any,{}),
"y_labels": (any,{}),
}
}
@@ -124,25 +124,29 @@ class StackImages:
INPUT_IS_LIST = (True,)
OUTPUT_IS_LIST = (False,)
OUTPUT_NODE = True
FUNCTION = "stack_images"
FUNCTION = "xy_image"
CATEGORY = "List Stuff"
def stack_images(
def xy_image(
self,
images: List[Tensor],
splits: List[int],
stack_mode: List[str],
flip_axis: List[str],
batch_stack_mode: List[str],
horizontal_labels: Optional[List[str]] = None,
vertical_labels: Optional[List[str]] = None,
x_labels: Optional[List[str]] = None,
y_labels: Optional[List[str]] = None,
) -> Tuple[Tensor]:
if len(stack_mode) != 1:
raise Exception("Only single stack mode supported.")
if len(flip_axis) != 1:
raise Exception("Only single flip_axis value supported.")
if len(batch_stack_mode) != 1:
raise Exception("Only single batch stack mode supported.")
stack_direction = stack_mode[0]
stack_direction = "horizontal"
if flip_axis[0] == "True":
stack_direction = "vertical"
x_labels, y_labels = y_labels, x_labels
batch_stack_direction = batch_stack_mode[0]
if len(splits) == 1:
@@ -175,13 +179,13 @@ class StackImages:
y_label_offset = 0
has_horizontal_labels = False
if horizontal_labels is not None:
horizontal_labels = [str(lbl) for lbl in horizontal_labels]
if x_labels is not None:
x_labels = [str(lbl) for lbl in x_labels]
if stack_direction == "horizontal":
if len(horizontal_labels) != len(splits):
if len(x_labels) != len(splits):
raise Exception("Number of horizontal labels must match number of splits.")
else:
if len(horizontal_labels) != max(splits):
if len(x_labels) != max(splits):
raise Exception("Number of horizontal labels must match maximum split size.")
full_h += 60
y_label_offset = 60
@@ -189,14 +193,14 @@ class StackImages:
x_label_offset = 0
has_vertical_labels = False
if vertical_labels is not None:
vertical_labels = [str(lbl) for lbl in vertical_labels]
if y_labels is not None:
y_labels = [str(lbl) for lbl in y_labels]
if stack_direction == "horizontal":
if len(vertical_labels) != max(splits):
raise Exception("Number of vertical labels must match maximum split size.")
if len(y_labels) != max(splits):
raise Exception(f"Number of vertical labels must match maximum split size. Got {len(y_labels)} labels for {max(splits)} splits.")
else:
if len(vertical_labels) != len(splits):
raise Exception("Number of vertical labels must match number of splits.")
if len(y_labels) != len(splits):
raise Exception(f"Number of vertical labels must match number of splits. Got {len(y_labels)} labels for {len(splits)} splits.")
full_w += 60
x_label_offset = 60
has_vertical_labels = True
@@ -207,18 +211,18 @@ class StackImages:
batch_idx = 0
if has_horizontal_labels:
assert horizontal_labels is not None
assert x_labels is not None
font = ImageFont.truetype(fm.findfont(fm.FontProperties()), 60)
for label_idx, label in enumerate(horizontal_labels):
for label_idx, label in enumerate(x_labels):
x_offset = (batch_w * label_idx) + x_label_offset
draw = ImageDraw.Draw(full_image)
draw.rectangle((x_offset, 0, x_offset + batch_w, 60), fill="#ffffff")
draw.text((x_offset + (batch_w / 2), 0), label, fill="red", font=font)
if has_vertical_labels:
assert vertical_labels is not None
assert y_labels is not None
font = ImageFont.truetype(fm.findfont(fm.FontProperties()), 60)
for label_idx, label in enumerate(vertical_labels):
for label_idx, label in enumerate(y_labels):
y_offset = (batch_h * label_idx) + y_label_offset
draw = ImageDraw.Draw(full_image)
draw.rectangle((0, y_offset, 60, y_offset + batch_h), fill="#ffffff")
+67
View File
@@ -0,0 +1,67 @@
import itertools
from typing import List, Any
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
# Our any instance wants to be a wildcard string
ANY = AnyType("*")
class UnzippedProductAny:
def __init__(self) -> None:
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"X": (ANY, {}),
"Y": (ANY, {}),
"X_Label_Fallback": (["str()", "Numbers"], {"default": "str()", "name": "X Label Fallback"}),
"Y_Label_Fallback": (["str()", "Numbers"], {"default": "str()", "name": "Y Label Fallback"})
},
"optional": {"X_Labels": (ANY, {"name": "X Labels"}), "Y_Labels": (ANY, {"name": "Y Labels"})},
}
RETURN_TYPES = (ANY, "STRING", ANY, "STRING", "INT", "INT")
RETURN_NAMES = ("X Values", "X Labels", "Y Values", "Y Labels", "Total Images", "Split Every")
OUTPUT_IS_LIST = (True, True, True, True, False, False)
INPUT_IS_LIST = True
FUNCTION = "to_xy"
CATEGORY = "List Stuff"
def to_xy(self, X: List[Any], Y: List[Any], X_Label_Fallback: List[str], Y_Label_Fallback: List[str], X_Labels: List[Any] = None, Y_Labels: List[Any] = None):
#region Validation
if len(X_Label_Fallback) != 1:
raise Exception("X_Label_Fallback must be a single value")
if len(Y_Label_Fallback) != 1:
raise Exception("Y_Label_Fallback must be a single value")
#region Labels
if X_Label_Fallback[0] == "str()":
get_x_fallback_labels = lambda x: [str(i) for i in x]
else:
get_x_fallback_labels = lambda x: list(range(len(x)))
if Y_Label_Fallback[0] == "str()":
get_y_fallback_labels = lambda x: [str(i) for i in x]
else:
get_y_fallback_labels = lambda x: list(range(len(x)))
if X_Labels is None:
X_Labels = get_x_fallback_labels(X)
if Y_Labels is None:
Y_Labels = get_y_fallback_labels(Y)
#endregion
product = itertools.product(X, Y)
X_out, Y_out = zip(*product)
X_out = list(X_out)
Y_out = list(Y_out)
return (X_out, X_Labels, Y_out, Y_Labels, len(X_out), len(Y))