diff --git a/utility_nodes/list_utility_nodes.py b/utility_nodes/list_utility_nodes.py index d9c8252..fcfda3b 100644 --- a/utility_nodes/list_utility_nodes.py +++ b/utility_nodes/list_utility_nodes.py @@ -1,5 +1,6 @@ import torch from comfy_api.latest import io +from typing import Iterable class BatchFromImageList(io.ComfyNode): @classmethod @@ -43,19 +44,23 @@ class ImageListFromBatch(io.ComfyNode): def execute(cls, images): # type: ignore image_list = list( i.unsqueeze(0) for i in images ) return io.NodeOutput(image_list,) - + +template = io.MatchType.Template("pick_from_list") + class PickFromList(io.ComfyNode): @classmethod def define_schema(cls): return io.Schema( node_id = "Pick from List", display_name = "Pick from List", + description = "Given a comma separated list of indexes, return only those entries from the input list.", inputs = [ - io.AnyType.Input("anything"), - io.String.Input("indexes", display_name="indexes", tooltip="comma separated list of indexes. Whitespace stripped. Only these entries will be included. Zero indexed.") + io.MatchType.Input("anything", template=template, tooltip="note that a list is expected (or an image batch)"), + io.String.Input("indexes", display_name="indexes", tooltip="comma separated list of indexes. Whitespace stripped. Only these entries will be included."), + io.Combo.Input("indexing", options=['0','1'], default='0', tooltip="is the first item number 0 or number 1?") ], outputs = [ - io.String.Output("picks", display_name="picks", is_output_list=True) + io.MatchType.Output(template=template, display_name="picks", is_output_list=True) ], category = "image_filter/helpers", is_input_list=True @@ -63,19 +68,25 @@ class PickFromList(io.ComfyNode): @classmethod - def execute(cls, anything:list, indexes:list[str]): # type: ignore + def execute(cls, anything:list, indexes:list[str], indexing:list[str|int]=[0,]): # type: ignore - if len(anything)==1 and isinstance(anything[0],list): - print("Warning: received list of lists. Processing just anything[0]") - anything = anything[0] + is_image_batch = False + if len(anything)==1: + if isinstance(anything[0],list): + print("Warning: PickFromList received length 1 list of lists. Processing anything[0]") + anything = anything[0] + elif isinstance(anything[0], torch.Tensor) and len(anything[0].shape) == 4: + anything = anything[0] # type: ignore + is_image_batch = True index_str:str = indexes[0] + offset:int = int(indexing[0]) - result = [] - for x in [x.strip() for x in index_str.split(',')]: - try: - result.append(anything[int(x)]) - except Exception as e: - print(f"{e} when processing {x} from {index_str}") - - return io.NodeOutput(result, ) \ No newline at end of file + try: + index_ns = [int(x.strip())-offset for x in index_str.split(',') if x.strip()] + result = [anything[x] for x in index_ns] + if is_image_batch: result = [torch.stack(result),] + return io.NodeOutput(result, ) + except Exception as e: + print(f"{e} when processing {index_str}") + raise