786 lines
31 KiB
Python
786 lines
31 KiB
Python
from multiprocessing.shared_memory import SharedMemory
|
|
from multiprocessing import Event
|
|
from dataclasses import dataclass
|
|
from threading import Thread
|
|
from comfy import utils
|
|
import numpy.random
|
|
import numpy as np
|
|
import torch
|
|
import cv2
|
|
|
|
from .quilting import generate_texture, generate_texture_parallel
|
|
from .misc.validation_utils import validate_array_shape
|
|
from .guess_block_size import guess_nice_block_size
|
|
from .types import UiCoordData
|
|
|
|
# TODO add nodes where user defines output's height and width instead of scale
|
|
|
|
NODES_CATEGORY = "Bmad/CV/Quilting"
|
|
SEAMLESS_DIRS = ["H", "V", "H & V"] # options for seamless nodes
|
|
|
|
|
|
def get_quilting_shared_input_types():
|
|
return {
|
|
# block size is given in pixels.
|
|
# if inferior to 3, guess_nice_block_size is used instead.
|
|
# if negative, guess_nice_block_size only does freq. analysis (should be considerably faster)
|
|
"block_size": ("INT", {"default": -1, "min": -1, "max": 512, "step": 1}),
|
|
|
|
# the percentage of pixels that overlap between each block_sized block
|
|
"overlap": ("FLOAT", {"default": 1 / 5.0, "min": .1, "max": .9, "step": .01}),
|
|
|
|
# this is a percentage relative to min error when searching for a patch.
|
|
# the ones that within tolerance are potential candidates to be selected.
|
|
# tolerance equal to 1 means a tolerance of 2 times the min error.
|
|
# my interpretation regarding its application is that
|
|
# tolerance can help prevent too much sameness in the texture
|
|
# due to some subset of patches being better at minimizing the error.
|
|
"tolerance": ("FLOAT", {"default": .1, "min": 0, "max": 2, "step": .01}),
|
|
|
|
# 0 -> no parallelization; uses the reference implementation
|
|
# 1 -> 4 jobs used ( divides generation into 4 sections )
|
|
# 2 and above -> the number of jobs per section
|
|
# parallelization lvl also affects output; even if seed is fixed, changing this will change the output
|
|
"parallelization_lvl": ("INT", {"default": 1, "min": 0, "max": 6, "step": 1}),
|
|
|
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
|
|
|
"version": ("INT", {"default": 3, "min": 0, "max": 3}),
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class QuiltingFuncWrapper:
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
block_sizes: list[int]
|
|
overlap: float
|
|
out_h: int
|
|
out_w: int
|
|
tolerance: float
|
|
parallelization_lvl: int
|
|
rng: numpy.random.Generator
|
|
version: int
|
|
jobs_shm_name: str
|
|
|
|
|
|
# region AUX FUNCTIONS
|
|
def waiting_loop(abort_loop_event: Event, pbar: utils.ProgressBar, total_steps, shm_name, ntasks=1):
|
|
"""
|
|
Listens for interrupts and propagates to Problem running using interruption_proxy.
|
|
Updates progress_bar via ticker_proxy, updated within Problem instances.
|
|
|
|
@param abort_loop_event: to be triggered in the main thread once the job's done
|
|
@param pbar: comfyui progress bar to update every 100 milliseconds
|
|
@param total_steps: the total number of patches to quilt
|
|
@param shm_name: shared memory name for a numpy integer array
|
|
the first index is used to stop the processes in case of an interruption
|
|
the remaining indexes store the number of patches places by each process
|
|
"""
|
|
from time import sleep
|
|
from comfy.model_management import processing_interrupted
|
|
shm = SharedMemory(name=shm_name)
|
|
procs_statuses = np.ndarray((1 + ntasks,), dtype=np.dtype('uint32'), buffer=shm.buf)
|
|
while not abort_loop_event.is_set():
|
|
sleep(.1) # pause for 1 second
|
|
if processing_interrupted():
|
|
procs_statuses[0] = 1
|
|
return
|
|
pbar.update_absolute(int(np.sum(procs_statuses[1:ntasks + 1])), total_steps)
|
|
|
|
|
|
def terminate_task(finished_event, jobs_shared_memory, pbt: Thread):
|
|
"""
|
|
1. triggers finished_event so that UI thread can stop looping
|
|
2. check if the process stopped due to an interruption
|
|
3. closes/unlinks shared memory
|
|
4. if process stopped due to interruption executes throw_exception_if_processing_interrupted()
|
|
"""
|
|
coord_jobs_array = np.ndarray((1,), dtype=np.dtype('uint32'), buffer=jobs_shared_memory.buf)
|
|
interrupted = coord_jobs_array[0]
|
|
finished_event.set()
|
|
pbt.join()
|
|
jobs_shared_memory.close()
|
|
jobs_shared_memory.unlink()
|
|
if interrupted:
|
|
from comfy.model_management import throw_exception_if_processing_interrupted
|
|
throw_exception_if_processing_interrupted()
|
|
|
|
|
|
def setup_pbar_quilting(block_sizes: list[int], overlap_percentage: float, out_height, out_width, par_lvl):
|
|
"""
|
|
@param block_sizes: block size for each item in the batch. len(block_sizes) = batch length
|
|
"""
|
|
total_steps: int = 0
|
|
batch_len: int = len(block_sizes)
|
|
|
|
for block_size in block_sizes:
|
|
overlap = overlap_percentage_to_pixels(block_size, overlap_percentage)
|
|
n_rows = int(np.ceil((out_height - block_size) * 1.0 / (block_size - overlap)))
|
|
n_columns = int(np.ceil((out_width - block_size) * 1.0 / (block_size - overlap)))
|
|
total_steps += ((n_rows + 1) * (n_columns + 1) - 1) # ignores first corner/center
|
|
# might be less than the generated when using parallel solution, but not by far, so will leave it like this for now
|
|
|
|
return setup_pbar(total_steps, par_lvl, batch_len)
|
|
|
|
|
|
def setup_pbar_seamless(ori, block_sizes: list[int], overlap_percentage: float, height, width, batch_len):
|
|
"""
|
|
here batch_len may differ from block_sizes due to the number of lookup textures
|
|
"""
|
|
from .make_seamless import get_numb_of_blocks_to_fill_stripe
|
|
total_steps: int = 0
|
|
for i in range(batch_len):
|
|
block_size = block_sizes[min(i, len(block_sizes) - 1)]
|
|
overlap = overlap_percentage_to_pixels(block_size, overlap_percentage)
|
|
match ori:
|
|
case "H":
|
|
total_steps += get_numb_of_blocks_to_fill_stripe(block_size, overlap, width)
|
|
case "V":
|
|
total_steps += get_numb_of_blocks_to_fill_stripe(block_size, overlap, height)
|
|
case _:
|
|
total_steps += (
|
|
2 +
|
|
get_numb_of_blocks_to_fill_stripe(block_size, overlap, height) +
|
|
get_numb_of_blocks_to_fill_stripe(block_size, overlap, width)
|
|
)
|
|
return setup_pbar(total_steps, 0, batch_len)
|
|
|
|
|
|
def setup_pbar_seamless_v2(ori, batch_len):
|
|
# 3 increments per big block + 2 for the "H & V" H Seam patch
|
|
match ori:
|
|
case "H & V":
|
|
total_steps = 2 + 3 * 2
|
|
case _:
|
|
total_steps = 3
|
|
total_steps = total_steps * batch_len
|
|
return setup_pbar(total_steps, 0, batch_len)
|
|
|
|
|
|
def setup_pbar(total_steps, par_lvl, batch_len):
|
|
"""
|
|
@param total_steps: aggregate from all the jobs/batch items; not for a single job or batch item
|
|
"""
|
|
pbar: utils.ProgressBar = utils.ProgressBar(total_steps)
|
|
finished_event = Event()
|
|
|
|
if batch_len > 1:
|
|
n_jobs = batch_len
|
|
n_jobs *= 1 if par_lvl == 0 else 4
|
|
else:
|
|
n_jobs = 1 if par_lvl == 0 else 4
|
|
n_jobs *= par_lvl if par_lvl > 0 else 1
|
|
|
|
size = (1 + n_jobs) * np.dtype('uint32').itemsize
|
|
shm_jobs = SharedMemory(create=True, size=size)
|
|
|
|
t = Thread(target=waiting_loop, args=(finished_event, pbar, total_steps, shm_jobs.name, n_jobs))
|
|
t.start()
|
|
|
|
return finished_event, t, shm_jobs.name, shm_jobs
|
|
|
|
|
|
def unwrap_to_grey(image, is_latent: bool = False) -> np.ndarray:
|
|
squeeze = len(image.shape) > 3
|
|
image = image.cpu().numpy()
|
|
image = image.squeeze() if squeeze else image
|
|
image = np.moveaxis(image, 0, -1) if is_latent else image
|
|
return cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
|
|
|
|
|
|
def unwrap_and_quilt(wrapped_func, image, job_id, is_latent: bool = False):
|
|
"""quilting job when using batches"""
|
|
squeeze = len(image.shape) > 3
|
|
image = image.cpu().numpy()
|
|
image = image.squeeze() if squeeze else image
|
|
image = np.moveaxis(image, 0, -1) if is_latent else image
|
|
result = wrapped_func(image, job_id)
|
|
result = np.moveaxis(result, -1, 0) if is_latent else result
|
|
result = torch.from_numpy(result)
|
|
result = result.unsqueeze(0) if squeeze else result
|
|
return result
|
|
|
|
|
|
def unwrap_and_quilt_seamless(wrapped_func, image, lookup, job_id, is_latent: bool = False):
|
|
"""seamless quilting job when using batches"""
|
|
image = image.cpu().numpy()
|
|
squeeze = len(image.shape) > 3
|
|
image = image.squeeze() if squeeze else image
|
|
image = np.moveaxis(image, 0, -1) if is_latent else image
|
|
if lookup is not None:
|
|
lookup = lookup.cpu().numpy()
|
|
lookup = lookup.squeeze() if len(lookup.shape) > 3 else lookup
|
|
lookup = np.moveaxis(lookup, 0, -1) if is_latent else lookup
|
|
result = wrapped_func(image, lookup, job_id)
|
|
result = np.moveaxis(result, -1, 0) if is_latent else result
|
|
result = torch.from_numpy(result)
|
|
result = result.unsqueeze(0) if squeeze else result
|
|
return result
|
|
|
|
|
|
def overlap_percentage_to_pixels(block_size: int, overlap: float):
|
|
return int(block_size * overlap) if overlap > 0 else int(
|
|
block_size * get_quilting_shared_input_types()["overlap"][1]["default"])
|
|
|
|
|
|
def get_block_sizes(src, block_size_input: int, block_size_upper_bound: int | None = None):
|
|
match block_size_input:
|
|
case _ if block_size_input in [-1, 0]:
|
|
print(f"block size set to {block_size_input}!\nguessing nice block size...")
|
|
only_do_freq_analysis = block_size_input < 0
|
|
sizes = [guess_nice_block_size(unwrap_to_grey(src[i]), only_do_freq_analysis, block_size_upper_bound)
|
|
for i in range(src.shape[0])]
|
|
case 1:
|
|
print(f"src shape = {src.shape}")
|
|
sizes = [round(min(src.shape[1:3]) * 1 / 3)] * src.shape[0] # a "medium" block size, w/ respect to src
|
|
case 2:
|
|
sizes = [round(min(src.shape[1:3]) * 3 / 4)] * src.shape[0] # a "big" block size w/ respect to src
|
|
case __:
|
|
sizes = [block_size_input] * src.shape[0]
|
|
print(f"block sizes: {sizes}")
|
|
return sizes
|
|
|
|
|
|
def block_size_upper_bound_for_seamless(ori, tex_h, tex_w, overlap_percentage):
|
|
"""
|
|
if using guess_block_size in seamless nodes, mind H seam.
|
|
note that if the default computed upper bound is lower than the value here obtained, it will be used instead.
|
|
"""
|
|
match ori:
|
|
case "H & V":
|
|
# mind H seam w/ the following related assert -> image.shape[1] >= block_size + overlap * 2
|
|
# -> H >= S + S * OP * 2 <-> H / (1 + 2 * OP ) >= S
|
|
return round(tex_h / (1 + 2 * overlap_percentage) / 1.1)
|
|
case _______:
|
|
return None # texture default dims are fine as bounds
|
|
|
|
|
|
def batch_seamless_using_jobs(wrapped_func, src, lookup, is_latent: bool = False):
|
|
from joblib import Parallel, delayed
|
|
# process in the same fashion as lists
|
|
number_of_processes = src.shape[0] if lookup is None else max(src.shape[0], lookup.shape[0])
|
|
results = Parallel(n_jobs=-1, backend="loky", timeout=None)(
|
|
delayed(unwrap_and_quilt_seamless)(
|
|
wrapped_func,
|
|
src[min(i, src.shape[0] - 1)],
|
|
lookup[min(i, lookup.shape[0] - 1)]
|
|
if lookup is not None else None,
|
|
i, is_latent) for i in range(number_of_processes))
|
|
return torch.stack(results)
|
|
|
|
|
|
def validate_gen_args(source, block_size):
|
|
validate_array_shape(source, min_height=block_size, min_width=block_size,
|
|
help="Change the block size.")
|
|
|
|
|
|
def validate_seamless_args(orientation, source, lookup, block_size, overlap):
|
|
if not (overlap * 2 <= block_size):
|
|
raise ValueError(f"Overlap ({overlap}) needs to be 50% or less of the block size ({block_size}).")
|
|
|
|
if orientation == "H & V":
|
|
validate_array_shape(source, min_height=block_size, min_width=block_size + overlap * 2,
|
|
help="Change the block size or the overlap.")
|
|
else:
|
|
validate_array_shape(source, min_height=block_size, min_width=block_size,
|
|
help="Change the block size.")
|
|
|
|
if lookup is not None:
|
|
validate_array_shape(lookup, min_height=block_size, min_width=block_size,
|
|
help="Use a bigger lookup or change the block size.")
|
|
if lookup.dtype != source.dtype:
|
|
raise TypeError("lookup_texture dtype does not match image dtype")
|
|
|
|
|
|
# endregion AUX FUNCTIONS & CLASSES
|
|
|
|
|
|
# region NODES
|
|
|
|
class ImageQuilting:
|
|
class ImageQuiltingFuncWrapper(QuiltingFuncWrapper):
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
def __call__(self, image, job_id):
|
|
if self.version == 2:
|
|
image = cv2.cvtColor(image, cv2.COLOR_RGB2Lab)
|
|
|
|
block_size = self.block_sizes[job_id]
|
|
overlap = overlap_percentage_to_pixels(block_size, self.overlap)
|
|
validate_gen_args(image, block_size)
|
|
|
|
if self.parallelization_lvl == 0:
|
|
result = generate_texture(
|
|
image, block_size, overlap, self.out_h, self.out_w, self.tolerance,
|
|
self.version, self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
else:
|
|
result = generate_texture_parallel(
|
|
image, block_size, overlap, self.out_h, self.out_w, self.tolerance,
|
|
self.version, self.parallelization_lvl, self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
if self.version == 2:
|
|
result = cv2.cvtColor(result, cv2.COLOR_LAB2RGB)
|
|
|
|
return result
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"src": ("IMAGE",),
|
|
# can be a single image or a batch of images
|
|
# using an image batch limits max parallelization lvl to 1 ( values above are ignored ).
|
|
|
|
# self explanatory
|
|
"scale": ("FLOAT", {"default": 4, "min": 2, "max": 32, "step": .1}),
|
|
|
|
**get_quilting_shared_input_types()
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, block_size, scale, overlap, tolerance, parallelization_lvl, seed, version):
|
|
h, w = src.shape[1:3]
|
|
out_h, out_w = int(scale * h), int(scale * w)
|
|
block_sizes = get_block_sizes(src, block_size)
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
# note: the input src should have normalized values, not 0 to 255
|
|
|
|
finish_event, t, shm_name, shm_jobs = \
|
|
setup_pbar_quilting(block_sizes, overlap, out_h, out_w, parallelization_lvl)
|
|
|
|
try:
|
|
func = self.ImageQuiltingFuncWrapper(
|
|
block_sizes, overlap, out_h, out_w, tolerance, parallelization_lvl, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1
|
|
output = self.batch_using_jobs(func, src) if is_batch else unwrap_and_quilt(func, src, 0)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return (output,)
|
|
|
|
@staticmethod
|
|
def batch_using_jobs(wrapped_func, src):
|
|
from joblib import Parallel, delayed
|
|
results = Parallel(n_jobs=-1, backend="loky", timeout=None)(
|
|
delayed(unwrap_and_quilt)(wrapped_func, src[i], i) for i in range(src.shape[0]))
|
|
return torch.stack(results)
|
|
|
|
|
|
class LatentQuilting:
|
|
class LatentQuiltingFuncWrapper(QuiltingFuncWrapper):
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
def __call__(self, latent_image, job_id):
|
|
block_size = self.block_sizes[job_id]
|
|
overlap = overlap_percentage_to_pixels(block_size, self.overlap)
|
|
validate_gen_args(latent_image, block_size)
|
|
if self.parallelization_lvl == 0:
|
|
return generate_texture(
|
|
latent_image, block_size, overlap, self.out_h, self.out_w, self.tolerance,
|
|
self.version, self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
else:
|
|
return generate_texture_parallel(
|
|
latent_image, block_size, overlap, self.out_h, self.out_w, self.tolerance,
|
|
self.version, self.parallelization_lvl, self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
shared_inputs = get_quilting_shared_input_types()
|
|
# don't allow for auto block size
|
|
shared_inputs["block_size"][1]["default"] = 18
|
|
shared_inputs["block_size"][1]["min"] = 3
|
|
# note: could try running over individual channels; not sure if worth it thought.
|
|
|
|
return {
|
|
"required": {
|
|
"src": ("LATENT",),
|
|
|
|
# self explanatory
|
|
"scale": ("FLOAT", {"default": 4, "min": 2, "max": 32, "step": .1}),
|
|
|
|
**shared_inputs
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("LATENT",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, block_size, scale, overlap, tolerance, parallelization_lvl, seed, version):
|
|
src = src["samples"]
|
|
h, w = src.shape[2:4]
|
|
out_h, out_w = int(scale * h), int(scale * w)
|
|
block_sizes = [block_size] * src.shape[0]
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
finish_event, t, shm_name, shm_jobs = \
|
|
setup_pbar_quilting(block_sizes, overlap, out_h, out_w, parallelization_lvl)
|
|
|
|
try:
|
|
func = self.LatentQuiltingFuncWrapper(
|
|
block_sizes, overlap, out_h, out_w, tolerance, parallelization_lvl, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1
|
|
output = self.batch_using_jobs(func, src) if is_batch else \
|
|
unwrap_and_quilt(func, src, 0, is_latent=True)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return ({"samples": output},)
|
|
|
|
@staticmethod
|
|
def batch_using_jobs(wrapped_func, src):
|
|
from joblib import Parallel, delayed
|
|
results = Parallel(n_jobs=-1, backend="loky", timeout=None)(
|
|
delayed(unwrap_and_quilt)(wrapped_func, src[i], i, True) for i in range(src.shape[0]))
|
|
return torch.stack(results)
|
|
|
|
|
|
class ImageMakeSeamlessMB:
|
|
"""Transition stripe is built using overlapping square patches."""
|
|
|
|
@dataclass
|
|
class SeamlessFuncWrapper:
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
ori: str
|
|
block_sizes: list[int]
|
|
overlap: int
|
|
tolerance: float
|
|
rng: numpy.random.Generator
|
|
version: int
|
|
jobs_shm_name: str
|
|
|
|
def __call__(self, image, lookup, job_id):
|
|
from .make_seamless import make_seamless_horizontally, make_seamless_vertically, make_seamless_both
|
|
|
|
block_size = self.block_sizes[min(job_id, len(self.block_sizes) - 1)]
|
|
overlap = overlap_percentage_to_pixels(block_size, self.overlap)
|
|
|
|
if self.version == 2:
|
|
image = cv2.cvtColor(image, cv2.COLOR_RGB2Lab)
|
|
if lookup is not None:
|
|
lookup = cv2.cvtColor(lookup, cv2.COLOR_RGB2Lab) if lookup is not None else None
|
|
|
|
match self.ori:
|
|
case "H":
|
|
func = make_seamless_horizontally
|
|
case "V":
|
|
func = make_seamless_vertically
|
|
case ___:
|
|
func = make_seamless_both
|
|
|
|
validate_seamless_args(self.ori, image, lookup, block_size, overlap)
|
|
result = func(image, block_size, overlap, self.tolerance,
|
|
self.rng, self.version, lookup, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
if self.version == 2:
|
|
result = cv2.cvtColor(result, cv2.COLOR_LAB2RGB)
|
|
|
|
return result
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
inputs = get_quilting_shared_input_types()
|
|
inputs.pop("parallelization_lvl")
|
|
inputs["version"][1]["min"] = 1
|
|
inputs["overlap"][1]["max"] = .5
|
|
return {
|
|
"required": {
|
|
"src": ("IMAGE",),
|
|
"ori": (SEAMLESS_DIRS, {"default": SEAMLESS_DIRS[0]}),
|
|
**inputs,
|
|
},
|
|
"optional": {
|
|
"lookup": ("IMAGE",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, ori, block_size, overlap, tolerance, seed, version, lookup=None):
|
|
# note that src = lookup is the current algorithm policy when lookup is not provided.
|
|
# this policy could change in the future, so do not apply it here too despite being idempotent.
|
|
h, w = src.shape[1:3]
|
|
blk_size_upper_bound = block_size_upper_bound_for_seamless(ori, h, w, overlap)
|
|
block_sizes = get_block_sizes(src, block_size, blk_size_upper_bound)
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
lookup_batch_size = lookup.shape[0] if lookup is not None else 0
|
|
finish_event, t, shm_name, shm_jobs = setup_pbar_seamless(
|
|
ori, block_sizes, overlap, h, w, max(src.shape[0], lookup_batch_size))
|
|
|
|
try:
|
|
func = self.SeamlessFuncWrapper(
|
|
ori, block_sizes, overlap, tolerance, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1 or lookup_batch_size > 1
|
|
output = batch_seamless_using_jobs(func, src, lookup, is_latent=False) if is_batch else \
|
|
unwrap_and_quilt_seamless(func, src, lookup, 0)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return (output,)
|
|
|
|
|
|
class ImageMakeSeamlessSB:
|
|
"""Transition stripe is built via a single rectangular block."""
|
|
|
|
@dataclass
|
|
class SeamlessFuncWrapper:
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
ori: str
|
|
block_sizes: list[int]
|
|
overlap: int
|
|
rng: numpy.random.Generator
|
|
version: int
|
|
jobs_shm_name: str
|
|
|
|
def __call__(self, image, lookup, job_id):
|
|
from .make_seamless2 import seamless_horizontal, seamless_vertical, seamless_both
|
|
|
|
block_size = self.block_sizes[min(job_id, len(self.block_sizes) - 1)]
|
|
overlap = overlap_percentage_to_pixels(block_size, self.overlap)
|
|
|
|
if self.version == 2:
|
|
image = cv2.cvtColor(image, cv2.COLOR_RGB2Lab)
|
|
if lookup is not None:
|
|
lookup = cv2.cvtColor(lookup, cv2.COLOR_RGB2Lab) if lookup is not None else None
|
|
|
|
match self.ori:
|
|
case "H":
|
|
func = seamless_horizontal
|
|
case "V":
|
|
func = seamless_vertical
|
|
case ___:
|
|
func = seamless_both
|
|
|
|
validate_seamless_args(self.ori, image, lookup, block_size, overlap)
|
|
result = func(image, block_size, overlap,
|
|
self.version, lookup, self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
if self.version == 2:
|
|
result = cv2.cvtColor(result, cv2.COLOR_LAB2RGB)
|
|
|
|
return result
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
inputs = get_quilting_shared_input_types()
|
|
inputs.pop("parallelization_lvl")
|
|
inputs.pop("tolerance")
|
|
inputs["version"][1]["min"] = 1
|
|
inputs["overlap"][1]["max"] = .5
|
|
return {
|
|
"required": {
|
|
"src": ("IMAGE",),
|
|
"ori": (SEAMLESS_DIRS, {"default": SEAMLESS_DIRS[0]}),
|
|
**inputs,
|
|
},
|
|
"optional": {
|
|
"lookup": ("IMAGE",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, ori, block_size, overlap, seed, version, lookup=None):
|
|
# note that src = lookup is the current algorithm policy when lookup is not provided.
|
|
# this policy could change in the future, so do not apply it here too despite being idempotent.
|
|
h, w = src.shape[1:3]
|
|
blk_size_upper_bound = block_size_upper_bound_for_seamless(ori, h, w, overlap)
|
|
block_sizes = get_block_sizes(src, block_size, blk_size_upper_bound)
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
lookup_batch_size = lookup.shape[0] if lookup is not None else 0
|
|
finish_event, t, shm_name, shm_jobs = setup_pbar_seamless_v2(ori, max(src.shape[0], lookup_batch_size))
|
|
|
|
try:
|
|
func = self.SeamlessFuncWrapper(
|
|
ori, block_sizes, overlap, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1 or lookup_batch_size > 1
|
|
output = batch_seamless_using_jobs(func, src, lookup, is_latent=False) if is_batch else \
|
|
unwrap_and_quilt_seamless(func, src, lookup, 0)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return (output,)
|
|
|
|
|
|
class GuessNiceBlockSize:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"src": ("IMAGE",),
|
|
"simple_and_fast": ("BOOLEAN", {"default": False})
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("INT",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, simple_and_fast):
|
|
return (guess_nice_block_size(unwrap_to_grey(src[0]), freq_analysis_only=simple_and_fast),)
|
|
|
|
|
|
class LatentMakeSeamlessMB:
|
|
"""Transition stripe is built using overlapping square patches."""
|
|
|
|
@dataclass
|
|
class SeamlessFuncWrapper:
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
ori: str
|
|
block_size: int
|
|
overlap: int
|
|
tolerance: float
|
|
rng: numpy.random.Generator
|
|
version: int
|
|
jobs_shm_name: str
|
|
|
|
def __call__(self, latent, lookup, job_id):
|
|
from .make_seamless import make_seamless_horizontally, make_seamless_vertically, make_seamless_both
|
|
|
|
match self.ori:
|
|
case "H":
|
|
func = make_seamless_horizontally
|
|
case "V":
|
|
func = make_seamless_vertically
|
|
case ___:
|
|
func = make_seamless_both
|
|
|
|
overlap = overlap_percentage_to_pixels(self.block_size, self.overlap)
|
|
validate_seamless_args(self.ori, latent, lookup, self.block_size, overlap)
|
|
result = func(latent, self.block_size, overlap, self.tolerance,
|
|
self.rng, self.version, lookup, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
return result
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
inputs = LatentQuilting.INPUT_TYPES()["required"] #get_quilting_shared_input_types()
|
|
for to_remove in ["parallelization_lvl", "scale", "src"]:
|
|
inputs.pop(to_remove)
|
|
inputs["version"][1]["min"] = 1
|
|
inputs["overlap"][1]["max"] = .5
|
|
return {
|
|
"required": {
|
|
"src": ("LATENT",),
|
|
"ori": (SEAMLESS_DIRS, {"default": SEAMLESS_DIRS[0]}),
|
|
**inputs,
|
|
},
|
|
"optional": {
|
|
"lookup": ("LATENT",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("LATENT",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, ori, block_size, overlap, tolerance, seed, version, lookup=None):
|
|
src = src["samples"]
|
|
lookup = lookup["samples"] if lookup is not None else None
|
|
h, w = src.shape[2:4]
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
lookup_batch_size = lookup.shape[0] if lookup is not None else 0
|
|
finish_event, t, shm_name, shm_jobs = setup_pbar_seamless(
|
|
ori, [block_size] * src.shape[0], overlap, h, w, max(src.shape[0], lookup_batch_size))
|
|
|
|
try:
|
|
func = self.SeamlessFuncWrapper(
|
|
ori, block_size, overlap, tolerance, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1 or lookup_batch_size > 1
|
|
output = batch_seamless_using_jobs(func, src, lookup, is_latent=True) if is_batch else \
|
|
unwrap_and_quilt_seamless(func, src, lookup, 0, is_latent=True)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return ({"samples": output},)
|
|
|
|
|
|
class LatentMakeSeamlessSB:
|
|
"""Transition stripe is built using overlapping square patches."""
|
|
|
|
@dataclass
|
|
class SeamlessFuncWrapper:
|
|
"""Wraps node functionality for easy re-use when using jobs."""
|
|
|
|
ori: str
|
|
block_size: int
|
|
overlap: int
|
|
rng: numpy.random.Generator
|
|
version: int
|
|
jobs_shm_name: str
|
|
|
|
def __call__(self, latent, lookup, job_id):
|
|
from .make_seamless2 import seamless_horizontal, seamless_vertical, seamless_both
|
|
|
|
match self.ori:
|
|
case "H":
|
|
func = seamless_horizontal
|
|
case "V":
|
|
func = seamless_vertical
|
|
case ___:
|
|
func = seamless_both
|
|
|
|
overlap = overlap_percentage_to_pixels(self.block_size, self.overlap)
|
|
validate_seamless_args(self.ori, latent, lookup, self.block_size, overlap)
|
|
result = func(latent, self.block_size, overlap, self.version, lookup,
|
|
self.rng, UiCoordData(self.jobs_shm_name, job_id))
|
|
|
|
return result
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
inputs = LatentQuilting.INPUT_TYPES()["required"] #get_quilting_shared_input_types()
|
|
for to_remove in ["tolerance", "parallelization_lvl", "scale", "src"]:
|
|
inputs.pop(to_remove)
|
|
inputs["version"][1]["min"] = 1
|
|
inputs["overlap"][1]["max"] = .5
|
|
return {
|
|
"required": {
|
|
"src": ("LATENT",),
|
|
"ori": (SEAMLESS_DIRS, {"default": SEAMLESS_DIRS[0]}),
|
|
**inputs,
|
|
},
|
|
"optional": {
|
|
"lookup": ("LATENT",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("LATENT",)
|
|
FUNCTION = "compute"
|
|
CATEGORY = NODES_CATEGORY
|
|
|
|
def compute(self, src, ori, block_size, overlap, seed, version, lookup=None):
|
|
src = src["samples"]
|
|
lookup = lookup["samples"] if lookup is not None else None
|
|
rng: numpy.random.Generator = np.random.default_rng(seed=seed)
|
|
|
|
lookup_batch_size = lookup.shape[0] if lookup is not None else 0
|
|
finish_event, t, shm_name, shm_jobs = setup_pbar_seamless_v2(ori, max(src.shape[0], lookup_batch_size))
|
|
|
|
try:
|
|
func = self.SeamlessFuncWrapper(ori, block_size, overlap, rng, version, shm_name)
|
|
|
|
is_batch = src.shape[0] > 1 or lookup_batch_size > 1
|
|
output = batch_seamless_using_jobs(func, src, lookup, is_latent=True) if is_batch else \
|
|
unwrap_and_quilt_seamless(func, src, lookup, 0, is_latent=True)
|
|
finally:
|
|
terminate_task(finish_event, shm_jobs, t)
|
|
return ({"samples": output},)
|
|
|
|
# endregion NODES
|