improve plotparameter node

This commit is contained in:
cubiq
2024-08-20 17:21:32 +02:00
parent 15ca881d4b
commit 635322c1e4
+45 -20
View File
@@ -429,13 +429,14 @@ class PlotParameters:
"cols_value": (["none", "time", "seed", "steps", "denoise", "sampler", "scheduler", "guidance", "max_shift", "base_shift"], ),
"cols_num": ("INT", {"default": -1, "min": -1, "max": 1024 }),
"add_prompt": (["false", "true", "excerpt"], ),
"add_params": (["false", "true", "changes only"], {"default": "true"}),
}}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "execute"
CATEGORY = "essentials/sampling"
def execute(self, images, params, order_by, cols_value, cols_num, add_prompt):
def execute(self, images, params, order_by, cols_value, cols_num, add_prompt, add_params):
from PIL import Image, ImageDraw, ImageFont
import math
import textwrap
@@ -443,15 +444,17 @@ class PlotParameters:
if images.shape[0] != len(params):
raise ValueError("Number of images and number of parameters do not match.")
_params = params.copy()
if order_by != "none":
sorted_params = sorted(params, key=lambda x: x[order_by])
indices = [params.index(item) for item in sorted_params]
sorted_params = sorted(_params, key=lambda x: x[order_by])
indices = [_params.index(item) for item in sorted_params]
images = images[torch.tensor(indices)]
params = sorted_params
_params = sorted_params
if cols_value != "none" and cols_num > -1:
groups = {}
for p in params:
for p in _params:
value = p[cols_value]
if value not in groups:
groups[value] = []
@@ -463,9 +466,9 @@ class PlotParameters:
for g in zip(*groups):
sorted_params.extend(g)
indices = [params.index(item) for item in sorted_params]
indices = [_params.index(item) for item in sorted_params]
images = images[torch.tensor(indices)]
params = sorted_params
_params = sorted_params
elif cols_num == 0:
cols_num = int(math.sqrt(images.shape[0]))
cols_num = max(1, min(cols_num, 1024))
@@ -478,22 +481,44 @@ class PlotParameters:
line_height = font.getmask('Q').getbbox()[3] + font.getmetrics()[1] + text_padding*2
char_width = font.getbbox('M')[2]+1 # using monospace font
for (image, param) in zip(images, params):
if add_params == "changes only":
value_tracker = {}
for p in _params:
for key, value in p.items():
if key != "time":
if key not in value_tracker:
value_tracker[key] = set()
value_tracker[key].add(value)
changing_keys = {key for key, values in value_tracker.items() if len(values) > 1 or key == "prompt"}
result = []
for p in _params:
changing_params = {key: value for key, value in p.items() if key in changing_keys}
result.append(changing_params)
_params = result
for (image, param) in zip(images, _params):
image = image.permute(2, 0, 1)
text = f"time: {param['time']:.2f}s, seed: {param['seed']}, steps: {param['steps']}, size: {param['width']}×{param['height']}\ndenoise: {param['denoise']}, sampler: {param['sampler']}, sched: {param['scheduler']}\nguidance: {param['guidance']}, max/base shift: {param['max_shift']}/{param['base_shift']}"
lines = text.split("\n")
text_height = line_height * len(lines)
text_image = Image.new('RGB', (width, text_height), color=(0, 0, 0))
if add_params != "false":
if add_params == "changes only":
text = "\n".join([f"{key}: {value}" for key, value in param.items() if key != "prompt"])
else:
text = f"time: {param['time']:.2f}s, seed: {param['seed']}, steps: {param['steps']}, size: {param['width']}×{param['height']}\ndenoise: {param['denoise']}, sampler: {param['sampler']}, sched: {param['scheduler']}\nguidance: {param['guidance']}, max/base shift: {param['max_shift']}/{param['base_shift']}"
for i, line in enumerate(lines):
draw = ImageDraw.Draw(text_image)
draw.text((text_padding, i * line_height + text_padding), line, font=font, fill=(255, 255, 255))
text_image = T.ToTensor()(text_image).to(image.device)
image = torch.cat([image, text_image], 1)
lines = text.split("\n")
text_height = line_height * len(lines)
text_image = Image.new('RGB', (width, text_height), color=(0, 0, 0))
if param['prompt'] and add_prompt != "false":
for i, line in enumerate(lines):
draw = ImageDraw.Draw(text_image)
draw.text((text_padding, i * line_height + text_padding), line, font=font, fill=(255, 255, 255))
text_image = T.ToTensor()(text_image).to(image.device)
image = torch.cat([image, text_image], 1)
if 'prompt' in param and param['prompt'] and add_prompt != "false":
prompt = param['prompt']
if add_prompt == "excerpt":
prompt = " ".join(param['prompt'].split()[:64])
@@ -516,7 +541,7 @@ class PlotParameters:
out_image.append(image)
# ensure all images have the same height
if add_prompt != "false":
if add_prompt != "false" or add_params == "changes only":
max_height = max([image.shape[1] for image in out_image])
out_image = [F.pad(image, (0, 0, 0, max_height - image.shape[1])) for image in out_image]