This commit is contained in:
@@ -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, )
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user