This commit is contained in:
LEv145
2023-04-05 19:32:49 +02:00
parent 66651c816a
commit 38a0ae0fd0
11 changed files with 177 additions and 89 deletions
+2 -3
View File
@@ -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,
}
+1 -2
View File
@@ -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
+7 -1
View File
@@ -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]
+9 -5
View File
@@ -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,)
+3 -6
View File
@@ -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],)
+62
View File
@@ -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
)
+29
View File
@@ -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},)
+4 -4
View File
@@ -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)
+16 -24
View File
@@ -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
+21 -21
View File
@@ -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,
+21 -21
View File
@@ -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,