Made SparseIndex and SparseSpread much more robust, added usable spread types for SparseSpread

This commit is contained in:
Jedrzej Kosinski
2023-12-20 23:43:42 -06:00
parent 1f66b7c581
commit be0460ec04
3 changed files with 75 additions and 24 deletions
+1 -1
View File
@@ -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
+63 -16
View File
@@ -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
+11 -7
View File
@@ -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: