diff --git a/.gitignore b/.gitignore index 281d761..b5a4568 100644 --- a/.gitignore +++ b/.gitignore @@ -3,3 +3,4 @@ /luts/*.CUBE /fonts/*.ttf /fonts/*.otf +!/fonts/ShareTechMono-Regular.ttf \ No newline at end of file diff --git a/fonts/ShareTechMono-Regular.ttf b/fonts/ShareTechMono-Regular.ttf new file mode 100644 index 0000000..0ae0b19 Binary files /dev/null and b/fonts/ShareTechMono-Regular.ttf differ diff --git a/sampling.py b/sampling.py index ac33f31..89f8b86 100644 --- a/sampling.py +++ b/sampling.py @@ -1,8 +1,11 @@ +import os import comfy.samplers import comfy.sample import torch from nodes import common_ksampler -from .utils import expand_mask +from .utils import expand_mask, FONTS_DIR, parse_string_to_list +import torchvision.transforms.v2 as T +import torch.nn.functional as F class KSamplerVariationsWithNoise: @classmethod @@ -167,14 +170,249 @@ class InjectLatentNoise: return (noise_latent, ) +class FluxSamplerParams: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ("MODEL", ), + "conditioning": ("CONDITIONING", ), + "latent_image": ("LATENT", ), + + "noise": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "?" }), + "sampler": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "ipndm" }), + "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" }), + "denoise": ("STRING", { "multiline": False, "dynamicPrompts": False, "default": "1.0" }), + }} + + RETURN_TYPES = ("LATENT","SAMPLER_PARAMS") + RETURN_NAMES = ("latent", "params") + FUNCTION = "execute" + CATEGORY = "essentials/sampling" + + def execute(self, model, conditioning, latent_image, noise, sampler, scheduler, steps, guidance, max_shift, base_shift, split_sigmas, denoise): + import random + 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 node_helpers import conditioning_set_values + + noise = noise.replace("\n", ",").split(",") + noise = [random.randint(0, 999999) if "?" in n else int(n) for n in noise] + if not noise: + noise = [random.randint(0, 999999)] + + if sampler == '*': + sampler = comfy.samplers.KSampler.SAMPLERS + elif sampler.startswith("!"): + sampler = sampler.replace("\n", ",").split(",") + sampler = [s.strip("! ") for s in sampler] + sampler = [s for s in comfy.samplers.KSampler.SAMPLERS if s not in sampler] + else: + sampler = sampler.replace("\n", ",").split(",") + sampler = [s.strip() for s in sampler if s.strip() in comfy.samplers.KSampler.SAMPLERS] + if not sampler: + sampler = ['ipndm'] + + if scheduler == '*': + scheduler = comfy.samplers.KSampler.SCHEDULERS + elif scheduler.startswith("!"): + scheduler = scheduler.replace("\n", ",").split(",") + scheduler = [s.strip("! ") for s in scheduler] + scheduler = [s for s in comfy.samplers.KSampler.SCHEDULERS if s not in scheduler] + else: + scheduler = scheduler.replace("\n", ",").split(",") + scheduler = [s.strip() for s in scheduler] + scheduler = [s for s in scheduler if s in comfy.samplers.KSampler.SCHEDULERS] + if not scheduler: + scheduler = ['simple'] + + steps = steps.replace("\n", ",").split(",") + steps = [int(s) for s in steps] + if not steps: + steps = [20] + + denoise = parse_string_to_list(denoise) + if not denoise: + denoise = [1.0] + + 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] + + base_shift = parse_string_to_list(base_shift) + if not base_shift: + base_shift = [0.5] + + split_sigmas = parse_string_to_list(split_sigmas) + if not split_sigmas: + split_sigmas = [1.0] + + out_latent = None + out_params = [] + + basicschedueler = BasicScheduler() + basicguider = BasicGuider() + samplercustomadvanced = SamplerCustomAdvanced() + latentbatch = LatentBatch() + modelsamplingflux = ModelSamplingFlux() + 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}) + + if out_latent is None: + out_latent = latent + else: + out_latent = latentbatch.batch(out_latent, latent)[0] + + return (out_latent, out_params) + +class PlotParameters: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "images": ("IMAGE", ), + "params": ("SAMPLER_PARAMS", ), + "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 }), + }} + + RETURN_TYPES = ("IMAGE", ) + FUNCTION = "execute" + CATEGORY = "essentials/sampling" + + def execute(self, images, params, order_by, cols_value, cols_num): + from PIL import Image, ImageDraw, ImageFont + import math + + if images.shape[0] != len(params): + raise ValueError("Number of images and number of parameters do not match.") + + if order_by != "none": + if cols_value != "none" and cols_num < 1: + cols_num = len(set(p[cols_value] for p in params)) + sorted_params = sorted(params, key=lambda x: x[order_by]) + indices = [params.index(item) for item in sorted_params] + params = sorted_params + images = images[torch.tensor(indices)] + + width = images.shape[2] + out_image = None + + 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 + + for (image, param) in zip(images, params): + 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) + text_image = Image.new('RGB', (width, text_height), color=(0, 0, 0, 0)) + + 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).unsqueeze(0).permute([0,2,3,1]).to(image.device) + image = torch.cat([image.unsqueeze(0), text_image], 1) + + if out_image is None: + out_image = image + else: + out_image = torch.cat([out_image, image], 0) + + 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 = min(mosaic_columns, out_image.shape[0]) + b, h, w, c = out_image.shape + rows = math.ceil(b / cols) + + # Pad the tensor if necessary + if b % cols != 0: + padding = cols - (b % cols) + out_image = F.pad(out_image, (0, 0, 0, 0, 0, 0, 0, padding)) + b = out_image.shape[0] + + # Reshape and transpose + out_image = out_image.reshape(rows, cols, h, w, c) + out_image = out_image.permute(0, 2, 1, 3, 4) + out_image = out_image.reshape(rows * h, cols * w, c).unsqueeze(0) + + """ + width = out_image.shape[2] + # add the title and notes on top + if title and export_labels: + title_font = ImageFont.truetype(os.path.join(FONTS_DIR, 'ShareTechMono-Regular.ttf'), 48) + title_width = title_font.getbbox(title)[2] + title_padding = 6 + title_line_height = title_font.getmask(title).getbbox()[3] + title_font.getmetrics()[1] + title_padding*2 + title_text_height = title_line_height + title_text_image = Image.new('RGB', (width, title_text_height), color=(0, 0, 0, 0)) + + draw = ImageDraw.Draw(title_text_image) + draw.text((width//2 - title_width//2, title_padding), title, font=title_font, fill=(255, 255, 255)) + + title_text_image = T.ToTensor()(title_text_image).unsqueeze(0).permute([0,2,3,1]).to(out_image.device) + out_image = torch.cat([title_text_image, out_image], 1) + """ + + return (out_image, ) + SAMPLING_CLASS_MAPPINGS = { "KSamplerVariationsStochastic+": KSamplerVariationsStochastic, "KSamplerVariationsWithNoise+": KSamplerVariationsWithNoise, "InjectLatentNoise+": InjectLatentNoise, + "FluxSamplerParams+": FluxSamplerParams, + "PlotParameters+": PlotParameters, } SAMPLING_NAME_MAPPINGS = { "KSamplerVariationsStochastic+": "🔧 KSampler Stochastic Variations", "KSamplerVariationsWithNoise+": "🔧 KSampler Variations with Noise Injection", - "InjectLatentNoise+": "🔧 Inject Latent Noise" + "InjectLatentNoise+": "🔧 Inject Latent Noise", + "FluxSamplerParams+": "🔧 Flux Sampler Parameters", + "PlotParameters+": "🔧 Plot Sampler Parameters", } \ No newline at end of file diff --git a/text.py b/text.py index 132ce03..643a1f3 100644 --- a/text.py +++ b/text.py @@ -2,8 +2,8 @@ import os import torch from nodes import MAX_RESOLUTION import torchvision.transforms.v2 as T +from .utils import FONTS_DIR -FONTS_DIR = os.path.join(os.path.dirname(os.path.realpath(__file__)), "fonts") class DrawText: @classmethod def INPUT_TYPES(s): diff --git a/utils.py b/utils.py index 790006c..fa36e31 100644 --- a/utils.py +++ b/utils.py @@ -1,6 +1,10 @@ import torch import numpy as np import scipy +import os +import re + +FONTS_DIR = os.path.join(os.path.dirname(os.path.realpath(__file__)), "fonts") # from https://github.com/pythongosssss/ComfyUI-Custom-Scripts class AnyType(str): @@ -37,3 +41,41 @@ def expand_mask(mask, expand, tapered_corners): out.append(output) return torch.stack(out, dim=0) + +def parse_string_to_list(s): + elements = s.split(',') + result = [] + + def parse_number(s): + try: + if '.' in s: + return float(s) + else: + return int(s) + except ValueError: + return 0 + + def decimal_places(s): + if '.' in s: + return len(s.split('.')[1]) + return 0 + + for element in elements: + element = element.strip() + if '...' in element: + start, rest = element.split('...') + end, step = rest.split('+') + decimals = decimal_places(step) + start = parse_number(start) + end = parse_number(end) + step = parse_number(step) + current = start + if (start > end and step > 0) or (start < end and step < 0): + step = -step + while current <= end: + result.append(round(current, decimals)) + current += step + else: + result.append(round(parse_number(element), decimal_places(element))) + + return result \ No newline at end of file