diff --git a/__init__.py b/__init__.py index 6f2ba5d..5815486 100644 --- a/__init__.py +++ b/__init__.py @@ -1,8 +1,7 @@ -from .src import ImageSetAreaNode, FloatImageCombineNode, XYPlotNode +from .src import LatentCombineNode, XYPlotNode NODE_CLASS_MAPPINGS = { - "ImageSetArea": ImageSetAreaNode, - "FloatImageCombine": FloatImageCombineNode, "XYPlot": XYPlotNode, + "LatentCombine": LatentCombineNode, } diff --git a/src/__init__.py b/src/__init__.py index 5124cd7..b3ad20d 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -1,3 +1,2 @@ -from .nodes.image_set_area import ImageSetAreaNode -from .nodes.float_image_combine import FloatImageCombineNode from .nodes.xy_plot import XYPlotNode +from .nodes.latent_combine import LatentCombineNode diff --git a/src/base.py b/src/base.py index 2e35db6..de60d1a 100644 --- a/src/base.py +++ b/src/base.py @@ -1,4 +1,5 @@ import typing as t +from dataclasses import dataclass class BasePlotNode(): @@ -6,5 +7,10 @@ class BasePlotNode(): FUNCTION: str = "execute" +@dataclass +class KSamplerXYPlotInput(): + setting: str + value: int + + Image = t.Any -FloatImage = list[Image] diff --git a/src/nodes/float_image_combine.py b/src/nodes/float_image_combine.py index 9eac417..967b10a 100644 --- a/src/nodes/float_image_combine.py +++ b/src/nodes/float_image_combine.py @@ -1,19 +1,23 @@ import typing as t -from ..base import BasePlotNode, FloatImage +from ..base import BasePlotNode, Image class FloatImageCombineNode(BasePlotNode): - RETURN_TYPES: t.Tuple[str] = ("FLOAT_IMAGE",) + RETURN_TYPES: t.Tuple[str] = ("IMAGES",) @classmethod def INPUT_TYPES(cls) -> t.Dict[str, t.Any]: return { "required": { - "float_image_1": ("FLOAT_IMAGE",), - "float_image_2": ("FLOAT_IMAGE",), + "float_image_1": ("IMAGES",), + "float_image_2": ("IMAGES",), }, } - def execute(self, float_image_1: FloatImage, float_image_2: FloatImage) -> tuple[FloatImage]: + def execute( + self, + float_image_1: t.List[Image], + float_image_2: t.List[Image], + ) -> t.Tuple[t.List[Image]]: return (float_image_1 + float_image_2,) diff --git a/src/nodes/image_set_area.py b/src/nodes/image_set_area.py index acb25a4..41355d1 100644 --- a/src/nodes/image_set_area.py +++ b/src/nodes/image_set_area.py @@ -1,13 +1,10 @@ import typing as t -from ..base import BasePlotNode, FloatImage, Image +from ..base import BasePlotNode, Image class ImageSetAreaNode(BasePlotNode): - RETURN_TYPES: t.Tuple[str] = ("FLOAT_IMAGE",) - - def __init__(self): - pass + RETURN_TYPES: t.Tuple[str] = ("IMAGES",) @classmethod def INPUT_TYPES(cls) -> t.Dict[str, t.Any]: @@ -17,5 +14,5 @@ class ImageSetAreaNode(BasePlotNode): }, } - def execute(self, image: Image) -> tuple[FloatImage]: + def execute(self, image: Image) -> t.Tuple[t.List[Image]]: return ([image],) diff --git a/src/nodes/k_sampler_xy_plot.py b/src/nodes/k_sampler_xy_plot.py new file mode 100644 index 0000000..e4c52eb --- /dev/null +++ b/src/nodes/k_sampler_xy_plot.py @@ -0,0 +1,62 @@ +import typing as t + +from nodes import KSamplerAdvanced # type: ignore + +from ..base import BasePlotNode, Image, KSamplerXYPlotInput + + +class KSamplerXYPlotNode(BasePlotNode): + RETURN_TYPES: t.Tuple[str] = ("IMAGES",) + + def __init__(self) -> None: + self._sampler = KSamplerAdvanced() + + @classmethod + def INPUT_TYPES(cls): + result = KSamplerAdvanced.INPUT_TYPES() + result["required"]["vae"] = ("VAE", ) + #result["required"]["x_items"] = ("XYPlotItem",) + #result["required"]["y_items"] = ("XYPlotItem",) + return result + + def execute( + self, + vae, + #x_items, + #y_items, + **sampler_kw, + ) -> tuple[t.List[Image]]: + x_items = [ + KSamplerXYPlotInput(value=1, setting="cfg"), + KSamplerXYPlotInput(value=2, setting="cfg"), + ] + y_items = [ + KSamplerXYPlotInput(value=1, setting="noise_seed"), + KSamplerXYPlotInput(value=2, setting="noise_seed"), + ] + + latents = self._sample_latents( + x_items=x_items, + y_items=y_items, + sampler_kw=sampler_kw, + ) + result = list(self._decode_latents(latents=latents, vae=vae)) + print(result) + print(type(result[0])) + + return (result,) + + def _sample_latents(self, x_items, y_items, sampler_kw): + for x in x_items: + for y in y_items: + sampler_settings = sampler_kw.copy() + sampler_settings[x.setting] = x.value + sampler_settings[y.setting] = y.value + + yield self._sampler.sample(**sampler_settings)[0] + + def _decode_latents(self, latents, vae) -> t.Iterable[Image]: + return ( + vae.decode(i["samples"]) + for i in latents + ) diff --git a/src/nodes/latent_combine.py b/src/nodes/latent_combine.py new file mode 100644 index 0000000..6172cb2 --- /dev/null +++ b/src/nodes/latent_combine.py @@ -0,0 +1,29 @@ +import typing as t + +import torch + +from ..base import BasePlotNode, Image + + +class LatentCombineNode(BasePlotNode): + RETURN_TYPES: t.Tuple[str] = ("LATENT",) + + @classmethod + def INPUT_TYPES(cls) -> t.Dict[str, t.Any]: + return { + "required": { + "latent_1": ("LATENT",), + "latent_2": ("LATENT",), + }, + } + + def execute( + self, + latent_1: t.Dict[str, t.Any], + latent_2: t.Dict[str, t.Any], + ) -> t.Tuple[t.Dict[str, t.Any]]: + latent_1_samples = latent_1["samples"] + latent_2_samples = latent_2["samples"] + samples = torch.cat((latent_1_samples, latent_2_samples), 0) + + return ({"samples": samples},) diff --git a/src/nodes/xy_plot.py b/src/nodes/xy_plot.py index ce5e722..94ecbc9 100644 --- a/src/nodes/xy_plot.py +++ b/src/nodes/xy_plot.py @@ -1,6 +1,6 @@ import typing as t -from ..base import BasePlotNode, FloatImage, Image +from ..base import BasePlotNode, Image from ..utils import tensor_to_pillow, pillow_to_tensor, create_image_grid @@ -11,7 +11,7 @@ class XYPlotNode(BasePlotNode): def INPUT_TYPES(cls) -> t.Dict[str, t.Any]: return { "required": { - "float_image": ("FLOAT_IMAGE",), + "images": ("IMAGE",), "gap": ("INT", {"default": 0, "min": 0}), "nrow": ("INT", {"default": 1, "min": 1}), }, @@ -19,11 +19,11 @@ class XYPlotNode(BasePlotNode): def execute( self, - float_image: FloatImage, + images: Image, nrow: int, gap: int ) -> tuple[Image]: - pillow_images = [tensor_to_pillow(i) for i in float_image] + pillow_images = [tensor_to_pillow(i) for i in images] pillow_grid = create_image_grid(pillow_images, nrow=nrow, gap=gap) tensor_grid = pillow_to_tensor(pillow_grid) diff --git a/src/utils.py b/src/utils.py index d9a07c3..5e4bae0 100644 --- a/src/utils.py +++ b/src/utils.py @@ -13,35 +13,27 @@ def pillow_to_tensor(image): return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) -def create_image_grid(images, gap, nrow): - """ - Create a grid of images with a specified gap and number of rows. +def create_image_grid(images: t.List[Image.Image], gap: int, ncol: int): + # Calculate the number of rows needed based on the number of images and columns + nrow = (len(images) + ncol - 1) // ncol - Args: - images (List[PIL.Image.Image]): List of images to be placed in the grid. - gap (int, optional): Gap between images in pixels. Defaults to 10. - nrow (int, optional): Number of rows in the grid. Defaults to 3. + # Get the size of the first image to use as a template for the grid + size = images[0].size - Returns: - PIL.Image.Image: The merged image grid. - """ - # Calculate number of columns based on number of rows and images - ncol = (len(images) + nrow - 1) // nrow + # Calculate the total size of the grid with gaps + width = size[0] * ncol + gap * (ncol - 1) + height = size[1] * nrow + gap * (nrow - 1) - # Get size of each image in pixels - image_width, image_height = images[0].size + # Create a new image for the grid + grid_image = Image.new("RGB", (width, height), color="white") - # Create new image to hold the grid - grid_width = ncol * image_width + (ncol - 1) * gap - grid_height = nrow * image_height + (nrow - 1) * gap - grid_image = Image.new("RGB", (grid_width, grid_height), color="white") - - # Paste images into grid + # Iterate over each image and paste it into the grid for i, image in enumerate(images): - row = i // ncol - col = i % ncol - x = col * (image_width + gap) - y = row * (image_height + gap) + # Calculate the position of the image in the grid + x = (i % ncol) * (size[0] + gap) + y = (i // ncol) * (size[1] + gap) + + # Paste the image into the grid grid_image.paste(image, (x, y)) return grid_image diff --git a/workflows/xy_plot_base.json b/workflows/xy_plot_base.json index 54f5291..f2aca6e 100644 --- a/workflows/xy_plot_base.json +++ b/workflows/xy_plot_base.json @@ -131,8 +131,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 13 ], @@ -637,8 +637,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 22 ], @@ -672,8 +672,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 12 ], @@ -701,19 +701,19 @@ "inputs": [ { "name": "float_image_1", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 12 }, { "name": "float_image_2", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 13 } ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 32 ], @@ -741,19 +741,19 @@ "inputs": [ { "name": "float_image_1", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 32 }, { "name": "float_image_2", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 22 } ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 24 ], @@ -780,8 +780,8 @@ "mode": 0, "inputs": [ { - "name": "float_image", - "type": "FLOAT_IMAGE", + "name": "images", + "type": "IMAGES", "link": 24 } ], @@ -870,7 +870,7 @@ 0, 12, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 13, @@ -878,7 +878,7 @@ 0, 12, 1, - "FLOAT_IMAGE" + "IMAGES" ], [ 20, @@ -894,7 +894,7 @@ 0, 19, 1, - "FLOAT_IMAGE" + "IMAGES" ], [ 24, @@ -902,7 +902,7 @@ 0, 17, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 32, @@ -910,7 +910,7 @@ 0, 19, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 35, @@ -1133,4 +1133,4 @@ "config": {}, "extra": {}, "version": 0.4 -} \ No newline at end of file +} diff --git a/workflows/xy_plot_mini.json b/workflows/xy_plot_mini.json index 406dbd0..9d12f01 100644 --- a/workflows/xy_plot_mini.json +++ b/workflows/xy_plot_mini.json @@ -50,8 +50,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 12 ], @@ -79,19 +79,19 @@ "inputs": [ { "name": "float_image_1", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 12 }, { "name": "float_image_2", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 13 } ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 32 ], @@ -119,19 +119,19 @@ "inputs": [ { "name": "float_image_1", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 32 }, { "name": "float_image_2", - "type": "FLOAT_IMAGE", + "type": "IMAGES", "link": 22 } ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 24 ], @@ -166,8 +166,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 13 ], @@ -201,8 +201,8 @@ ], "outputs": [ { - "name": "FLOAT_IMAGE", - "type": "FLOAT_IMAGE", + "name": "IMAGES", + "type": "IMAGES", "links": [ 22 ], @@ -303,8 +303,8 @@ "mode": 0, "inputs": [ { - "name": "float_image", - "type": "FLOAT_IMAGE", + "name": "images", + "type": "IMAGES", "link": 24 } ], @@ -371,7 +371,7 @@ 0, 12, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 13, @@ -379,7 +379,7 @@ 0, 12, 1, - "FLOAT_IMAGE" + "IMAGES" ], [ 20, @@ -395,7 +395,7 @@ 0, 19, 1, - "FLOAT_IMAGE" + "IMAGES" ], [ 24, @@ -403,7 +403,7 @@ 0, 17, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 32, @@ -411,7 +411,7 @@ 0, 19, 0, - "FLOAT_IMAGE" + "IMAGES" ], [ 73, @@ -442,4 +442,4 @@ "config": {}, "extra": {}, "version": 0.4 -} \ No newline at end of file +}