fix: make ImageSelector work again

This commit is contained in:
SLAPaper
2023-08-14 02:21:51 +08:00
parent 4796a09e50
commit 4a7cf2481e
+31 -16
View File
@@ -12,9 +12,10 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
import torch
import typing as tg import typing as tg
import torch
class ImageSelector: class ImageSelector:
""" """
@@ -49,20 +50,24 @@ class ImageSelector:
CATEGORY = "image" CATEGORY = "image"
def run(self, images: tg.Sequence[tg.Mapping[tg.Text, tg.Any]], def run(self, images: torch.Tensor, selected_indexes: tg.Text):
selected_indexes: tg.Text): shape = images.shape
res_images: tg.List[tg.Any] = [] len_first_dim = shape[0]
selected_index: tg.List[int] = []
for s in selected_indexes.strip().split(','): for s in selected_indexes.strip().split(','):
try: try:
x: int = int(s.strip()) - 1 x: int = int(s.strip()) - 1
if x < len(images): if x < len_first_dim:
res_images.append(images[x]) selected_index.append(x)
except: except:
pass pass
if res_images: if selected_index:
return (res_images, ) print(f"ImageSelector: selected: {len(selected_index)} latents")
return (images[selected_index, :, :, :], )
print(f"ImageSelector: selected no latents, passthrough")
return (images, ) return (images, )
@@ -72,6 +77,7 @@ class ImageDuplicator:
""" """
def __init__(self): def __init__(self):
self._name = "ImageDuplicator"
pass pass
@classmethod @classmethod
@@ -100,12 +106,16 @@ class ImageDuplicator:
CATEGORY = "image" CATEGORY = "image"
def run(self, images: tg.Sequence[tg.Any], dup_times: int): def run(self, images: torch.Tensor, dup_times: int):
res_images: tg.List[tg.Any] = []
for _ in range(dup_times):
res_images.extend(images)
return (res_images, ) tensor_list = [images
] + [torch.clone(images) for _ in range(dup_times - 1)]
print(
f"ImageDuplicator: dup {dup_times} times,",
f"return {len(tensor_list)} images",
)
return (torch.cat(tensor_list), )
class LatentSelector: class LatentSelector:
@@ -141,7 +151,7 @@ class LatentSelector:
CATEGORY = "latent" CATEGORY = "latent"
def run(self, latent_image: tg.Sequence[tg.Any], def run(self, latent_image: tg.Mapping[tg.Text, torch.Tensor],
selected_indexes: tg.Text): selected_indexes: tg.Text):
samples = latent_image['samples'] samples = latent_image['samples']
shape = samples.shape shape = samples.shape
@@ -157,8 +167,10 @@ class LatentSelector:
pass pass
if selected_index: if selected_index:
print(f"LatentSelector: selected: {len(selected_index)} latents")
return ({'samples': samples[selected_index, :, :, :]}, ) return ({'samples': samples[selected_index, :, :, :]}, )
print(f"LatentSelector: selected no latents, passthrough")
return (latent_image, ) return (latent_image, )
@@ -196,15 +208,18 @@ class LatentDuplicator:
CATEGORY = "latent" CATEGORY = "latent"
def run(self, latent_image: tg.Sequence[tg.Mapping[tg.Text, torch.Tensor]], def run(self, latent_image: tg.Mapping[tg.Text, torch.Tensor],
dup_times: int): dup_times: int):
samples = latent_image['samples'] samples = latent_image['samples']
shape = samples.shape
sample_list = [samples] + [ sample_list = [samples] + [
torch.clone(samples) for _ in range(dup_times - 1) torch.clone(samples) for _ in range(dup_times - 1)
] ]
print(
f"LatentDuplicator: dup {dup_times} times,",
f"return {len(sample_list)} images",
)
return ({ return ({
'samples': torch.cat(sample_list), 'samples': torch.cat(sample_list),
}, ) }, )