implement

This commit is contained in:
hnmr293
2023-03-30 02:09:15 +09:00
parent 9ff9e89c46
commit b6be594939
11 changed files with 1024 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
.vscode/
__pycache__/
+74
View File
@@ -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,
}
+11
View File
@@ -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), {}]], )
+116
View File
@@ -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
+299
View File
@@ -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,)
+43
View File
@@ -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,)
+19
View File
@@ -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}, )
+317
View File
@@ -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
View File
@@ -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,)
+20
View File
@@ -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,)
+84
View File
@@ -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}, )