264 lines
7.1 KiB
Python
264 lines
7.1 KiB
Python
"""Workflow-compatible implementations of the ETK and SWAIN list selectors."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Mapping
|
|
from copy import deepcopy
|
|
|
|
|
|
class AnyType(str):
|
|
"""ComfyUI wildcard socket type for a list whose item type is not fixed."""
|
|
|
|
def __ne__(self, other):
|
|
return False
|
|
|
|
|
|
ANY_TYPE = AnyType("*")
|
|
|
|
|
|
class ETKListToComfyList:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"items": ("LIST", {"default": []}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = (ANY_TYPE,)
|
|
RETURN_NAMES = ("items",)
|
|
OUTPUT_IS_LIST = (True,)
|
|
FUNCTION = "convert"
|
|
CATEGORY = "ETK/List"
|
|
|
|
def convert(self, items):
|
|
if not isinstance(items, list):
|
|
raise TypeError(f"items must be a LIST, got {type(items).__name__}")
|
|
return (items,)
|
|
|
|
|
|
class ImageFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"images": ("LIST", {"default": []}),
|
|
"index": ("INT", {"default": 0, "min": 0, "max": 10e10}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE",)
|
|
FUNCTION = "get_image"
|
|
CATEGORY = "ETK/Image"
|
|
|
|
def get_image(self, images, index):
|
|
return (images[index],)
|
|
|
|
|
|
class LatentFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"latents": ("LIST", {"default": []}),
|
|
"index": ("INT", {"default": 0, "min": -10e10, "max": 10e10}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("LATENT",)
|
|
FUNCTION = "get_latent"
|
|
CATEGORY = "ETK/Image"
|
|
|
|
def get_latent(self, latents, index):
|
|
return (latents[index],)
|
|
|
|
|
|
class ConditioningFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"conditionings": ("LIST", {"default": []}),
|
|
"index": ("INT", {"default": 0, "min": 0, "max": 10e10}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("CONDITIONING",)
|
|
FUNCTION = "get_conditioning"
|
|
CATEGORY = "ETK/Image"
|
|
|
|
def get_conditioning(self, conditionings, index):
|
|
try:
|
|
return (conditionings[index],)
|
|
except IndexError as exc:
|
|
raise IndexError(
|
|
f"ConditioningFromListIdx index {index} out of range for list length {len(conditionings)}"
|
|
) from exc
|
|
|
|
|
|
def _normalize_comfy_audio(audio):
|
|
import torch
|
|
|
|
if audio is None:
|
|
raise ValueError("audio input must be provided")
|
|
if not isinstance(audio, Mapping):
|
|
raise ValueError("AUDIO input must be a mapping with 'waveform' and 'sample_rate' keys")
|
|
|
|
try:
|
|
waveform = audio["waveform"]
|
|
except KeyError as exc:
|
|
raise ValueError("AUDIO input must contain a 'waveform' key") from exc
|
|
|
|
if not torch.is_tensor(waveform):
|
|
waveform = torch.tensor(waveform)
|
|
if waveform.dim() == 1:
|
|
waveform = waveform.unsqueeze(0).unsqueeze(0)
|
|
if waveform.dim() == 2:
|
|
waveform = waveform.unsqueeze(0)
|
|
if waveform.dim() != 3:
|
|
raise ValueError(f"AUDIO waveform must have shape [B, C, T], got {tuple(waveform.shape)}")
|
|
|
|
sample_rate = int(audio.get("sample_rate", 44100) or 44100)
|
|
if sample_rate <= 0:
|
|
raise ValueError(f"AUDIO sample_rate must be positive, got {sample_rate}")
|
|
|
|
audio_dict = dict(audio)
|
|
audio_dict["waveform"] = waveform
|
|
audio_dict["sample_rate"] = sample_rate
|
|
return audio_dict
|
|
|
|
|
|
class AudioFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"audios": ("LIST", {"default": []}),
|
|
"index": ("INT", {"default": 0, "min": -10e10, "max": 10e10}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("AUDIO",)
|
|
FUNCTION = "get_audio"
|
|
CATEGORY = "ETK/audio"
|
|
|
|
def get_audio(self, audios, index):
|
|
return (_normalize_comfy_audio(audios[int(index)]),)
|
|
|
|
|
|
def _swain_input_types(*, bounded_index: bool):
|
|
index_config = {"default": 0}
|
|
if bounded_index:
|
|
index_config.update({"min": -1000, "max": 1000})
|
|
return {
|
|
"required": {"list": ("LIST", {"default": []})},
|
|
"optional": {"int": ("INT", index_config)},
|
|
}
|
|
|
|
|
|
def _selected(kwargs):
|
|
values = kwargs.get("list", None)
|
|
index = kwargs.get("int", None)
|
|
if values is None:
|
|
raise ValueError("Must provide an input")
|
|
if not isinstance(values, list):
|
|
raise ValueError("input must be a list")
|
|
if not isinstance(index, int):
|
|
raise ValueError("index must be an int")
|
|
if index >= len(values):
|
|
raise ValueError("index must be in the range of the list")
|
|
return values, values[index]
|
|
|
|
|
|
class DictFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return _swain_input_types(bounded_index=True)
|
|
|
|
RETURN_TYPES = ("DICT",)
|
|
RETURN_NAMES = ("dictionary",)
|
|
FUNCTION = "handler"
|
|
CATEGORY = "SWAIN/text"
|
|
|
|
def handler(self, **kwargs):
|
|
_, value = _selected(deepcopy(kwargs))
|
|
return (value,)
|
|
|
|
|
|
class StrFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return _swain_input_types(bounded_index=True)
|
|
|
|
RETURN_TYPES = ("STRING", "LIST")
|
|
RETURN_NAMES = ("string", "list")
|
|
FUNCTION = "handler"
|
|
CATEGORY = "SWAIN/text"
|
|
|
|
def handler(self, **kwargs):
|
|
values, value = _selected(deepcopy(kwargs))
|
|
return value, values
|
|
|
|
|
|
class BytesFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return _swain_input_types(bounded_index=False)
|
|
|
|
RETURN_TYPES = ("BYTES",)
|
|
RETURN_NAMES = ("bytes",)
|
|
FUNCTION = "handler"
|
|
CATEGORY = "SWAIN/text"
|
|
|
|
def handler(self, **kwargs):
|
|
_, value = _selected(deepcopy(kwargs))
|
|
if isinstance(value, list):
|
|
if len(value) > 1:
|
|
raise ValueError("bytes must be a list of length 1")
|
|
value = value[0]
|
|
return (value,)
|
|
|
|
|
|
class IntFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return _swain_input_types(bounded_index=False)
|
|
|
|
RETURN_TYPES = ("INT", "LIST")
|
|
RETURN_NAMES = ("int", "list")
|
|
FUNCTION = "handler"
|
|
CATEGORY = "SWAIN/text"
|
|
|
|
def handler(self, **kwargs):
|
|
values, value = _selected(deepcopy(kwargs))
|
|
if not isinstance(value, int):
|
|
try:
|
|
value = int(value)
|
|
except Exception as exc:
|
|
raise ValueError(
|
|
"list at idx must be an int or a string that can be converted to an int"
|
|
) from exc
|
|
return value, values
|
|
|
|
|
|
class FloatFromListIdx:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return _swain_input_types(bounded_index=False)
|
|
|
|
RETURN_TYPES = ("FLOAT", "LIST")
|
|
RETURN_NAMES = ("float", "list")
|
|
FUNCTION = "handler"
|
|
CATEGORY = "SWAIN/text"
|
|
|
|
def handler(self, **kwargs):
|
|
values, value = _selected(deepcopy(kwargs))
|
|
if not isinstance(value, float):
|
|
try:
|
|
value = float(value)
|
|
except Exception as exc:
|
|
raise ValueError(
|
|
"list at idx must be an int or a string that can be converted to an int"
|
|
) from exc
|
|
return value, values
|