diff --git a/node.py b/node.py index 698bcb5..2d71700 100644 --- a/node.py +++ b/node.py @@ -9,133 +9,16 @@ from torchvision import transforms import torch import base64 import time +from .schema_to_node import ( + schema_to_comfyui_input_types, + get_return_type, + name_and_version, +) -def convert_to_comfyui_input_type(openapi_type, openapi_format=None): - type_mapping = { - "string": "STRING", - "integer": "INT", - "number": "FLOAT", - "boolean": "BOOLEAN", - } - if openapi_type == "string" and openapi_format == "uri": - return "IMAGE" - return type_mapping.get(openapi_type, "STRING") - - -def resolve_schema(prop_data, schemas): - if "$ref" in prop_data: - ref_path = prop_data["$ref"].split("/") - current = schemas - for path in ref_path[1:]: # Skip the first '#' element - current = current[path] - return current - return prop_data - - -def convert_schema_to_comfyui(schema, schemas): - input_types = {"required": {}, "optional": {}} - - required_props = schema.get("required", []) - - for prop_name, prop_data in schema["properties"].items(): - prop_data = resolve_schema(prop_data, schemas) - default_value = prop_data.get("default", None) - - if "allOf" in prop_data: - prop_data = resolve_schema(prop_data["allOf"][0], schemas) - - if "enum" in prop_data: - input_type = prop_data["enum"] - elif "type" in prop_data: - input_type = convert_to_comfyui_input_type( - prop_data["type"], prop_data.get("format") - ) - else: - input_type = "STRING" - - input_config = {"default": default_value} - - if "minimum" in prop_data: - input_config["min"] = prop_data["minimum"] - if "maximum" in prop_data: - input_config["max"] = prop_data["maximum"] - if input_type == "FLOAT": - input_config["step"] = 0.01 - input_config["round"] = 0.001 - - if "prompt" in prop_name and prop_data.get("type") == "string": - input_config["multiline"] = True - - # Meta prompt_template needs `{prompt}` to be sent through - if "template" not in prop_name: - input_config["dynamicPrompts"] = True - - if prop_name in required_props: - input_types["required"][prop_name] = (input_type, input_config) - else: - input_types["optional"][prop_name] = (input_type, input_config) - - input_types["optional"]["force_rerun"] = ("BOOLEAN", {"default": False}) - - return reorder_input_types(input_types, schema) - - -def reorder_input_types(input_types, schema): - ordered_input_types = {"required": {}, "optional": {}} - sorted_properties = sorted( - schema["properties"].items(), key=lambda x: x[1].get("x-order", float("inf")) - ) - - for prop_name, _ in sorted_properties: - if prop_name in input_types["required"]: - ordered_input_types["required"][prop_name] = input_types["required"][ - prop_name - ] - elif prop_name in input_types["optional"]: - ordered_input_types["optional"][prop_name] = input_types["optional"][ - prop_name - ] - - # force_rerun is always at the end of optional inputs - ordered_input_types["optional"]["force_rerun"] = input_types["optional"][ - "force_rerun" - ] - - return ordered_input_types - - -def get_return_type(schemas, model_info): - output_schema = schemas["components"]["schemas"].get("Output") - default_example = model_info.get("default_example") - default_example_output = ( - default_example.get("output", []) if default_example else [] - ) - - if ( - output_schema - and output_schema.get("type") == "array" - and output_schema["items"].get("type") == "string" - and output_schema["items"].get("format") == "uri" - and default_example_output - and default_example_output[0] - and default_example_output[0] - .lower() - .endswith((".png", ".jpg", ".jpeg", ".gif", ".webp")) - ): - return "IMAGE" - return "STRING" - - -def create_comfyui_node(schemas, model_info): - author = model_info["owner"] - name = model_info["name"] - version = model_info["latest_version"]["id"] - - replicate_model = f"{author}/{name}:{version}" - node_name = f"Replicate {author}/{name}" - input_schema = schemas["components"]["schemas"]["Input"] - return_type = get_return_type(schemas, model_info) +def create_comfyui_node(schema): + replicate_model, node_name = name_and_version(schema) + return_type = get_return_type(schema) class ReplicateToComfyUI: @classmethod @@ -144,13 +27,13 @@ def create_comfyui_node(schemas, model_info): @classmethod def INPUT_TYPES(cls): - return convert_schema_to_comfyui(input_schema, schemas) + return schema_to_comfyui_input_types(schema) RETURN_TYPES = (return_type,) FUNCTION = "run_openapi_to_comfyui" CATEGORY = "Replicate" - def handle_image_input(self, image): + def convert_image_to_base64(self, image): if isinstance(image, torch.Tensor): image = image.permute(0, 3, 1, 2).squeeze(0) to_pil = transforms.ToPILImage() @@ -164,18 +47,7 @@ def create_comfyui_node(schemas, model_info): img_str = base64.b64encode(buffer.getvalue()).decode() return f"data:image/png;base64,{img_str}" - def run_openapi_to_comfyui(self, **kwargs): - # Convert IMAGE inputs to base64 URI - for key, value in kwargs.items(): - if value is not None: - input_type = ( - self.INPUT_TYPES()["required"].get(key, (None,))[0] - or self.INPUT_TYPES().get("optional", {}).get(key, (None,))[0] - ) - if input_type == "IMAGE": - kwargs[key] = self.handle_image_input(value) - - # Truncate any base64 values in kwargs for logging + def log_input(self, kwargs): truncated_kwargs = { k: v[:20] + "..." if isinstance(v, str) and v.startswith("data:image") @@ -183,43 +55,54 @@ def create_comfyui_node(schemas, model_info): for k, v in kwargs.items() } print(f"Running {replicate_model} with {truncated_kwargs}") + + def handle_image_output(self, output): + # Handle both string and list outputs + output_list = [output] if isinstance(output, str) else list(output) + if output_list: + output_tensors = [] + transform = transforms.ToTensor() + for image_url in output_list: + response = requests.get(image_url) + if response.status_code == 200: + image = Image.open(BytesIO(response.content)) + if image.mode != "RGB": + image = image.convert("RGB") + + tensor_image = transform(image) + tensor_image = tensor_image.unsqueeze(0) + tensor_image = tensor_image.permute(0, 2, 3, 1).cpu().float() + output_tensors.append(tensor_image) + else: + print( + f"Failed to download image. Status code: {response.status_code}" + ) + # Combine all tensors into a single batch if multiple images + return ( + torch.cat(output_tensors, dim=0) + if len(output_tensors) > 1 + else output_tensors[0] + ) + else: + print("No output received from the model") + return None + + def run_openapi_to_comfyui(self, **kwargs): + for key, value in kwargs.items(): + if value is not None: + input_type = ( + self.INPUT_TYPES()["required"].get(key, (None,))[0] + or self.INPUT_TYPES().get("optional", {}).get(key, (None,))[0] + ) + if input_type == "IMAGE": + kwargs[key] = self.convert_image_to_base64(value) + + self.log_input(kwargs) output = replicate.run(replicate_model, input=kwargs) print(f"Output: {output}") if return_type == "IMAGE": - # Convert generator to list - output_list = list(output) - if output_list: - output_tensors = [] - transform = transforms.ToTensor() - for image_url in output_list: - # Download the image from the URL - response = requests.get(image_url) - if response.status_code == 200: - image = Image.open(BytesIO(response.content)) - # Convert image to RGB if it's not already - if image.mode != "RGB": - image = image.convert("RGB") - # Convert to tensor and reshape - tensor_image = transform(image) - tensor_image = tensor_image.unsqueeze(0) - tensor_image = ( - tensor_image.permute(0, 2, 3, 1).cpu().float() - ) - output_tensors.append(tensor_image) - else: - print( - f"Failed to download image. Status code: {response.status_code}" - ) - # Combine all tensors into a single batch if multiple images - output = ( - torch.cat(output_tensors, dim=0) - if len(output_tensors) > 1 - else output_tensors[0] - ) - else: - print("No output received from the model") - output = None + output = self.handle_image_output(output) else: output = "".join(list(output)).strip() @@ -236,9 +119,7 @@ def create_comfyui_nodes_from_schemas(schemas_dir): if schema_file.endswith(".json"): with open(os.path.join(schemas_dir_path, schema_file), "r") as f: schema = json.load(f) - openapi_schema = schema["latest_version"]["openapi_schema"] - model_info = schema - node_name, node_class = create_comfyui_node(openapi_schema, model_info) + node_name, node_class = create_comfyui_node(schema) nodes[node_name] = node_class return nodes diff --git a/schema_to_node.py b/schema_to_node.py new file mode 100644 index 0000000..b60a109 --- /dev/null +++ b/schema_to_node.py @@ -0,0 +1,141 @@ +def convert_to_comfyui_input_type(openapi_type, openapi_format=None): + type_mapping = { + "string": "STRING", + "integer": "INT", + "number": "FLOAT", + "boolean": "BOOLEAN", + } + if openapi_type == "string" and openapi_format == "uri": + return "IMAGE" + return type_mapping.get(openapi_type, "STRING") + + +def name_and_version(schema): + author = schema["owner"] + name = schema["name"] + version = schema["latest_version"]["id"] + replicate_model = f"{author}/{name}:{version}" + node_name = f"Replicate {author}/{name}" + return replicate_model, node_name + + +def resolve_schema(prop_data, openapi_schemas): + if "$ref" in prop_data: + ref_path = prop_data["$ref"].split("/") + current = openapi_schemas + for path in ref_path[1:]: # Skip the first '#' element + current = current[path] + return current + return prop_data + + +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": {}} + + required_props = input_schema.get("required", []) + + for prop_name, prop_data in input_schema["properties"].items(): + prop_data = resolve_schema(prop_data, openapi_schema) + default_value = prop_data.get("default", None) + + if "allOf" in prop_data: + prop_data = resolve_schema(prop_data["allOf"][0], openapi_schema) + + if "enum" in prop_data: + input_type = prop_data["enum"] + elif "type" in prop_data: + input_type = convert_to_comfyui_input_type( + prop_data["type"], prop_data.get("format") + ) + else: + input_type = "STRING" + + input_config = {"default": default_value} + + if "minimum" in prop_data: + input_config["min"] = prop_data["minimum"] + if "maximum" in prop_data: + input_config["max"] = prop_data["maximum"] + if input_type == "FLOAT": + input_config["step"] = 0.01 + input_config["round"] = 0.001 + + if "prompt" in prop_name and prop_data.get("type") == "string": + input_config["multiline"] = True + + # Meta prompt_template needs `{prompt}` to be sent through + if "template" not in prop_name: + input_config["dynamicPrompts"] = True + + if prop_name in required_props: + input_types["required"][prop_name] = (input_type, input_config) + else: + input_types["optional"][prop_name] = (input_type, input_config) + + input_types["optional"]["force_rerun"] = ("BOOLEAN", {"default": False}) + + return reorder_input_types(input_types, input_schema) + + +def reorder_input_types(input_types, input_schema): + ordered_input_types = {"required": {}, "optional": {}} + sorted_properties = sorted( + input_schema["properties"].items(), + key=lambda x: x[1].get("x-order", float("inf")), + ) + + for prop_name, _ in sorted_properties: + if prop_name in input_types["required"]: + ordered_input_types["required"][prop_name] = input_types["required"][ + prop_name + ] + elif prop_name in input_types["optional"]: + ordered_input_types["optional"][prop_name] = input_types["optional"][ + prop_name + ] + + # force_rerun is always at the end of optional inputs + ordered_input_types["optional"]["force_rerun"] = input_types["optional"][ + "force_rerun" + ] + + return ordered_input_types + + +def get_return_type(schema): + image_extensions = (".png", ".jpg", ".jpeg", ".gif", ".webp") + openapi_schema = schema["latest_version"]["openapi_schema"] + output_schema = openapi_schema["components"]["schemas"].get("Output") + default_example = schema.get("default_example") + default_example_output = default_example.get("output") if default_example else None + + if isinstance( + default_example_output, str + ) and default_example_output.lower().endswith(image_extensions): + return "IMAGE" + elif ( + isinstance(default_example_output, list) + and default_example_output + and isinstance(default_example_output[0], str) + and default_example_output[0].lower().endswith(image_extensions) + ): + return "IMAGE" + + if output_schema: + if ( + output_schema.get("type") == "string" + and output_schema.get("format") == "uri" + ): + # Handle single image output + return "IMAGE" + elif ( + output_schema.get("type") == "array" + and output_schema["items"].get("type") == "string" + and output_schema["items"].get("format") == "uri" + ): + # Handle multiple image output + return "IMAGE" + + return "STRING"