implement
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
.vscode/
|
||||
__pycache__/
|
||||
+74
@@ -0,0 +1,74 @@
|
||||
from .text import Text
|
||||
from .number import Integer, Float
|
||||
from .sampler import SamplerName, SchedulerName
|
||||
from .cliptextencode2 import CLIPTextEncode2
|
||||
from .randomlatent import RandomLatentImage
|
||||
from .vae import VAEDecodeBatched, VAEEncodeBatched
|
||||
from .sample import KSamplerSetting, KSamplerOverrided, KSamplerXYZ
|
||||
from .model import StateDictLoader, Dict2Model, StateDictMerger, StateDictMergerBlockWeighted
|
||||
from .image import GridImage
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# basic nodes
|
||||
|
||||
## text output
|
||||
'Text': Text,
|
||||
|
||||
## integer output
|
||||
'Integer': Integer,
|
||||
|
||||
## float output
|
||||
'Float': Float,
|
||||
|
||||
## sampler selection
|
||||
'SamplerName': SamplerName,
|
||||
|
||||
## scheduler selection
|
||||
'SchedulerName': SchedulerName,
|
||||
|
||||
# conditioning
|
||||
|
||||
## same as CLIPTextEncode, but the prompt and CLIP are external inputs
|
||||
'CLIPTextEncode2': CLIPTextEncode2,
|
||||
|
||||
# latent
|
||||
|
||||
'RandomLatentImage': RandomLatentImage,
|
||||
|
||||
## pass latents to VAE separately
|
||||
'VAEDecodeBatched': VAEDecodeBatched,
|
||||
'VAEEncodeBatched': VAEEncodeBatched,
|
||||
|
||||
# sampling
|
||||
|
||||
## put parameters for sampler into a dict
|
||||
'KSamplerSetting': KSamplerSetting,
|
||||
|
||||
## KSampler with a dict as default setting
|
||||
'KSamplerOverrided': KSamplerOverrided,
|
||||
|
||||
## XYZ plotting
|
||||
'KSamplerXYZ': KSamplerXYZ,
|
||||
|
||||
# loader
|
||||
|
||||
## loads state_dict of the specified checkpoint and returns it
|
||||
'StateDictLoader': StateDictLoader,
|
||||
|
||||
# model
|
||||
|
||||
## creates model from state_dict loaded by `StateDictLoader`
|
||||
'Dict2Model': Dict2Model,
|
||||
|
||||
## merge two (weighted sum) or three (add difference) state_dict
|
||||
'StateDictMerger': StateDictMerger,
|
||||
|
||||
## merge block weighted
|
||||
## weights should be specified by Text
|
||||
'StateDictMergerBlockWeighted': StateDictMergerBlockWeighted,
|
||||
|
||||
# image
|
||||
|
||||
## rearrange images to single image with specified columns and gap
|
||||
'GridImage': GridImage,
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
class CLIPTextEncode2:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {"clip": ("CLIP", ), "text": ("TEXT",)}}
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "encode"
|
||||
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
def encode(self, clip, text):
|
||||
return ([[clip.encode(text), {}]], )
|
||||
@@ -0,0 +1,116 @@
|
||||
import os
|
||||
import math
|
||||
import json
|
||||
from typing import List
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
from nodes import SaveImage
|
||||
|
||||
class GridImage(SaveImage):
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'images': ('IMAGE',),
|
||||
'filename_prefix': ('STRING', {'default': 'ComfyUI-Grid'}),
|
||||
'x': ('INT', {
|
||||
'default': 1,
|
||||
'min': 1,
|
||||
'max': 64,
|
||||
'step': 1
|
||||
}),
|
||||
'gap': ('INT', {
|
||||
'default': 0,
|
||||
'min': 0,
|
||||
'max': 32,
|
||||
'step': 1
|
||||
}),
|
||||
},
|
||||
'hidden': {
|
||||
'prompt': 'PROMPT',
|
||||
'extra_pnginfo': 'EXTRA_PNGINFO'
|
||||
},
|
||||
}
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'image'
|
||||
|
||||
def execute(self, images: List[torch.Tensor], filename_prefix: str = 'ComfyUI-Grid', x: int = 1, gap: int = 0, prompt=None, extra_pnginfo=None):
|
||||
y = max([math.ceil(len(images) / x), 1])
|
||||
|
||||
def map_filename(filename):
|
||||
prefix_len = len(os.path.basename(filename_prefix))
|
||||
prefix = filename[:prefix_len + 1]
|
||||
try:
|
||||
digits = int(filename[prefix_len + 1:].split('_')[0])
|
||||
except:
|
||||
digits = 0
|
||||
return (digits, prefix)
|
||||
|
||||
subfolder = os.path.dirname(os.path.normpath(filename_prefix))
|
||||
filename = os.path.basename(os.path.normpath(filename_prefix))
|
||||
|
||||
full_output_folder = os.path.join(self.output_dir, subfolder)
|
||||
|
||||
if os.path.commonpath((self.output_dir, os.path.realpath(full_output_folder))) != self.output_dir:
|
||||
print("Saving image outside the output folder is not allowed.")
|
||||
return {}
|
||||
|
||||
try:
|
||||
counter = max(filter(lambda a: a[1][:-1] == filename and a[1][-1] == "_", map(map_filename, os.listdir(full_output_folder))))[0] + 1
|
||||
except ValueError:
|
||||
counter = 1
|
||||
except FileNotFoundError:
|
||||
os.makedirs(full_output_folder, exist_ok=True)
|
||||
counter = 1
|
||||
|
||||
if not os.path.exists(self.output_dir):
|
||||
os.makedirs(self.output_dir)
|
||||
|
||||
results = list()
|
||||
|
||||
canvas = self.grid_image(images, x, y, gap)
|
||||
metadata = PngInfo()
|
||||
if prompt is not None:
|
||||
metadata.add_text("prompt", json.dumps(prompt))
|
||||
if extra_pnginfo is not None:
|
||||
for x in extra_pnginfo:
|
||||
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
|
||||
|
||||
file = f"{filename}_{counter:05}_.png"
|
||||
canvas.save(os.path.join(full_output_folder, file), pnginfo=metadata, compress_level=4)
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": 'output'
|
||||
})
|
||||
counter += 1
|
||||
|
||||
return { "ui": { "images": results } }
|
||||
|
||||
def grid_image(self, images: List[torch.Tensor], x: int, y: int, gap: int):
|
||||
width, height, _ = images[0].shape
|
||||
canvas = Image.new('RGB', (x*(width+gap)-gap, y*(height+gap)-gap), color='black')
|
||||
|
||||
for Y in range(y):
|
||||
for X in range(x):
|
||||
idx = Y * x + X
|
||||
if len(images) <= idx:
|
||||
return canvas
|
||||
|
||||
image = images[idx]
|
||||
|
||||
i = 255. * image.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||
|
||||
canvas.paste(img, (X*(width+gap), Y*(height+gap)))
|
||||
|
||||
return canvas
|
||||
@@ -0,0 +1,299 @@
|
||||
import re
|
||||
from typing import Dict, Union, List, Callable
|
||||
import torch
|
||||
import tqdm
|
||||
import folder_paths
|
||||
from comfy.sd import load_torch_file, load_checkpoint
|
||||
|
||||
#class ModelName:
|
||||
#
|
||||
# @classmethod
|
||||
# def INPUT_TYPES(cls):
|
||||
# return {
|
||||
# 'required': {
|
||||
# 'value': (folder_paths.get_filename_list("checkpoints"),)
|
||||
# }
|
||||
# }
|
||||
#
|
||||
# RETURN_TYPES = ('STRING',)
|
||||
#
|
||||
# FUNCTION = 'execute'
|
||||
#
|
||||
# CATEGORY = 'value'
|
||||
#
|
||||
# def execute(self, value: str):
|
||||
# return (value,)
|
||||
|
||||
|
||||
class StateDictLoader:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'ckpt_name': (folder_paths.get_filename_list("checkpoints"), )
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('DICT',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'loaders'
|
||||
|
||||
def execute(self, ckpt_name: str):
|
||||
ckpt_path = folder_paths.get_full_path('checkpoints', ckpt_name)
|
||||
sd = load_torch_file(ckpt_path)
|
||||
return (sd,)
|
||||
|
||||
|
||||
class Dict2Model:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'weights': ('DICT', ),
|
||||
'config_name': (folder_paths.get_filename_list('configs'), ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('MODEL', 'CLIP', 'VAE')
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'model'
|
||||
|
||||
def execute(self, weights: dict, config_name: str):
|
||||
config_path = folder_paths.get_full_path("configs", config_name)
|
||||
|
||||
def load_torch_file_hook(*args, **kwargs):
|
||||
return weights
|
||||
|
||||
import comfy.sd as sd
|
||||
load_torch_file_org = sd.load_torch_file
|
||||
setattr(sd, 'load_torch_file', load_torch_file_hook)
|
||||
|
||||
try:
|
||||
return sd.load_checkpoint(config_path, None, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||
finally:
|
||||
setattr(sd, 'load_torch_file', load_torch_file_org)
|
||||
|
||||
|
||||
def merge(
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
merge_fn: Callable[[str, torch.Tensor, torch.Tensor], torch.Tensor],
|
||||
position_ids: str,
|
||||
half: str,
|
||||
ignore_keys_only_in_B: bool = False,
|
||||
):
|
||||
result = dict()
|
||||
for key in tqdm.tqdm(model_A.keys()):
|
||||
if key not in model_B:
|
||||
print(f' key {key} is found in model_A but not model_B')
|
||||
result[key] = model_A[key]
|
||||
continue
|
||||
|
||||
if key.endswith('.position_ids'):
|
||||
continue
|
||||
|
||||
#print(f' key : {key}')
|
||||
|
||||
val = merge_fn(key, model_A[key], model_B[key])
|
||||
|
||||
if half == 'True':
|
||||
val = val.half()
|
||||
else:
|
||||
val = val.float()
|
||||
|
||||
result[key] = val
|
||||
|
||||
for key in model_B.keys():
|
||||
if not ignore_keys_only_in_B and key not in model_A:
|
||||
if key.endswith('.position_ids'):
|
||||
continue
|
||||
|
||||
print(f' key {key} is not found in model_A but model_B')
|
||||
|
||||
val = model_B[key]
|
||||
if half == 'True':
|
||||
val = val.half()
|
||||
else:
|
||||
val = val.float()
|
||||
|
||||
result[key] = val
|
||||
|
||||
print('position_ids')
|
||||
if position_ids == 'A':
|
||||
position_ids_key = next(x for x in model_A.keys() if '.position_ids' in x)
|
||||
position_ids_val = model_A[position_ids_key]
|
||||
print(f" using model_A's one ({position_ids_val.dtype}: {position_ids_val.shape})")
|
||||
elif position_ids == 'B':
|
||||
position_ids_key = next(x for x in model_B.keys() if '.position_ids' in x)
|
||||
position_ids_val = model_B[position_ids_key]
|
||||
print(f" using model_B's one ({position_ids_val.dtype}: {position_ids_val.shape})")
|
||||
elif position_ids == 'Reset':
|
||||
position_ids_key = next(x for x in model_A.keys() if '.position_ids' in x)
|
||||
position_ids_val = torch.LongTensor(list(range(77))).reshape(model_A[position_ids_key].shape)
|
||||
print(f" reset ({position_ids_val.dtype}: {position_ids_val.shape})")
|
||||
else:
|
||||
raise ValueError('must not happen')
|
||||
result[position_ids_key] = position_ids_val
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def weighted_sum(
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
alpha1: float,
|
||||
alpha2: float,
|
||||
position_ids: str,
|
||||
half: str,
|
||||
):
|
||||
print('merging ...')
|
||||
print('mode: Weighted Sum')
|
||||
|
||||
def merge_fn(key, t1, t2):
|
||||
return alpha1 * t1 + alpha2 * t2
|
||||
|
||||
return merge(model_A, model_B, merge_fn, position_ids, half)
|
||||
|
||||
def add_diff(
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
model_C: Dict[str,torch.Tensor],
|
||||
alpha: float,
|
||||
position_ids: str,
|
||||
half: str,
|
||||
):
|
||||
print('merging ...')
|
||||
print('mode: Add Difference')
|
||||
|
||||
print('X = B - C')
|
||||
def merge_fn1(key, t1, t2):
|
||||
return t1 - t2
|
||||
model_X = merge(model_B, model_C, merge_fn1, 'A', half, ignore_keys_only_in_B=True)
|
||||
|
||||
print('A + alpha*X')
|
||||
def merge_fn2(key, t1, t2):
|
||||
return t1 + t2*alpha
|
||||
result = merge(model_A, model_X, merge_fn2, position_ids, half)
|
||||
|
||||
return result
|
||||
|
||||
re_inp = re.compile(r'\.input_blocks\.(\d+)\.')
|
||||
re_mid = re.compile(r'\.middle_block\.(\d+)\.')
|
||||
re_out = re.compile(r'\.output_blocks\.(\d+)\.')
|
||||
|
||||
def weighted_sum_block(
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
base_alpha: float,
|
||||
alphas: List[float],
|
||||
position_ids: str,
|
||||
half: str,
|
||||
):
|
||||
print('merging ...')
|
||||
print('mode: Block Weighted')
|
||||
|
||||
def index(key: str):
|
||||
if not key.startswith('model.diffusion_model.'):
|
||||
return None
|
||||
if 'time_embed' in key:
|
||||
return 0
|
||||
if '.out.' in key:
|
||||
return 24
|
||||
m = re_inp.search(key)
|
||||
if m: return int(m.group(1))
|
||||
m = re_mid.search(key)
|
||||
if m: return 12 + int(m.group(1))
|
||||
m = re_out.search(key)
|
||||
if m: return 13 + int(m.group(1))
|
||||
return None
|
||||
|
||||
def merge_fn(key, t1, t2):
|
||||
weight_index = index(key)
|
||||
if weight_index is None:
|
||||
alpha = base_alpha
|
||||
elif 25 <= weight_index:
|
||||
raise ValueError('must not happen')
|
||||
else:
|
||||
alpha = alphas[weight_index]
|
||||
#print(key, alpha)
|
||||
return (1.0 - alpha) * t1 + alpha * t2
|
||||
|
||||
return merge(model_A, model_B, merge_fn, position_ids, half)
|
||||
|
||||
|
||||
class StateDictMerger:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'model_A': ('DICT',),
|
||||
'model_B': ('DICT',),
|
||||
'alpha': ('FLOAT', {
|
||||
'default': 0,
|
||||
'min': -1,
|
||||
'max': 2,
|
||||
'step': 0.001
|
||||
}),
|
||||
'position_ids': (['A', 'B', 'Reset'], ),
|
||||
'half': (['True', 'False'], ),
|
||||
},
|
||||
'optional': {
|
||||
'model_C': ('DICT',),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('DICT',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'model'
|
||||
|
||||
def execute(
|
||||
self,
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
alpha: float,
|
||||
position_ids: str,
|
||||
half: str,
|
||||
model_C: Union[Dict[str,torch.Tensor],None] = None,
|
||||
):
|
||||
if model_C is None:
|
||||
result = weighted_sum(model_A, model_B, 1.0 - alpha, alpha, position_ids, half)
|
||||
else:
|
||||
result = add_diff(model_A, model_B, model_C, alpha, position_ids, half)
|
||||
return (result,)
|
||||
|
||||
|
||||
class StateDictMergerBlockWeighted(StateDictMerger):
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
d = StateDictMerger.INPUT_TYPES()
|
||||
d['required']['base_alpha'] = d['required']['alpha']
|
||||
del d['required']['alpha']
|
||||
del d['optional']['model_C']
|
||||
d['required']['alphas'] = ('TEXT',)
|
||||
return d
|
||||
|
||||
def execute(
|
||||
self,
|
||||
model_A: Dict[str,torch.Tensor],
|
||||
model_B: Dict[str,torch.Tensor],
|
||||
position_ids: str,
|
||||
half: str,
|
||||
base_alpha: float,
|
||||
alphas: str,
|
||||
):
|
||||
alphas_ = [float(x.strip()) for x in alphas.split(',')]
|
||||
if len(alphas_) != 25:
|
||||
raise ValueError(f'given {len(alphas_)} values, expected 25.')
|
||||
result = weighted_sum_block(model_A, model_B, base_alpha, alphas_, position_ids, half)
|
||||
|
||||
return (result,)
|
||||
@@ -0,0 +1,43 @@
|
||||
import re
|
||||
|
||||
re_range = re.compile(r"\s*([+-]?\s*\d+)\s*-\s*([+-]?\s*\d+)(?:\s*\(([+-]\d+)\s*\))?\s*")
|
||||
re_range_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*-\s*([+-]?\s*\d+(?:.\d*)?)(?:\s*\(([+-]\d+(?:.\d*)?)\s*\))?\s*")
|
||||
|
||||
class Integer:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'value': ('INT', {})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('Integer',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'value'
|
||||
|
||||
def execute(self, value: int):
|
||||
return (value,)
|
||||
|
||||
|
||||
class Float:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'value': ('FLOAT', {})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('Float',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'value'
|
||||
|
||||
def execute(self, value: float):
|
||||
return (value,)
|
||||
@@ -0,0 +1,19 @@
|
||||
import torch
|
||||
|
||||
class RandomLatentImage:
|
||||
def __init__(self, device="cpu"):
|
||||
self.device = device
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "width": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 64}),
|
||||
"height": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 64}),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 64})}}
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "generate"
|
||||
|
||||
CATEGORY = "latent"
|
||||
|
||||
def generate(self, width, height, batch_size=1):
|
||||
latent = torch.randn(batch_size, 4, height // 8, width // 8)
|
||||
return ({"samples":latent}, )
|
||||
@@ -0,0 +1,317 @@
|
||||
import re
|
||||
from itertools import product
|
||||
from typing import Callable, List, Dict, Any, Union, Iterable
|
||||
import torch
|
||||
import model_management # type: ignore
|
||||
import comfy.samplers
|
||||
from nodes import common_ksampler
|
||||
from comfy.sd import ModelPatcher
|
||||
|
||||
re_int = re.compile(r"\s*([+-]?\s*\d+)\s*")
|
||||
re_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*")
|
||||
re_range = re.compile(r"\s*([+-]?\s*\d+)\s*-\s*([+-]?\s*\d+)(?:\s*\(([+-]\d+)\s*\))?\s*")
|
||||
re_range_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*-\s*([+-]?\s*\d+(?:.\d*)?)(?:\s*\(([+-]\d+(?:.\d*)?)\s*\))?\s*")
|
||||
|
||||
def frange(start, end, step):
|
||||
x = float(start)
|
||||
end = float(end)
|
||||
step = float(step)
|
||||
while x < end:
|
||||
yield x
|
||||
x += step
|
||||
|
||||
def get_noise(seeds: List[int], latent_image: torch.Tensor, disable_noise: bool):
|
||||
noises: List[torch.Tensor] = []
|
||||
latents: List[torch.Tensor] = []
|
||||
|
||||
if latent_image.dim() == 3:
|
||||
latent_image = latent_image.unsqueeze(0) # add batch dim
|
||||
|
||||
if disable_noise:
|
||||
noise_ = torch.zeros([len(seeds)]+list(latent_image.size())[-3:], dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
noises.append(noise_)
|
||||
latents.extend([latent_image] * (len(seeds) // latent_image.shape[0]))
|
||||
else:
|
||||
for s in seeds:
|
||||
noise_ = torch.randn(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=torch.manual_seed(s), device="cpu")
|
||||
noises.append(noise_)
|
||||
latents.append(latent_image)
|
||||
|
||||
return torch.cat(noises), torch.cat(latents)
|
||||
|
||||
def get_cfg(noises: torch.Tensor, latent_image: torch.Tensor, cfgs: List[float]):
|
||||
# batch_size = noises.shape[0] * len(cfgs)
|
||||
ns = [noises] * len(cfgs)
|
||||
lat = [latent_image] * len(cfgs)
|
||||
cf = torch.FloatTensor(cfgs * noises.shape[0])
|
||||
return torch.cat(ns), torch.cat(lat), cf[...,None,None,None]
|
||||
|
||||
def common_ksampler_xyz(
|
||||
model: Union[ModelPatcher,Iterable[ModelPatcher]],
|
||||
seed: Union[int,List[int]],
|
||||
steps: Union[int,List[int]],
|
||||
cfg: Union[float,List[float]],
|
||||
sampler_name: Union[str,List[str]],
|
||||
scheduler: Union[str,List[str]],
|
||||
positive,
|
||||
negative,
|
||||
latent,
|
||||
denoise=1.0,
|
||||
disable_noise=False,
|
||||
start_step=None,
|
||||
last_step=None,
|
||||
force_full_denoise=False
|
||||
):
|
||||
latent_image = latent["samples"]
|
||||
noise_mask = None
|
||||
device = model_management.get_torch_device()
|
||||
|
||||
if not isinstance(model, Iterable):
|
||||
model = (model,)
|
||||
|
||||
if not isinstance(seed, list):
|
||||
seed = [seed]
|
||||
|
||||
if not isinstance(steps, list):
|
||||
steps = [steps]
|
||||
|
||||
if not isinstance(cfg, list):
|
||||
cfg = [cfg]
|
||||
|
||||
if not isinstance(sampler_name, list):
|
||||
sampler_name = [sampler_name]
|
||||
|
||||
if not isinstance(scheduler, list):
|
||||
scheduler = [scheduler]
|
||||
|
||||
noise, latent_image = get_noise(seed, latent_image, disable_noise)
|
||||
noise, latent_image, cfg_ = get_cfg(noise, latent_image, cfg)
|
||||
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent['noise_mask']
|
||||
noise_mask = torch.nn.functional.interpolate(noise_mask[None,None,], size=(noise.shape[2], noise.shape[3]), mode="bilinear")
|
||||
noise_mask = noise_mask.round()
|
||||
noise_mask = torch.cat([noise_mask] * noise.shape[1], dim=1)
|
||||
noise_mask = torch.cat([noise_mask] * noise.shape[0])
|
||||
noise_mask = noise_mask.to(device)
|
||||
|
||||
noise = noise.to(device)
|
||||
latent_image = latent_image.to(device)
|
||||
cfg_ = cfg_.to(device)
|
||||
|
||||
positive_copy = []
|
||||
negative_copy = []
|
||||
|
||||
control_nets = []
|
||||
for p in positive:
|
||||
t = p[0]
|
||||
if t.shape[0] < noise.shape[0]:
|
||||
t = torch.cat([t] * noise.shape[0])
|
||||
t = t.to(device)
|
||||
if 'control' in p[1]:
|
||||
control_nets += [p[1]['control']]
|
||||
positive_copy += [[t] + p[1:]]
|
||||
for n in negative:
|
||||
t = n[0]
|
||||
if t.shape[0] < noise.shape[0]:
|
||||
t = torch.cat([t] * noise.shape[0])
|
||||
t = t.to(device)
|
||||
if 'control' in n[1]:
|
||||
control_nets += [n[1]['control']]
|
||||
negative_copy += [[t] + n[1:]]
|
||||
|
||||
control_net_models = []
|
||||
for x in control_nets:
|
||||
control_net_models += x.get_control_models()
|
||||
model_management.load_controlnet_gpu(control_net_models)
|
||||
|
||||
#samplers: List[comfy.samplers.KSampler] = []
|
||||
samplers: List[Dict[str,Any]] = []
|
||||
for model_, sampler_name_, scheduler_, steps_ in product(model, sampler_name, scheduler, steps):
|
||||
if sampler_name_ not in comfy.samplers.KSampler.SAMPLERS:
|
||||
raise ValueError(f'unknown sampler name: {sampler_name_}')
|
||||
if scheduler_ not in comfy.samplers.KSampler.SCHEDULERS:
|
||||
raise ValueError(f'unknown scheduler name: {scheduler_}')
|
||||
samplers.append(dict(
|
||||
model=model_,
|
||||
steps=steps_,
|
||||
device=device,
|
||||
sampler=sampler_name_,
|
||||
scheduler=scheduler_,
|
||||
denoise=denoise,
|
||||
))
|
||||
|
||||
all_samples: List[torch.Tensor] = []
|
||||
for sampler_args in samplers:
|
||||
model_ = sampler_args['model']
|
||||
model_management.load_model_gpu(model_)
|
||||
sampler_args['model'] = model_.model
|
||||
|
||||
sampler = comfy.samplers.KSampler(**sampler_args)
|
||||
print(f'XYZ sampler={sampler.sampler}/{sampler.scheduler} {sampler.steps}steps')
|
||||
|
||||
samples = sampler.sample(noise, positive_copy, negative_copy, cfg=cfg_, latent_image=latent_image, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise, denoise_mask=noise_mask)
|
||||
samples = samples.cpu()
|
||||
all_samples.append(samples)
|
||||
for c in control_nets:
|
||||
c.cleanup()
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = torch.cat(all_samples)
|
||||
return (out, )
|
||||
|
||||
|
||||
class KSamplerSetting:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'model': ('MODEL',),
|
||||
'seed': ('INT', {'default': 0, 'min': 0, 'max': 0xffffffffffffffff}),
|
||||
'steps': ('INT', {'default': 20, 'min': 1, 'max': 10000}),
|
||||
'cfg': ('FLOAT', {'default': 8.0, 'min': 0.0, 'max': 100.0}),
|
||||
'sampler_name': (comfy.samplers.KSampler.SAMPLERS, ),
|
||||
'scheduler': (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
'positive': ('CONDITIONING', ),
|
||||
'negative': ('CONDITIONING', ),
|
||||
'latent_image': ('LATENT', ),
|
||||
'denoise': ('FLOAT', {'default': 1.0, 'min': 0.0, 'max': 1.0, 'step': 0.01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('DICT',)
|
||||
|
||||
FUNCTION = 'sample'
|
||||
|
||||
CATEGORY = 'sampling'
|
||||
|
||||
def sample(self, **kwargs):
|
||||
return kwargs,
|
||||
|
||||
|
||||
class KSamplerOverrided:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'setting': ('DICT',),
|
||||
},
|
||||
'optional': {
|
||||
'model': ('MODEL',),
|
||||
'seed': ('Integer', {'default': 0, 'min': 0, 'max': 0xffffffffffffffff}),
|
||||
'steps': ('Integer', {'default': 20, 'min': 1, 'max': 10000}),
|
||||
'cfg': ('Float', {'default': 8.0, 'min': 0.0, 'max': 100.0}),
|
||||
'sampler_name': ('SamplerName',),
|
||||
'scheduler': ('SchedulerName', ),
|
||||
'positive': ('CONDITIONING', ),
|
||||
'negative': ('CONDITIONING', ),
|
||||
'latent_image': ('LATENT', ),
|
||||
'denoise': ('Float', {'default': 1.0, 'min': 0.0, 'max': 1.0, 'step': 0.01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('LATENT',)
|
||||
FUNCTION = 'sample'
|
||||
|
||||
CATEGORY = 'sampling'
|
||||
|
||||
def sample(self, setting: dict, **kwargs):
|
||||
if 'latent_image' in setting:
|
||||
setting['latent'] = setting['latent_image']
|
||||
del setting['latent_image']
|
||||
|
||||
setting.update(kwargs)
|
||||
|
||||
return common_ksampler(**setting)
|
||||
|
||||
class KSamplerXYZ:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'setting': ('DICT',),
|
||||
},
|
||||
'optional': {
|
||||
'model': ('MODEL',),
|
||||
'seed': ('TEXT',),
|
||||
'steps': ('TEXT',),
|
||||
'cfg': ('TEXT',),
|
||||
'sampler_name': ('TEXT',),
|
||||
'scheduler': ('TEXT', ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('LATENT',)
|
||||
FUNCTION = 'sample'
|
||||
|
||||
CATEGORY = 'sampling'
|
||||
|
||||
def sample(self, setting: dict, **kwargs):
|
||||
if 'latent_image' in setting:
|
||||
setting['latent'] = setting['latent_image']
|
||||
del setting['latent_image']
|
||||
|
||||
setting = { **setting, **kwargs }
|
||||
|
||||
if isinstance(setting.get('seed', None), str):
|
||||
setting['seed'] = self.parse(setting['seed'], self.parse_int)
|
||||
|
||||
if isinstance(setting.get('steps', None), str):
|
||||
setting['steps'] = self.parse(setting['steps'], self.parse_int)
|
||||
|
||||
if isinstance(setting.get('cfg', None), str):
|
||||
setting['cfg'] = self.parse(setting['cfg'], self.parse_float)
|
||||
|
||||
if isinstance(setting.get('sampler_name', None), str):
|
||||
setting['sampler_name'] = self.parse(setting['sampler_name'], None)
|
||||
if len(setting['sampler_name']) == 1:
|
||||
setting['sampler_name'] = setting['sampler_name'][0]
|
||||
|
||||
if isinstance(setting.get('scheduler', None), str):
|
||||
setting['scheduler'] = self.parse(setting['scheduler'], None)
|
||||
if len(setting['scheduler']) == 1:
|
||||
setting['scheduler'] = setting['scheduler'][0]
|
||||
|
||||
for k, v in setting.items():
|
||||
if k in kwargs and isinstance(v, (list, tuple)):
|
||||
print(f'XYZ {k}: {v}')
|
||||
|
||||
return common_ksampler_xyz(**setting)
|
||||
|
||||
def parse(self, input: str, cont: Union[Callable[[str],Any],None]):
|
||||
vs = [ x.strip() for x in input.split(',') ]
|
||||
if cont is not None:
|
||||
vs = [cont(v) for v in vs ]
|
||||
return vs
|
||||
|
||||
def parse_int(self, input: str):
|
||||
m = re_int.fullmatch(input)
|
||||
if m is not None:
|
||||
return int(m.group(1))
|
||||
|
||||
m = re_range.fullmatch(input)
|
||||
if m is None:
|
||||
raise ValueError(f'failed to process: {input}')
|
||||
|
||||
start, end, step = m.group(1), m.group(2), m.group(3)
|
||||
if step is None:
|
||||
step = 1
|
||||
|
||||
return list(range(int(start), int(end), int(step)))
|
||||
|
||||
def parse_float(self, input: str):
|
||||
m = re_float.fullmatch(input)
|
||||
if m is not None:
|
||||
return float(m.group(1))
|
||||
|
||||
m = re_range_float.fullmatch(input)
|
||||
if m is None:
|
||||
raise ValueError(f'failed to process: {input}')
|
||||
|
||||
start, end, step = m.group(1), m.group(2), m.group(3)
|
||||
if step is None:
|
||||
step = 1.0
|
||||
|
||||
return list(frange(float(start), float(end), float(step)))
|
||||
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
import comfy.samplers
|
||||
|
||||
class SamplerName:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'value': (comfy.samplers.KSampler.SAMPLERS,)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('SamplerName',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'value'
|
||||
|
||||
def execute(self, value: str):
|
||||
return (value,)
|
||||
|
||||
class SchedulerName:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'value': (comfy.samplers.KSampler.SCHEDULERS,)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('SchedulerName',)
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'value'
|
||||
|
||||
def execute(self, value: str):
|
||||
return (value,)
|
||||
@@ -0,0 +1,20 @@
|
||||
class Text:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
'required': {
|
||||
'text': ('STRING', {
|
||||
'multiline': True,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('TEXT',) # currently no widgets can receive 'STRING' inputs...
|
||||
|
||||
FUNCTION = 'execute'
|
||||
|
||||
CATEGORY = 'value'
|
||||
|
||||
def execute(self, text: str):
|
||||
return (text,)
|
||||
@@ -0,0 +1,84 @@
|
||||
import torch
|
||||
from tqdm import trange
|
||||
|
||||
class VAEDecodeBatched:
|
||||
def __init__(self, device="cpu"):
|
||||
self.device = device
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"samples": ("LATENT", ),
|
||||
"vae": ("VAE", ),
|
||||
"batch_size": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 32,
|
||||
"step": 1
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "decode"
|
||||
|
||||
CATEGORY = "latent"
|
||||
|
||||
def decode(self, vae, samples, batch_size: int):
|
||||
s = samples['samples']
|
||||
n = s.shape[0]
|
||||
|
||||
results = []
|
||||
for i in trange(0, n, batch_size):
|
||||
e = min([i+batch_size, n])
|
||||
t = s[i:e, ...]
|
||||
v = vae.decode(t)
|
||||
results.append(v)
|
||||
|
||||
vs = torch.cat(results)
|
||||
return (vs,)
|
||||
|
||||
|
||||
class VAEEncodeBatched:
|
||||
def __init__(self, device="cpu"):
|
||||
self.device = device
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"pixels": ("IMAGE", ),
|
||||
"vae": ("VAE", ),
|
||||
"batch_size": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 32,
|
||||
"step": 1
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "encode"
|
||||
|
||||
CATEGORY = "latent"
|
||||
|
||||
def encode(self, vae, pixels, batch_size: int):
|
||||
n = pixels.shape[0]
|
||||
x = (pixels.shape[1] // 64) * 64
|
||||
y = (pixels.shape[2] // 64) * 64
|
||||
if pixels.shape[1] != x or pixels.shape[2] != y:
|
||||
pixels = pixels[:,:x,:y,:]
|
||||
|
||||
pixels = pixels[:,:,:,:3]
|
||||
|
||||
results = []
|
||||
for i in trange(0, n, batch_size):
|
||||
e = max([i+batch_size, n])
|
||||
t = pixels[i:e, ...]
|
||||
v = vae.encode(t)
|
||||
results.append(v)
|
||||
|
||||
vs = torch.cat(results)
|
||||
return ({"samples":vs}, )
|
||||
Reference in New Issue
Block a user