fix: make ImageSelector work again
This commit is contained in:
+31
-16
@@ -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),
|
||||||
}, )
|
}, )
|
||||||
|
|||||||
Reference in New Issue
Block a user