feat: support range selection
This commit is contained in:
@@ -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.
|
||||
|
||||
**New in 2023/12/2:** Support range selection with left bound included and right bound excluded, see example below.
|
||||
|
||||
For example:
|
||||
|
||||
1. `1`: select the first 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.
|
||||
|
||||
|
||||
+95
-15
@@ -12,7 +12,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import typing as tg
|
||||
import collections.abc as clabc
|
||||
|
||||
import torch
|
||||
|
||||
@@ -28,7 +28,8 @@ class ImageSelector:
|
||||
@classmethod
|
||||
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
|
||||
"""
|
||||
return {
|
||||
@@ -42,7 +43,7 @@ class ImageSelector:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
# RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -50,13 +51,42 @@ class ImageSelector:
|
||||
|
||||
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
|
||||
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(','):
|
||||
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
|
||||
if x < len_first_dim:
|
||||
selected_index.append(x)
|
||||
@@ -64,10 +94,10 @@ class ImageSelector:
|
||||
pass
|
||||
|
||||
if selected_index:
|
||||
print(f"ImageSelector: selected: {len(selected_index)} latents")
|
||||
print(f"ImageSelector: selected: {len(selected_index)} images")
|
||||
return (images[selected_index, :, :, :], )
|
||||
|
||||
print(f"ImageSelector: selected no latents, passthrough")
|
||||
print(f"ImageSelector: selected no images, passthrough")
|
||||
return (images, )
|
||||
|
||||
|
||||
@@ -98,7 +128,7 @@ class ImageDuplicator:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
# RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -107,6 +137,17 @@ class ImageDuplicator:
|
||||
CATEGORY = "image"
|
||||
|
||||
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
|
||||
] + [torch.clone(images) for _ in range(dup_times - 1)]
|
||||
@@ -129,7 +170,8 @@ class LatentSelector:
|
||||
@classmethod
|
||||
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
|
||||
"""
|
||||
return {
|
||||
@@ -143,7 +185,7 @@ class LatentSelector:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", )
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
# RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -151,15 +193,42 @@ class LatentSelector:
|
||||
|
||||
CATEGORY = "latent"
|
||||
|
||||
def run(self, latent_image: tg.Mapping[tg.Text, torch.Tensor],
|
||||
selected_indexes: tg.Text):
|
||||
def run(self, latent_image: clabc.Mapping[str, torch.Tensor],
|
||||
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']
|
||||
shape = samples.shape
|
||||
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(','):
|
||||
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
|
||||
if x < len_first_dim:
|
||||
selected_index.append(x)
|
||||
@@ -200,7 +269,7 @@ class LatentDuplicator:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", )
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
# RETURN_NAMES = ("image_output_name",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -208,8 +277,19 @@ class LatentDuplicator:
|
||||
|
||||
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):
|
||||
"""
|
||||
对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']
|
||||
|
||||
sample_list = [samples] + [
|
||||
|
||||
Reference in New Issue
Block a user