diff --git a/node.py b/node.py index 550fc26..3588199 100644 --- a/node.py +++ b/node.py @@ -8,6 +8,8 @@ from torchvision import transforms import torch import base64 import time +import torchaudio +import soundfile as sf from replicate.client import Client from .schema_to_node import ( schema_to_comfyui_input_types, @@ -32,7 +34,11 @@ def create_comfyui_node(schema): def INPUT_TYPES(cls): return schema_to_comfyui_input_types(schema) - RETURN_TYPES = (return_type,) + RETURN_TYPES = ( + tuple(return_type.values()) + if isinstance(return_type, dict) + else (return_type,) + ) FUNCTION = "run_replicate_model" CATEGORY = "Replicate" @@ -45,6 +51,8 @@ def create_comfyui_node(schema): ) if input_type == "IMAGE": kwargs[key] = self.image_to_base64(value) + elif input_type == "AUDIO": + kwargs[key] = self.audio_to_base64(value) def image_to_base64(self, image): if isinstance(image, torch.Tensor): @@ -60,6 +68,31 @@ def create_comfyui_node(schema): img_str = base64.b64encode(buffer.getvalue()).decode() return f"data:image/png;base64,{img_str}" + def audio_to_base64(self, audio): + if ( + isinstance(audio, dict) + and "waveform" in audio + and "sample_rate" in audio + ): + waveform = audio["waveform"] + sample_rate = audio["sample_rate"] + else: + waveform, sample_rate = audio + + # Ensure waveform is 2D + if waveform.dim() == 1: + waveform = waveform.unsqueeze(0) + elif waveform.dim() > 2: + waveform = waveform.squeeze() + if waveform.dim() > 2: + raise ValueError("Waveform must be 1D or 2D") + + buffer = io.BytesIO() + sf.write(buffer, waveform.numpy().T, sample_rate, format="wav") + buffer.seek(0) + audio_str = base64.b64encode(buffer.getvalue()).decode() + return f"data:audio/wav;base64,{audio_str}" + def handle_array_inputs(self, kwargs): array_inputs = inputs_that_need_arrays(schema) for input_name in array_inputs: @@ -75,13 +108,18 @@ def create_comfyui_node(schema): def log_input(self, kwargs): truncated_kwargs = { k: v[:20] + "..." - if isinstance(v, str) and v.startswith("data:image") + if isinstance(v, str) + and (v.startswith("data:image") or v.startswith("data:audio")) else v for k, v in kwargs.items() } print(f"Running {replicate_model} with {truncated_kwargs}") def handle_image_output(self, output): + if output is None: + print("No image output received") + return None + # Handle both string and list outputs output_list = [output] if isinstance(output, str) else list(output) if output_list: @@ -112,19 +150,90 @@ def create_comfyui_node(schema): print("No output received from the model") return None + def handle_audio_output(self, output): + if output is None: + print("No audio output received from the model") + return None + + if isinstance(output, str): + output_list = [output] + elif isinstance(output, (list, tuple)): + output_list = list(output) + else: + print(f"Unexpected output type: {type(output)}") + return None + + if output_list: + audio_data = [] + for audio_url in output_list: + if audio_url: + response = requests.get(audio_url) + if response.status_code == 200: + audio_content = BytesIO(response.content) + waveform, sample_rate = torchaudio.load(audio_content) + audio_data.append( + { + "waveform": waveform.unsqueeze(0), + "sample_rate": sample_rate, + } + ) + else: + print( + f"Failed to download audio. Status code: {response.status_code}" + ) + else: + print("Empty audio URL received") + + # If there's only one audio file, return it directly + if len(audio_data) == 1: + return audio_data[0] + # If there are multiple audio files, return them as a list + return audio_data + else: + print("No valid audio URLs in the output") + return None + + def remove_falsey_optional_inputs(self, kwargs): + optional_inputs = self.INPUT_TYPES().get("optional", {}) + for key in list(kwargs.keys()): + if key in optional_inputs and not kwargs[key]: + del kwargs[key] + def run_replicate_model(self, **kwargs): self.handle_array_inputs(kwargs) + self.remove_falsey_optional_inputs(kwargs) self.convert_input_images_to_base64(kwargs) self.log_input(kwargs) - output = replicate.run(replicate_model, input=kwargs) + kwargs_without_force_rerun = { + k: v for k, v in kwargs.items() if k != "force_rerun" + } + output = replicate.run(replicate_model, input=kwargs_without_force_rerun) print(f"Output: {output}") - if return_type == "IMAGE": - output = self.handle_image_output(output) + processed_outputs = [] + if isinstance(return_type, dict): + for prop_name, prop_type in return_type.items(): + if prop_type == "IMAGE": + processed_outputs.append( + self.handle_image_output(output.get(prop_name)) + ) + elif prop_type == "AUDIO": + processed_outputs.append( + self.handle_audio_output(output.get(prop_name)) + ) + elif prop_type == "STRING": + processed_outputs.append( + "".join(list(output.get(prop_name, ""))).strip() + ) else: - output = "".join(list(output)).strip() + if return_type == "IMAGE": + processed_outputs.append(self.handle_image_output(output)) + elif return_type == "AUDIO": + processed_outputs.append(self.handle_audio_output(output)) + else: + processed_outputs.append("".join(list(output)).strip()) - return (output,) + return tuple(processed_outputs) return node_name, ReplicateToComfyUI @@ -135,7 +244,9 @@ def create_comfyui_nodes_from_schemas(schemas_dir): schemas_dir_path = os.path.join(current_path, schemas_dir) for schema_file in os.listdir(schemas_dir_path): if schema_file.endswith(".json"): - with open(os.path.join(schemas_dir_path, schema_file), "r", encoding="utf-8") as f: + with open( + os.path.join(schemas_dir_path, schema_file), "r", encoding="utf-8" + ) as f: schema = json.load(f) node_name, node_class = create_comfyui_node(schema) nodes[node_name] = node_class diff --git a/schema_to_node.py b/schema_to_node.py index 76df66a..8e91a65 100644 --- a/schema_to_node.py +++ b/schema_to_node.py @@ -1,17 +1,41 @@ DEFAULT_STEP = 0.01 DEFAULT_ROUND = 0.001 +IMAGE_EXTENSIONS = (".png", ".jpg", ".jpeg", ".gif", ".webp") +VIDEO_EXTENSIONS = (".mp4", ".mkv", ".webm", ".mov", ".mpg", ".mpeg") +AUDIO_EXTENSIONS = (".mp3", ".wav", ".flac", ".mpga", ".m4a") -def convert_to_comfyui_input_type(openapi_type, openapi_format=None): - type_mapping = { - "string": "STRING", - "integer": "INT", - "number": "FLOAT", - "boolean": "BOOLEAN", - } +TYPE_MAPPING = { + "string": "STRING", + "integer": "INT", + "number": "FLOAT", + "boolean": "BOOLEAN", +} + + +def convert_to_comfyui_input_type( + input_name, openapi_type, openapi_format=None, default_example_input=None +): if openapi_type == "string" and openapi_format == "uri": - return "IMAGE" - return type_mapping.get(openapi_type, "STRING") + if ( + default_example_input + and isinstance(default_example_input, dict) + and input_name in default_example_input + ): + if is_type(default_example_input[input_name], IMAGE_EXTENSIONS): + return "IMAGE" + elif is_type(default_example_input[input_name], VIDEO_EXTENSIONS): + return "VIDEO" + elif is_type(default_example_input[input_name], AUDIO_EXTENSIONS): + return "AUDIO" + elif "image" in input_name.lower(): + return "IMAGE" + elif "audio" in input_name.lower(): + return "AUDIO" + else: + return "STRING" + + return TYPE_MAPPING.get(openapi_type, "STRING") def name_and_version(schema): @@ -23,11 +47,13 @@ def name_and_version(schema): return replicate_model, node_name -def resolve_schema(prop_data, openapi_schemas): +def resolve_schema(prop_data, openapi_schema): if "$ref" in prop_data: ref_path = prop_data["$ref"].split("/") - current = openapi_schemas + current = openapi_schema for path in ref_path[1:]: # Skip the first '#' element + if path not in current: + return prop_data # Return original if path is invalid current = current[path] return current return prop_data @@ -37,6 +63,7 @@ def schema_to_comfyui_input_types(schema): openapi_schema = schema["latest_version"]["openapi_schema"] input_schema = openapi_schema["components"]["schemas"]["Input"] input_types = {"required": {}, "optional": {}} + default_example_input = get_default_example_input(schema) required_props = input_schema.get("required", []) @@ -51,7 +78,10 @@ def schema_to_comfyui_input_types(schema): input_type = prop_data["enum"] elif "type" in prop_data: input_type = convert_to_comfyui_input_type( - prop_data["type"], prop_data.get("format") + prop_name, + prop_data["type"], + prop_data.get("format"), + default_example_input, ) else: input_type = "STRING" @@ -134,22 +164,65 @@ def is_type(default_example_output, extensions): return False -def get_return_type(schema): - image_extensions = (".png", ".jpg", ".jpeg", ".gif", ".webp") - video_extensions = (".mp4", ".mkv", ".webm", ".mov", ".mpg", ".mpeg") - audio_extensions = (".mp3", ".wav") - - openapi_schema = schema["latest_version"]["openapi_schema"] - output_schema = openapi_schema["components"]["schemas"].get("Output") +def get_default_example(schema): default_example = schema.get("default_example") - default_example_output = default_example.get("output") if default_example else None + return default_example if default_example else None - if is_type(default_example_output, image_extensions): + +def get_default_example_input(schema): + default_example = get_default_example(schema) + return default_example.get("input") if default_example else None + + +def get_default_example_output(schema): + default_example = get_default_example(schema) + return default_example.get("output") if default_example else None + + +def get_return_type(schema): + openapi_schema = schema["latest_version"]["openapi_schema"] + output_schema = ( + openapi_schema.get("components", {}).get("schemas", {}).get("Output") + ) + default_example_output = get_default_example_output(schema) + + if output_schema and "$ref" in output_schema: + output_schema = resolve_schema(output_schema, openapi_schema) + + if isinstance(output_schema, dict) and output_schema.get("properties"): + return_types = {} + for prop_name, prop_data in output_schema["properties"].items(): + if isinstance(default_example_output, dict): + prop_value = default_example_output.get(prop_name) + + if is_type(prop_value, IMAGE_EXTENSIONS): + return_types[prop_name] = "IMAGE" + elif is_type(prop_value, AUDIO_EXTENSIONS): + return_types[prop_name] = "AUDIO" + elif is_type(prop_value, VIDEO_EXTENSIONS): + return_types[prop_name] = "VIDEO_URI" + else: + return_types[prop_name] = "STRING" + elif prop_data.get("format") == "uri": + if "audio" in prop_name.lower(): + return_types[prop_name] = "AUDIO" + elif "image" in prop_name.lower(): + return_types[prop_name] = "IMAGE" + else: + return_types[prop_name] = "STRING" + elif prop_data.get("type") == "string": + return_types[prop_name] = "STRING" + else: + return_types[prop_name] = "STRING" + + return return_types + + if is_type(default_example_output, IMAGE_EXTENSIONS): return "IMAGE" - elif is_type(default_example_output, video_extensions): + elif is_type(default_example_output, VIDEO_EXTENSIONS): return "VIDEO_URI" - elif is_type(default_example_output, audio_extensions): - return "AUDIO_URI" + elif is_type(default_example_output, AUDIO_EXTENSIONS): + return "AUDIO" if output_schema: if ( @@ -160,8 +233,8 @@ def get_return_type(schema): return "IMAGE" elif ( output_schema.get("type") == "array" - and output_schema["items"].get("type") == "string" - and output_schema["items"].get("format") == "uri" + and output_schema.get("items", {}).get("type") == "string" + and output_schema.get("items", {}).get("format") == "uri" ): # Handle multiple image output return "IMAGE"