From 327ace887bb34ff3adaca21388af8cf648b5ea4a Mon Sep 17 00:00:00 2001 From: Ubuntu Date: Tue, 14 Oct 2025 08:57:45 +0000 Subject: [PATCH] WAI-4116 --- __init__.py | 32 ++++++- nodes/__init__.py | 4 +- nodes/text_2_image_gaia_generate_node.py | 115 +++++++++++++++++++++++ nodes/text_2_image_gaia_refine_node.py | 101 ++++++++++++++++++++ 4 files changed, 246 insertions(+), 6 deletions(-) create mode 100644 nodes/text_2_image_gaia_generate_node.py create mode 100644 nodes/text_2_image_gaia_refine_node.py diff --git a/__init__.py b/__init__.py index 48ade2d..904400f 100644 --- a/__init__.py +++ b/__init__.py @@ -1,6 +1,24 @@ -from .nodes import (EraserNode, GenFillNode, ImageExpansionNode, ReplaceBgNode, RmbgNode, RemoveForegroundNode, ShotByTextNode, ShotByImageNode, TailoredGenNode, - TailoredModelInfoNode, Text2ImageBaseNode, Text2ImageFastNode, Text2ImageHDNode, TailoredPortraitNode, - ReimagineNode, AttributionByImageNode) +from .nodes import ( + EraserNode, + GenFillNode, + ImageExpansionNode, + ReplaceBgNode, + RmbgNode, + RemoveForegroundNode, + ShotByTextNode, + ShotByImageNode, + TailoredGenNode, + TailoredModelInfoNode, + Text2ImageBaseNode, + Text2ImageFastNode, + Text2ImageHDNode, + TailoredPortraitNode, + ReimagineNode, + AttributionByImageNode, + Text2ImageGaiaGenerateNode, + Text2ImageGaiaRefineAndRegenerateNode, +) + # Map the node class to a name used internally by ComfyUI NODE_CLASS_MAPPINGS = { "BriaEraser": EraserNode, # Return the class, not an instance @@ -18,7 +36,9 @@ NODE_CLASS_MAPPINGS = { "Text2ImageFastNode": Text2ImageFastNode, "Text2ImageHDNode": Text2ImageHDNode, "ReimagineNode": ReimagineNode, - "AttributionByImageNode":AttributionByImageNode + "AttributionByImageNode": AttributionByImageNode, + "Text2ImageGaiaGenerateNode": Text2ImageGaiaGenerateNode, + "Text2ImageGaiaRefineAndRegenerateNode": Text2ImageGaiaRefineAndRegenerateNode, } # Map the node display name to the one shown in the ComfyUI node interface NODE_DISPLAY_NAME_MAPPINGS = { @@ -37,5 +57,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "Text2ImageFastNode": "Bria Text2Image Fast", "Text2ImageHDNode": "Bria Text2Image HD", "ReimagineNode": "Bria Reimagine", - "AttributionByImageNode":"Attribution By Image Node" + "AttributionByImageNode": "Attribution By Image Node", + "Text2ImageGaiaGenerateNode": "Bria GAIA - Generate", + "Text2ImageGaiaRefineAndRegenerateNode": "Bria GAIA - Refine and Regenerate", } diff --git a/nodes/__init__.py b/nodes/__init__.py index 3440ffc..bf09104 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -13,4 +13,6 @@ from .text_2_image_base_node import Text2ImageBaseNode from .text_2_image_fast_node import Text2ImageFastNode from .text_2_image_hd_node import Text2ImageHDNode from .reimagine_node import ReimagineNode -from .attribution_by_image_node import AttributionByImageNode \ No newline at end of file +from .attribution_by_image_node import AttributionByImageNode +from .text_2_image_gaia_generate_node import Text2ImageGaiaGenerateNode +from .text_2_image_gaia_refine_node import Text2ImageGaiaRefineAndRegenerateNode \ No newline at end of file diff --git a/nodes/text_2_image_gaia_generate_node.py b/nodes/text_2_image_gaia_generate_node.py new file mode 100644 index 0000000..f4028f3 --- /dev/null +++ b/nodes/text_2_image_gaia_generate_node.py @@ -0,0 +1,115 @@ +import requests +import torch + +from .common import ( + postprocess_image, + preprocess_image, + image_to_base64, + poll_status_until_completed, +) + + +class Text2ImageGaiaGenerateNode: + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "api_token": ("STRING", {"default": "BRIA_API_TOKEN"}), + "prompt": ("STRING", {"multiline": True}), + }, + "optional": { + "mode": (["GAIA"], {"default": "GAIA"}), + "negative_prompt": ("STRING", {"default": ""}), + "image": ("IMAGE",), + "aspect_ratio": ( + ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], + {"default": "1:1"}, + ), + "steps_num": ("INT", {"default": 30, "min": 20, "max": 50}), + "guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}), + "seed": ("INT", {"default": 1}), + }, + } + + RETURN_TYPES = ("IMAGE", "STRING", "INT") + RETURN_NAMES = ("image", "structured_prompt", "seed") + CATEGORY = "API Nodes" + FUNCTION = "execute" + + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/image/generate" + + def execute( + self, + api_token, + prompt, + mode, + negative_prompt, + image, + aspect_ratio, + steps_num, + guidance_scale, + seed, + ): + + if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": + raise Exception("Please insert a valid API token.") + + payload = { + "prompt": prompt, + "mode": mode, + "negative_prompt": negative_prompt, + "aspect_ratio": aspect_ratio, + "steps_num": steps_num, + "guidance_scale": guidance_scale, + "seed": seed, + } + + # Add image if provided + if image is not None: + if isinstance(image, torch.Tensor): + image = preprocess_image(image) + image_base64 = image_to_base64(image) + payload["image"] = image_base64 + + headers = {"Content-Type": "application/json", "api_token": api_token} + + try: + # Send initial request to get status URL + response = requests.post(self.api_url, json=payload, headers=headers) + + if response.status_code == 200 or response.status_code == 202: + print( + "Initial GAIA generate request successful, polling for completion..." + ) + response_dict = response.json() + status_url = response_dict.get("status_url") + request_id = response_dict.get("request_id") + + if not status_url: + raise Exception("No status_url returned from API") + + print(f"Request ID: {request_id}, Status URL: {status_url}") + + # Poll for completion + final_response = poll_status_until_completed(status_url, api_token) + + # Extract results + result = final_response.get("result", {}) + print(result) + result_image_url = result.get("image_url") + structured_prompt = result.get("structured_prompt", "") + used_seed = result.get("seed", seed) + + # Download and process the result image + image_response = requests.get(result_image_url) + result_image = postprocess_image(image_response.content) + + return (result_image, structured_prompt, used_seed) + else: + raise Exception( + f"Error: API request failed with status code {response.status_code} {response.text}" + ) + + except Exception as e: + raise Exception(f"{e}") diff --git a/nodes/text_2_image_gaia_refine_node.py b/nodes/text_2_image_gaia_refine_node.py new file mode 100644 index 0000000..18ec4f8 --- /dev/null +++ b/nodes/text_2_image_gaia_refine_node.py @@ -0,0 +1,101 @@ +import requests + +from .common import postprocess_image, poll_status_until_completed + + +class Text2ImageGaiaRefineAndRegenerateNode: + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "api_token": ("STRING", {"default": "BRIA_API_TOKEN"}), + "prompt": ("STRING",), + "structured_prompt": ("STRING",), + }, + "optional": { + "mode": (["GAIA"], {"default": "GAIA"}), + "negative_prompt": ("STRING", {"default": ""}), + "aspect_ratio": ( + ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], + {"default": "1:1"}, + ), + "steps_num": ("INT", {"default": 30, "min": 20, "max": 50}), + "guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}), + "seed": ("INT", {"default": 1}), + }, + } + + RETURN_TYPES = ("IMAGE", "STRING", "INT") + RETURN_NAMES = ("image", "structured_prompt", "seed") + CATEGORY = "API Nodes" + FUNCTION = "execute" + + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate" + + def execute( + self, + api_token, + prompt, + mode, + negative_prompt, + aspect_ratio, + steps_num, + guidance_scale, + seed, + ): + + if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": + raise Exception("Please insert a valid API token.") + + payload = { + "prompt": prompt, + "mode": mode, + "negative_prompt": negative_prompt, + "aspect_ratio": aspect_ratio, + "steps_num": steps_num, + "guidance_scale": guidance_scale, + "seed": seed, + } + + headers = {"Content-Type": "application/json", "api_token": api_token} + + try: + # Send initial request to get status URL + response = requests.post(self.api_url, json=payload, headers=headers) + + if response.status_code == 200 or response.status_code == 202: + print( + "Initial GAIA Refine and Regenerate request successful, polling for completion..." + ) + response_dict = response.json() + status_url = response_dict.get("status_url") + request_id = response_dict.get("request_id") + + if not status_url: + raise Exception("No status_url returned from API") + + print(f"Request ID: {request_id}, Status URL: {status_url}") + + # Poll for completion + final_response = poll_status_until_completed(status_url, api_token) + + # Extract results + result = final_response.get("result", {}) + print(result) + result_image_url = result.get("image_url") + structured_prompt = result.get("structured_prompt", "") + used_seed = result.get("seed", seed) + + # Download and process the result image + image_response = requests.get(result_image_url) + result_image = postprocess_image(image_response.content) + + return (result_image, structured_prompt, used_seed) + else: + raise Exception( + f"Error: API request failed with status code {response.status_code} {response.text}" + ) + + except Exception as e: + raise Exception(f"{e}")