Refactor modules and add some image nodes

This commit is contained in:
Alistair Buxton
2023-09-11 00:33:08 +01:00
parent db319d7e86
commit 3b2086ff06
5 changed files with 184 additions and 83 deletions
+1 -1
View File
@@ -10,7 +10,7 @@ def register_node(c):
return c
from . import sequence, paths, job, misc
from . import sequence, paths, job, image, debug
+20
View File
@@ -7,6 +7,25 @@ from . import register_node
@register_node
class Stringify:
"""Convert any input to str/repr."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"x": ("*", ),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("str", "repr")
FUNCTION = "go"
CATEGORY = "ali1234/debug"
def go(self, x):
return (str(x), repr(x))
class RestoreStdStreams(object):
# ComfyUI-Manager patches sys.stdout and sys.stder
# which breaks GNU Readline support and makes the
@@ -39,6 +58,7 @@ class Quitter:
print(MESSAGE)
@register_node
class Interact:
"""Opens an interactive REPL whenever the node is evaluated."""
@classmethod
+161
View File
@@ -0,0 +1,161 @@
import torch
import numpy as np
from . import register_node
@register_node
class JoinImageBatch:
"""Turns an image batch into one big image."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"mode": (("horizontal", "vertical"), {"default": "horizontal"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "join"
CATEGORY = "ali1234/image"
def join(self, images, mode):
n, h, w, c = images.shape
image = None
if mode == "vertical":
# for vertical we can just reshape
image = images.reshape(1, n * h, w, c)
elif mode == "horizontal":
# for horizontal we have to swap axes
image = torch.transpose(torch.transpose(images, 1, 2).reshape(1, n * w, h, c), 1, 2)
return (image,)
@register_node
class JoinImages:
"""Turns joins two images into one big image."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image_a": ("IMAGE",),
"image_b": ("IMAGE",),
"mode": (("horizontal", "vertical"), {"default": "horizontal"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "join"
CATEGORY = "ali1234/image"
def join(self, image_a, image_b, mode):
dim = {'horizontal': 2, 'vertical': 1}[mode]
return (torch.concat((image_a, image_b), dim), )
@register_node
class SelectImageBatch:
"""Selects one image from an image batch."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"select": ("INT", {"default": 0, "min": 0, "max": 99999, "step": 1}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "select"
CATEGORY = "ali1234/image"
def select(self, images, select):
n, h, w, c = images.shape
if select >= n:
select = n - 1
return (images[select].reshape(1, h, w, c),)
@register_node
class GetImageSize:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
},
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("width", "height")
FUNCTION = "go"
CATEGORY = "ali1234/image"
def go(self, images):
return (images.shape[2], images.shape[1])
@register_node
class StringToImage:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"default": "Hello world!"}),
"width": ("INT", {"default": 384}),
"height": ("INT", {"default": 16}),
"colour": ("COLOR", {"default": "white"}),
"background": ("COLOR", {"default": "black"}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "render"
CATEGORY = "ali1234/image"
def render(self, text, width, height, colour, background):
from PIL import Image, ImageDraw, ImageFont
font = ImageFont.load_default()
img = Image.new("RGB", (width, height), background)
draw = ImageDraw.Draw(img)
_, _, w, h = draw.textbbox((0, 0), text, font=font)
draw.text(((width - w) / 2, ((height - h) / 2) - 1), text, font=font, fill=colour)
tensor = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
return (tensor,)
@register_node
class ProgressBar:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"progress": ("FLOAT", {'default': 0, 'forceInput': True}),
"padding": ("INT", {'default': 3}),
"width": ("INT", {"default": 384}),
"height": ("INT", {"default": 16}),
"colour": ("COLOR", {"default": "white"}),
"background": ("COLOR", {"default": "black"}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "render"
CATEGORY = "ali1234/image"
def render(self, progress, padding, width, height, colour, background):
from PIL import Image, ImageDraw
img = Image.new("RGB", (width, height), background)
draw = ImageDraw.Draw(img)
draw.rectangle((padding, padding, width - padding - 1, height - padding - 1), outline=colour)
if progress > 0:
ip = padding + 2
draw.rectangle((ip, ip, max(ip+1, (width - ip - 1) * progress), height - ip - 1), outline=colour, fill=colour)
tensor = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
return (tensor,)
+2 -3
View File
@@ -26,7 +26,6 @@ class MakeJob:
def merge_dicts(self, *dicts):
#return collections.ChainMap(*reversed(dicts))
print(dicts)
return dict(itertools.chain.from_iterable(d.items() for d in dicts))
def go(self, sequence, name):
@@ -178,7 +177,7 @@ class JobIterator:
CATEGORY = "ali1234/job"
def go(self, job, start_step):
print(f'JobIterator: {start_step} / {len(job) - 1}')
print(f'JobIterator: {start_step + 1} / {len(job)}')
return (job[start_step], len(job), start_step)
@@ -188,7 +187,7 @@ orig_execute = PromptExecutor.execute
def execute(self, prompt, prompt_id, extra_data={}, execute_outputs=[]):
print("Prompt executor has been patched!")
print("Prompt executor has been patched by Job Iterator!")
orig_execute(self, prompt, prompt_id, extra_data, execute_outputs)
job_iterator = None
-79
View File
@@ -1,79 +0,0 @@
import torch
from . import register_node
@register_node
class Stringify:
"""Convert any input to str/repr."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"x": ("*", ),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("str", "repr")
FUNCTION = "go"
CATEGORY = "ali1234/debug"
def go(self, x):
return (str(x), repr(x))
@register_node
class JoinImageBatch:
"""Turns an image batch into one big image."""
@classmethod
def INPUT_TYPES(s):
"""
Joins an image batch into a single image.
"""
return {
"required": {
"images": ("IMAGE",),
"mode": (("horizontal", "vertical"), {"default": "horizontal"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "join"
CATEGORY = "ali1234/image"
def join(self, images, mode):
n, h, w, c = images.shape
image = None
if mode == "vertical":
# for vertical we can just reshape
image = images.reshape(1, n * h, w, c)
elif mode == "horizontal":
# for horizontal we have to swap axes
image = torch.transpose(torch.transpose(images, 1, 2).reshape(1, n * w, h, c), 1, 2)
return (image,)
@register_node
class SelectImageBatch:
"""Selects one image from an image batch."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"select": ("INT", {"default": 0, "min": 0, "max": 99999, "step": 1}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "select"
CATEGORY = "ali1234/image"
def select(self, images, select):
n, h, w, c = images.shape
if select >= n:
select = n - 1
return (images[select].reshape(1, h, w, c),)