diff --git a/__init__.py b/__init__.py index 6f6c7d4..7dda549 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,5 @@ -from .nodes import EraserNode, GenFillNode, ShotByTextNode, ShotByImageNode, TailoredGenNode, TailoredModelInfoNode +from .nodes import (EraserNode, GenFillNode, ShotByTextNode, ShotByImageNode, TailoredGenNode, + TailoredModelInfoNode, Text2ImageBaseNode) # Map the node class to a name used internally by ComfyUI NODE_CLASS_MAPPINGS = { "BriaEraser": EraserNode, # Return the class, not an instance @@ -7,6 +8,7 @@ NODE_CLASS_MAPPINGS = { "ShotByImageNode": ShotByImageNode, "BriaTailoredGen": TailoredGenNode, "TailoredModelInfoNode": TailoredModelInfoNode, + "Text2ImageBaseNode": Text2ImageBaseNode, } # Map the node display name to the one shown in the ComfyUI node interface NODE_DISPLAY_NAME_MAPPINGS = { @@ -16,4 +18,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ShotByImageNode": "Bria Shot By Image", "BriaTailoredGen": "Bria Tailored Gen", "TailoredModelInfoNode": "Bria Tailored Model Info", + "Text2ImageBaseNode": "Bria Text2Image Base", } diff --git a/nodes/__init__.py b/nodes/__init__.py index 6218cbe..0cb043e 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -3,4 +3,5 @@ from .generative_fill_node import GenFillNode from .shot_by_text_node import ShotByTextNode from .shot_by_image_node import ShotByImageNode from .tailored_gen_node import TailoredGenNode -from .tailored_model_info_node import TailoredModelInfoNode \ No newline at end of file +from .tailored_model_info_node import TailoredModelInfoNode +from .text_2_image_base_node import Text2ImageBaseNode \ No newline at end of file diff --git a/nodes/text_2_image_base_node.py b/nodes/text_2_image_base_node.py new file mode 100644 index 0000000..cb7fc29 --- /dev/null +++ b/nodes/text_2_image_base_node.py @@ -0,0 +1,82 @@ +import requests + +from .common import postprocess_image, preprocess_image, image_to_base64 + + +class Text2ImageBaseNode(): + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "api_key": ("STRING", ), + "prompt": ("STRING",), + }, + "optional": { + "aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], {"default": "4:3"}), + "seed": ("INT", {"default": -1}), + "negative_prompt": ("STRING", {"default": ""}), + "steps_num": ("INT", {"default": 30}), + "prompt_enhancement": ("INT", {"default": 0}), + "text_guidance_scale": ("INT", {"default": 5}), + "medium": (["photography", "art", "none"], {"default": "none"}), + "guidance_method_1": (["controlnet_canny", "controlnet_depth", "controlnet_recoloring", "controlnet_color_grid"], {"default": "controlnet_canny"}), + "guidance_method_1_scale": ("FLOAT", {"default": 1.0}), + "guidance_method_1_image": ("IMAGE", ), + "guidance_method_2": (["controlnet_canny", "controlnet_depth", "controlnet_recoloring", "controlnet_color_grid"], {"default": "controlnet_canny"}), + "guidance_method_2_scale": ("FLOAT", {"default": 1.0}), + "guidance_method_2_image": ("IMAGE", ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("output_image",) + CATEGORY = "API Nodes" + FUNCTION = "execute" # This is the method that will be executed + + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/base/2.3" #"http://0.0.0.0:5000/v1/text-to-image/base/2.3" + + def execute( + self, api_key, prompt, aspect_ratio, seed, negative_prompt, + steps_num, prompt_enhancement, text_guidance_scale, medium, + guidance_method_1=None, guidance_method_1_scale=None, guidance_method_1_image=None, + guidance_method_2=None, guidance_method_2_scale=None, guidance_method_2_image=None, + ): + prompt_enhancement = bool(prompt_enhancement) + payload = { + "prompt": prompt, + "num_results": 1, + "aspect_ratio": aspect_ratio, + "sync": True, + "seed": seed, + "negative_prompt": negative_prompt, + "steps_num": steps_num, + "text_guidance_scale": text_guidance_scale, + "prompt_enhancement": prompt_enhancement, + } + if medium != "none": + payload["medium"] = medium + if guidance_method_1_image is not None: + guidance_method_1_image = preprocess_image(guidance_method_1_image) + guidance_method_1_image = image_to_base64(guidance_method_1_image) + payload["guidance_method_1"] = guidance_method_1 + payload["guidance_method_1_scale"] = guidance_method_1_scale + payload["guidance_method_1_image_file"] = guidance_method_1_image + if guidance_method_2_image is not None: + guidance_method_2_image = preprocess_image(guidance_method_2_image) + guidance_method_2_image = image_to_base64(guidance_method_2_image) + payload["guidance_method_2"] = guidance_method_2 + payload["guidance_method_2_scale"] = guidance_method_2_scale + payload["guidance_method_2_image_file"] = guidance_method_2_image + response = requests.post( + self.api_url, + json=payload, + headers={"api_token": api_key} + ) + if response.status_code == 200: + response_dict = response.json() + image_response = requests.get(response_dict['result'][0]["urls"][0]) + result_image = postprocess_image(image_response.content) + return (result_image,) + else: + raise Exception(f"Error: API request failed with status code {response.status_code} and text {response.text}")