diff --git a/sampling.py b/sampling.py index 89f8b86..d17989a 100644 --- a/sampling.py +++ b/sampling.py @@ -2,7 +2,7 @@ import os import comfy.samplers import comfy.sample import torch -from nodes import common_ksampler +from nodes import common_ksampler, CLIPTextEncode from .utils import expand_mask, FONTS_DIR, parse_string_to_list import torchvision.transforms.v2 as T import torch.nn.functional as F @@ -170,6 +170,39 @@ class InjectLatentNoise: return (noise_latent, ) +class TextEncodeForSamplerParams: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True, "dynamicPrompts": True, "default": "Separate prompts with at least three dashes\n---\nLike so"}), + "clip": ("CLIP", ) + }} + + RETURN_TYPES = ("CONDITIONING", ) + FUNCTION = "execute" + CATEGORY = "essentials/sampling" + + def execute(self, text, clip): + import re + output_text = [] + output_encoded = [] + text = re.sub(r'[-*=~]{4,}\n', '---\n', text) + text = text.split("---\n") + + for t in text: + t = t.strip() + if t: + output_text.append(t) + output_encoded.append(CLIPTextEncode().encode(clip, t)[0]) + + #if len(output_encoded) == 1: + # output = output_encoded[0] + #else: + output = {"text": output_text, "encoded": output_encoded} + + return (output, ) + class FluxSamplerParams: @classmethod def INPUT_TYPES(s): @@ -183,9 +216,9 @@ class FluxSamplerParams: "scheduler": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "simple" }), "steps": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "20" }), "guidance": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "3.5" }), - "max_shift": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "1.15" }), - "base_shift": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "0.5" }), - "split_sigmas": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "1.0" }), + "max_shift": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "" }), + "base_shift": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "" }), + "split_sigmas": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "" }), "denoise": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "1.0" }), }} @@ -199,9 +232,11 @@ class FluxSamplerParams: import time from comfy_extras.nodes_custom_sampler import Noise_RandomNoise, BasicScheduler, BasicGuider, SamplerCustomAdvanced, SplitSigmasDenoise from comfy_extras.nodes_latent import LatentBatch - from comfy_extras.nodes_model_advanced import ModelSamplingFlux + from comfy_extras.nodes_model_advanced import ModelSamplingFlux, ModelSamplingAuraFlow from node_helpers import conditioning_set_values + is_schnell = model.model.model_type == comfy.model_base.ModelType.FLOW + noise = noise.replace("\n", ",").split(",") noise = [random.randint(0, 999999) if "?" in n else int(n) for n in noise] if not noise: @@ -232,30 +267,42 @@ class FluxSamplerParams: if not scheduler: scheduler = ['simple'] + if steps == "": + if is_schnell: + steps = "4" + else: + steps = "20" steps = steps.replace("\n", ",").split(",") steps = [int(s) for s in steps] - if not steps: - steps = [20] + denoise = "1.0" if denoise == "" else denoise denoise = parse_string_to_list(denoise) - if not denoise: - denoise = [1.0] + guidance = "3.5" if guidance == "" else guidance guidance = parse_string_to_list(guidance) - if not guidance: - guidance = [3.5] - max_shift = parse_string_to_list(max_shift) - if not max_shift: - max_shift = [1.15] + if not is_schnell: + max_shift = "1.15" if max_shift == "" else max_shift + max_shift = parse_string_to_list(max_shift) + else: + max_shift = [0] + if base_shift == "": + if is_schnell: + base_shift = "1.0" + else: + base_shift = "0.5" base_shift = parse_string_to_list(base_shift) - if not base_shift: - base_shift = [0.5] + split_sigmas = "1.0" if split_sigmas == "" else split_sigmas split_sigmas = parse_string_to_list(split_sigmas) - if not split_sigmas: - split_sigmas = [1.0] + + cond_text = None + if isinstance(conditioning, dict) and "encoded" in conditioning: + cond_text = conditioning["text"] + cond_encoded = conditioning["encoded"] + else: + cond_encoded = [conditioning] out_latent = None out_params = [] @@ -264,45 +311,52 @@ class FluxSamplerParams: basicguider = BasicGuider() samplercustomadvanced = SamplerCustomAdvanced() latentbatch = LatentBatch() - modelsamplingflux = ModelSamplingFlux() + modelsamplingflux = ModelSamplingFlux() if not is_schnell else ModelSamplingAuraFlow() splitsigmadenoise = SplitSigmasDenoise() width = latent_image["samples"].shape[3]*8 height = latent_image["samples"].shape[2]*8 - for n in noise: - randnoise = Noise_RandomNoise(n) - for ms in max_shift: - for bs in base_shift: - work_model = modelsamplingflux.patch(model, ms, bs, width, height)[0] - for g in guidance: - cond = conditioning_set_values(conditioning, {"guidance": g}) - guider = basicguider.get_guider(work_model, cond)[0] - for s in sampler: - samplerobj = comfy.samplers.sampler_object(s) - for sc in scheduler: - for st in steps: - for d in denoise: - sigmas = basicschedueler.get_sigmas(work_model, sc, st, d)[0] - for ss in split_sigmas: - sigmas = splitsigmadenoise.get_sigmas(sigmas, ss)[1] - start_time = time.time() - latent = samplercustomadvanced.sample(randnoise, guider, samplerobj, sigmas, latent_image)[1] - elapsed_time = time.time() - start_time - out_params.append({"time": elapsed_time, - "seed": n, - "sampler": s, - "scheduler": sc, - "steps": st, - "guidance": g, - "max_shift": ms, - "base_shift": bs, - "denoise": d, - "split_sigmas": ss}) + for i in range(len(cond_encoded)): + conditioning = cond_encoded[i] + ct = cond_text[i] if cond_text else None + for n in noise: + randnoise = Noise_RandomNoise(n) + for ms in max_shift: + for bs in base_shift: + if is_schnell: + work_model = modelsamplingflux.patch_aura(model, bs)[0] + else: + work_model = modelsamplingflux.patch(model, ms, bs, width, height)[0] + for g in guidance: + cond = conditioning_set_values(conditioning, {"guidance": g}) + guider = basicguider.get_guider(work_model, cond)[0] + for s in sampler: + samplerobj = comfy.samplers.sampler_object(s) + for sc in scheduler: + for st in steps: + for d in denoise: + sigmas = basicschedueler.get_sigmas(work_model, sc, st, d)[0] + for ss in split_sigmas: + sigmas = splitsigmadenoise.get_sigmas(sigmas, ss)[1] + start_time = time.time() + latent = samplercustomadvanced.sample(randnoise, guider, samplerobj, sigmas, latent_image)[1] + elapsed_time = time.time() - start_time + out_params.append({"time": elapsed_time, + "seed": n, + "sampler": s, + "scheduler": sc, + "steps": st, + "guidance": g, + "max_shift": ms, + "base_shift": bs, + "denoise": d, + "split_sigmas": ss, + "prompt": ct}) - if out_latent is None: - out_latent = latent - else: - out_latent = latentbatch.batch(out_latent, latent)[0] + if out_latent is None: + out_latent = latent + else: + out_latent = latentbatch.batch(out_latent, latent)[0] return (out_latent, out_params) @@ -315,15 +369,17 @@ class PlotParameters: "order_by": (["none", "time", "seed", "steps", "denoise", "sampler", "scheduler"], ), "cols_value": (["none", "time", "seed", "steps", "denoise", "sampler", "scheduler"], ), "cols_num": ("INT", {"default": -1, "min": -1, "max": 1024 }), + "add_prompt": (["false", "true", "excerpt"], ), }} RETURN_TYPES = ("IMAGE", ) FUNCTION = "execute" CATEGORY = "essentials/sampling" - def execute(self, images, params, order_by, cols_value, cols_num): + def execute(self, images, params, order_by, cols_value, cols_num, add_prompt): from PIL import Image, ImageDraw, ImageFont import math + import textwrap if images.shape[0] != len(params): raise ValueError("Number of images and number of parameters do not match.") @@ -337,13 +393,16 @@ class PlotParameters: images = images[torch.tensor(indices)] width = images.shape[2] - out_image = None + out_image = [] font = ImageFont.truetype(os.path.join(FONTS_DIR, 'ShareTechMono-Regular.ttf'), min(48, int(32*(width/1024)))) text_padding = 3 - line_height = font.getmask('WwMmQqlL1234567890').getbbox()[3] + font.getmetrics()[1] + text_padding*2 + 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): + image = image.permute(2, 0, 1) + text = f"time: {param['time']:.2f}s, seed: {param['seed']}, steps: {param['steps']}, denoise: {param['denoise']}\nsampler: {param['sampler']}, sched: {param['scheduler']}, sigmas at: {param['split_sigmas']}\nguidance: {param['guidance']}, max/base shift: {param['max_shift']}/{param['base_shift']}" lines = text.split("\n") text_height = line_height * len(lines) @@ -353,20 +412,42 @@ class PlotParameters: 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).unsqueeze(0).permute([0,2,3,1]).to(image.device) - image = torch.cat([image.unsqueeze(0), text_image], 1) + text_image = T.ToTensor()(text_image).to(image.device) + image = torch.cat([image, text_image], 1) - if out_image is None: - out_image = image - else: - out_image = torch.cat([out_image, image], 0) + if param['prompt'] and add_prompt != "false": + prompt = param['prompt'] + if add_prompt == "excerpt": + prompt = " ".join(param['prompt'].split()[:64]) + prompt += "..." + + cols = math.ceil(width / char_width) + prompt_lines = textwrap.wrap(prompt, width=cols) + prompt_height = line_height * len(prompt_lines) + prompt_image = Image.new('RGB', (width, prompt_height), color=(0, 0, 0, 0)) + + for i, line in enumerate(prompt_lines): + draw = ImageDraw.Draw(prompt_image) + draw.text((text_padding, i * line_height + text_padding), line, font=font, fill=(255, 255, 255)) + + prompt_image = T.ToTensor()(prompt_image).to(image.device) + image = torch.cat([image, prompt_image], 1) + + out_image.append(image) + + # ensure all images have the same height + if add_prompt != "false": + 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] + + out_image = torch.stack(out_image, 0).permute(0, 2, 3, 1) if cols_num > -1: if cols_num == 0: - mosaic_columns = int(math.sqrt(out_image.shape[0])) - mosaic_columns = max(1, min(mosaic_columns, 1024)) + cols_num = int(math.sqrt(out_image.shape[0])) + cols_num = max(1, min(cols_num, 1024)) - cols = min(mosaic_columns, out_image.shape[0]) + cols = min(cols_num, out_image.shape[0]) b, h, w, c = out_image.shape rows = math.ceil(b / cols) @@ -407,6 +488,7 @@ SAMPLING_CLASS_MAPPINGS = { "InjectLatentNoise+": InjectLatentNoise, "FluxSamplerParams+": FluxSamplerParams, "PlotParameters+": PlotParameters, + "TextEncodeForSamplerParams+": TextEncodeForSamplerParams, } SAMPLING_NAME_MAPPINGS = { @@ -415,4 +497,5 @@ SAMPLING_NAME_MAPPINGS = { "InjectLatentNoise+": "🔧 Inject Latent Noise", "FluxSamplerParams+": "🔧 Flux Sampler Parameters", "PlotParameters+": "🔧 Plot Sampler Parameters", + "TextEncodeForSamplerParams+": "🔧Text Encode for Sampler Params", } \ No newline at end of file