This commit is contained in:
Ubuntu
2025-10-14 08:57:45 +00:00
parent 9960a93044
commit 327ace887b
4 changed files with 246 additions and 6 deletions
+27 -5
View File
@@ -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",
}
+3 -1
View File
@@ -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
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
+115
View File
@@ -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}")
+101
View File
@@ -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}")