feat: support range selection

This commit is contained in:
SLAPaper
2023-12-02 12:00:00 +08:00
parent f2b87157f8
commit 291559780a
2 changed files with 106 additions and 21 deletions
+5
View File
@@ -22,10 +22,15 @@ Both nodes can be found in `image` category.
Selector takes a list of selected indexes, start with 1 (not 0, sorry), seperated by comma, and outputs only the selected images from input images. Selector takes a list of selected indexes, start with 1 (not 0, sorry), seperated by comma, and outputs only the selected images from input images.
**New in 2023/12/2:** Support range selection with left bound included and right bound excluded, see example below.
For example: For example:
1. `1`: select the first image 1. `1`: select the first image
2. `1,3,4,6,7`: select the 1st, 3rd, 4th, 6th and 7th image 2. `1,3,4,6,7`: select the 1st, 3rd, 4th, 6th and 7th image
3. `2:`: select 2nd, 3rd, ..., till the last image (omit the first image)
4. `:0`: select 1st, 2nd, ..., till the second last image (omit the last image)
5. `3:-1`: select 3rd, 4th, ..., till the third last image (omit first two and last two images)
All indexes that cannot convert to integer or out of bounds will be ignored. All indexes that cannot convert to integer or out of bounds will be ignored.
+95 -15
View File
@@ -12,7 +12,7 @@
# 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 typing as tg import collections.abc as clabc
import torch import torch
@@ -28,7 +28,8 @@ class ImageSelector:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
""" """
Input: list of index of selected image, seperated by comma Input: list of index of selected image, seperated by comma (",")
support colon (":") sperated range (left included, right excluded)
Indexes start with 1 for simplicity Indexes start with 1 for simplicity
""" """
return { return {
@@ -42,7 +43,7 @@ class ImageSelector:
} }
RETURN_TYPES = ("IMAGE", ) RETURN_TYPES = ("IMAGE", )
#RETURN_NAMES = ("image_output_name",) # RETURN_NAMES = ("image_output_name",)
FUNCTION = "run" FUNCTION = "run"
@@ -50,13 +51,42 @@ class ImageSelector:
CATEGORY = "image" CATEGORY = "image"
def run(self, images: torch.Tensor, selected_indexes: tg.Text): def run(self, images: torch.Tensor, selected_indexes: str):
"""
根据 selected_indexes 选择 images 中的图片,支持连续索引和范围索引
Args:
images (torch.Tensor): 输入的图像张量,维度为 [N, C, H, W], 其中 N 为图片数量, C 为通道数, H、W 为图片的高和宽。
selected_indexes (str): 选择的图片索引,支持连续索引和范围索引,例如:"0,2,4:6,8" 表示选择第1、3、5张和第2、4、6、8张图片。
Returns:
tuple: 选择的图片张量,维度为 [N', C, H, W],其中 N' 为选择的图片数量。
"""
shape = images.shape shape = images.shape
len_first_dim = shape[0] len_first_dim = shape[0]
selected_index: tg.List[int] = [] selected_index: list[int] = []
total_indexes: list[int] = list(range(len_first_dim))
for s in selected_indexes.strip().split(','): for s in selected_indexes.strip().split(','):
try: try:
if ":" in s:
_li = s.strip().split(':', maxsplit=1)
_start = _li[0]
_end = _li[1]
if _start and _end:
selected_index.extend(
total_indexes[int(_start) - 1:int(_end) - 1]
)
elif _start:
selected_index.extend(
total_indexes[int(_start) - 1:]
)
elif _end:
selected_index.extend(
total_indexes[:int(_end) - 1]
)
else:
x: int = int(s.strip()) - 1 x: int = int(s.strip()) - 1
if x < len_first_dim: if x < len_first_dim:
selected_index.append(x) selected_index.append(x)
@@ -64,10 +94,10 @@ class ImageSelector:
pass pass
if selected_index: if selected_index:
print(f"ImageSelector: selected: {len(selected_index)} latents") print(f"ImageSelector: selected: {len(selected_index)} images")
return (images[selected_index, :, :, :], ) return (images[selected_index, :, :, :], )
print(f"ImageSelector: selected no latents, passthrough") print(f"ImageSelector: selected no images, passthrough")
return (images, ) return (images, )
@@ -98,7 +128,7 @@ class ImageDuplicator:
} }
RETURN_TYPES = ("IMAGE", ) RETURN_TYPES = ("IMAGE", )
#RETURN_NAMES = ("image_output_name",) # RETURN_NAMES = ("image_output_name",)
FUNCTION = "run" FUNCTION = "run"
@@ -107,6 +137,17 @@ class ImageDuplicator:
CATEGORY = "image" CATEGORY = "image"
def run(self, images: torch.Tensor, dup_times: int): def run(self, images: torch.Tensor, dup_times: int):
"""
对输入的图像张量进行复制多次,并将复制后的张量拼接起来返回。
Args:
images (torch.Tensor): 输入的图像张量,维度为 (batch_size, channels, height, width)。
dup_times (int): 复制的次数。
Returns:
torch.Tensor: 拼接后的图像张量,维度为 (batch_size * dup_times, channels, height, width)。
"""
tensor_list = [images tensor_list = [images
] + [torch.clone(images) for _ in range(dup_times - 1)] ] + [torch.clone(images) for _ in range(dup_times - 1)]
@@ -129,7 +170,8 @@ class LatentSelector:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
""" """
Input: list of index of selected image, seperated by comma Input: list of index of selected image, seperated by comma (",")
support colon (":") sperated range (left included, right excluded)
Indexes start with 1 for simplicity Indexes start with 1 for simplicity
""" """
return { return {
@@ -143,7 +185,7 @@ class LatentSelector:
} }
RETURN_TYPES = ("LATENT", ) RETURN_TYPES = ("LATENT", )
#RETURN_NAMES = ("image_output_name",) # RETURN_NAMES = ("image_output_name",)
FUNCTION = "run" FUNCTION = "run"
@@ -151,15 +193,42 @@ class LatentSelector:
CATEGORY = "latent" CATEGORY = "latent"
def run(self, latent_image: tg.Mapping[tg.Text, torch.Tensor], def run(self, latent_image: clabc.Mapping[str, torch.Tensor],
selected_indexes: tg.Text): selected_indexes: str):
"""
对latent_image进行筛选,根据selected_indexes指定的索引进行筛选
Args:
latent_image: 待筛选的latent_image,Mapping[str, torch.Tensor],包含'samples'字段
selected_indexes: 待筛选的索引,以逗号分隔,支持连续索引范围以冒号分隔,例如'1,3,5:7,9'
Returns:
筛选后的latent_image,Mapping[str, torch.Tensor]
"""
samples = latent_image['samples'] samples = latent_image['samples']
shape = samples.shape shape = samples.shape
len_first_dim = shape[0] len_first_dim = shape[0]
selected_index: tg.List[int] = [] selected_index: list[int] = []
total_indexes: list[int] = list(range(len_first_dim))
for s in selected_indexes.strip().split(','): for s in selected_indexes.strip().split(','):
try: try:
if ":" in s:
_li = s.strip().split(':', maxsplit=1)
_start = _li[0]
_end = _li[1]
if _start and _end:
selected_index.extend(
total_indexes[int(_start) - 1:int(_end) - 1]
)
elif _start:
selected_index.extend(
total_indexes[int(_start) - 1:]
)
elif _end:
selected_index.extend(
total_indexes[:int(_end) - 1]
)
else:
x: int = int(s.strip()) - 1 x: int = int(s.strip()) - 1
if x < len_first_dim: if x < len_first_dim:
selected_index.append(x) selected_index.append(x)
@@ -200,7 +269,7 @@ class LatentDuplicator:
} }
RETURN_TYPES = ("LATENT", ) RETURN_TYPES = ("LATENT", )
#RETURN_NAMES = ("image_output_name",) # RETURN_NAMES = ("image_output_name",)
FUNCTION = "run" FUNCTION = "run"
@@ -208,8 +277,19 @@ class LatentDuplicator:
CATEGORY = "latent" CATEGORY = "latent"
def run(self, latent_image: tg.Mapping[tg.Text, torch.Tensor], def run(self, latent_image: clabc.Mapping[str, torch.Tensor],
dup_times: int): dup_times: int):
"""
对latent_image进行复制, 复制次数为dup_times。
Args:
latent_image (clabc.Mapping[str, torch.Tensor]): 输入的latent_image, 包含'samples'键。
dup_times (int): 复制次数。
Returns:
Tuple[Dict[str, torch.Tensor]]: 返回包含samples的字典, samples是一个长度为(dup_times+1)的样本张量。
"""
samples = latent_image['samples'] samples = latent_image['samples']
sample_list = [samples] + [ sample_list = [samples] + [