WAI-4116
This commit is contained in:
+27
-5
@@ -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
@@ -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
|
||||
@@ -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}")
|
||||
@@ -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}")
|
||||
Reference in New Issue
Block a user