From be0460ec04f1e41b4757f9257148280f4faccb66 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 20 Dec 2023 23:43:42 -0600 Subject: [PATCH] Made SparseIndex and SparseSpread much more robust, added usable spread types for SparseSpread --- control/control.py | 2 +- control/control_sparsectrl.py | 79 ++++++++++++++++++++++++++++------- control/nodes_sparsectrl.py | 18 ++++---- 3 files changed, 75 insertions(+), 24 deletions(-) diff --git a/control/control.py b/control/control.py index f81964f..c9bc138 100644 --- a/control/control.py +++ b/control/control.py @@ -257,7 +257,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): del self.cond_hint self.cond_hint = None # first, figure out which cond idxs are relevant, and where they fit in - cond_idxs = self.sparse_settings.sparse_method.get_indeces(hint_length=self.cond_hint_original.size(0), full_length=full_length) + cond_idxs = self.sparse_settings.sparse_method.get_indexes(hint_length=self.cond_hint_original.size(0), full_length=full_length) range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs hint_idxs = [] # idxs in cond_idxs diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 57bdec8..680a6d3 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -4,6 +4,7 @@ from abc import ABC, abstractmethod import math +import numpy as np from typing import Iterable, Union import torch import torch as th @@ -94,25 +95,53 @@ class SparseMethod(ABC): self.method = method @abstractmethod - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: pass class SparseSpreadMethod(SparseMethod): - def __init__(self, from_start=True): - super().__init__(self.SPREAD) - self.from_start = from_start + UNIFORM = "uniform" + STARTING = "starting" + ENDING = "ending" + CENTER = "center" - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + LIST = [UNIFORM, STARTING, ENDING, CENTER] + + def __init__(self, spread=UNIFORM): + super().__init__(self.SPREAD) + self.spread = spread + + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: + # if hint_length >= full_length, limit hints to full_length + if hint_length >= full_length: + return list(range(full_length)) # handle special case of 1 hint image if hint_length == 1: - return [0] if self.from_start else [full_length-1] - # handle special case of equal or more hint images than full length - if hint_length >= full_length: - return list(range(min(hint_length, full_length))) - if hint_length == 2: - return [0, full_length-1] - # TODO: other cases/modes + if self.spread in [self.UNIFORM, self.STARTING]: + return [0] + elif self.spread == self.ENDING: + return [full_length-1] + elif self.spread == self.CENTER: + # return second (of three) values as the center + return [np.linspace(0, full_length-1, 3, endpoint=True, dtype=int)[1]] + else: + raise ValueError(f"Unrecognized spread: {self.spread}") + # otherwise, handle other cases + if self.spread == self.UNIFORM: + return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int)) + elif self.spread == self.STARTING: + # make split 1 larger, remove last element + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1] + elif self.spread == self.ENDING: + # make split 1 larger, remove first element + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[1:] + elif self.spread == self.CENTER: + # if hint length is not 3 greater than full length, do STARTING behavior + if full_length-hint_length < 3: + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1] + # otherwise, get linspace of 2 greater than needed, then cut off first and last + return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int))[1:-1] + return ValueError(f"Unrecognized spread: {self.spread}") class SparseIndexMethod(SparseMethod): @@ -120,13 +149,31 @@ class SparseIndexMethod(SparseMethod): super().__init__(self.INDEX) self.idxs = idxs - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: + orig_hint_length = hint_length + if hint_length > full_length: + hint_length = full_length + # if idxs is less than hint_length, throw error + if len(self.idxs) < hint_length: + err_msg = f"There are not enough indexes ({len(self.idxs)}) provided to fit the usable {hint_length} input images." + if orig_hint_length != hint_length: + err_msg = f"{err_msg} (original input images: {orig_hint_length})" + raise ValueError(err_msg) + # cap idxs to hint_length + idxs = self.idxs[:hint_length] new_idxs = [] - for idx in self.idxs: + real_idxs = set() + for idx in idxs: if idx < 0: - new_idxs.append(full_length+idx) + real_idx = full_length+idx + if real_idx in real_idxs: + raise ValueError(f"Index '{idx}' maps to '{real_idx}' and is duplicate - indexes in Sparse Index Method must be unique.") else: - new_idxs.append(idx) + real_idx = idx + if real_idx in real_idxs: + raise ValueError(f"Index '{idx}' is duplicate (or a negative index is equivalent) - indexes in Sparse Index Method must be unique.") + real_idxs.add(real_idx) + new_idxs.append(real_idx) return new_idxs diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 63e0ea0..e396550 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -37,7 +37,7 @@ class SparseIndexMethodNode: def INPUT_TYPES(s): return { "required": { - "indeces": ("STRING", {"default": "0"}), + "indexes": ("STRING", {"default": "0"}), } } @@ -46,18 +46,22 @@ class SparseIndexMethodNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def get_method(self, indeces: str): + def get_method(self, indexes: str): idxs = [] + unique_idxs = set() # get indeces from string - str_idxs = [x.strip() for x in indeces.strip().split(",")] + str_idxs = [x.strip() for x in indexes.strip().split(",")] for str_idx in str_idxs: try: idx = int(str_idx) + if idx in unique_idxs: + raise ValueError(f"'{idx}' is duplicated; indexes must be unique.") idxs.append(idx) + unique_idxs.add(idx) except ValueError: raise ValueError(f"'{str_idx}' is not a valid integer index.") if len(idxs) == 0: - raise ValueError(f"No indeces were listed in Sparse Index Method.") + raise ValueError(f"No indexes were listed in Sparse Index Method.") return (SparseIndexMethod(idxs),) @@ -66,7 +70,7 @@ class SparseSpreadMethodNode: def INPUT_TYPES(s): return { "required": { - "from_start": ("BOOLEAN", {"default": True}), + "spread": (SparseSpreadMethod.LIST), } } @@ -75,8 +79,8 @@ class SparseSpreadMethodNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def get_method(self, from_start: bool): - return (SparseSpreadMethod(from_start=from_start),) + def get_method(self, spread: str): + return (SparseSpreadMethod(spread=spread),) class VAEEncodePreprocessor: