diff --git a/node.py b/node.py index 7c70fd8..6874479 100644 --- a/node.py +++ b/node.py @@ -13,6 +13,7 @@ from .schema_to_node import ( schema_to_comfyui_input_types, get_return_type, name_and_version, + inputs_that_need_arrays, ) @@ -33,7 +34,17 @@ def create_comfyui_node(schema): FUNCTION = "run_replicate_model" CATEGORY = "Replicate" - def convert_image_to_base64(self, image): + def convert_input_images_to_base64(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.image_to_base64(value) + + def image_to_base64(self, image): if isinstance(image, torch.Tensor): image = image.permute(0, 3, 1, 2).squeeze(0) to_pil = transforms.ToPILImage() @@ -47,6 +58,18 @@ def create_comfyui_node(schema): img_str = base64.b64encode(buffer.getvalue()).decode() return f"data:image/png;base64,{img_str}" + def handle_array_inputs(self, kwargs): + array_inputs = inputs_that_need_arrays(schema) + for input_name in array_inputs: + if input_name in kwargs: + if isinstance(kwargs[input_name], str): + if kwargs[input_name] == "": + kwargs[input_name] = [] + else: + kwargs[input_name] = kwargs[input_name].split("\n") + else: + kwargs[input_name] = [kwargs[input_name]] + def log_input(self, kwargs): truncated_kwargs = { k: v[:20] + "..." @@ -88,15 +111,8 @@ def create_comfyui_node(schema): return None def run_replicate_model(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.handle_array_inputs(kwargs) + self.convert_input_images_to_base64(kwargs) self.log_input(kwargs) output = replicate.run(replicate_model, input=kwargs) print(f"Output: {output}") diff --git a/schema_to_node.py b/schema_to_node.py index cdd96a6..944bfec 100644 --- a/schema_to_node.py +++ b/schema_to_node.py @@ -108,6 +108,16 @@ def order_inputs(input_types, input_schema): return ordered_input_types +def inputs_that_need_arrays(schema): + openapi_schema = schema["latest_version"]["openapi_schema"] + input_schema = openapi_schema["components"]["schemas"]["Input"] + array_inputs = [] + for prop_name, prop_data in input_schema["properties"].items(): + if prop_data.get("type") == "array": + array_inputs.append(prop_name) + + return array_inputs + def get_return_type(schema): image_extensions = (".png", ".jpg", ".jpeg", ".gif", ".webp") openapi_schema = schema["latest_version"]["openapi_schema"]