add fibo pro nodes
This commit is contained in:
+12
-6
@@ -15,8 +15,10 @@ from .nodes import (
|
||||
TailoredPortraitNode,
|
||||
ReimagineNode,
|
||||
AttributionByImageNode,
|
||||
FiboGenerateNode,
|
||||
FiboRefineAndRegenerateNode,
|
||||
GenerateImageNodeV2,
|
||||
GenerateImageProNodeV2,
|
||||
RefineImageNodeV2,
|
||||
RefineImageProNodeV2,
|
||||
)
|
||||
|
||||
# Map the node class to a name used internally by ComfyUI
|
||||
@@ -37,8 +39,10 @@ NODE_CLASS_MAPPINGS = {
|
||||
"Text2ImageHDNode": Text2ImageHDNode,
|
||||
"ReimagineNode": ReimagineNode,
|
||||
"AttributionByImageNode": AttributionByImageNode,
|
||||
"Text2ImageFiboGenerateNode": FiboGenerateNode,
|
||||
"Text2ImageFiboRefineAndRegenerateNode": FiboRefineAndRegenerateNode,
|
||||
"GenerateImageNodeV2": GenerateImageNodeV2,
|
||||
"RefineImageNodeV2": RefineImageNodeV2,
|
||||
"GenerateImageProNodeV2":GenerateImageProNodeV2,
|
||||
"RefineImageProNodeV2":RefineImageProNodeV2
|
||||
}
|
||||
# Map the node display name to the one shown in the ComfyUI node interface
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -58,6 +62,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Text2ImageHDNode": "Bria Text2Image HD",
|
||||
"ReimagineNode": "Bria Reimagine",
|
||||
"AttributionByImageNode": "Attribution By Image Node",
|
||||
"Text2ImageFiboGenerateNode": "Bria FIBO - Generate",
|
||||
"Text2ImageFiboRefineAndRegenerateNode": "Bria FIBO - Refine and Regenerate",
|
||||
"GenerateImageNodeV2": "Generate Image",
|
||||
"RefineImageNodeV2": "Refine and Regenerate Image",
|
||||
"GenerateImageProNodeV2":"Generate Image - Pro",
|
||||
"RefineImageProNodeV2":"Refine and Regenerate Image - Pro"
|
||||
}
|
||||
|
||||
+2
-2
@@ -14,5 +14,5 @@ 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 .fibo_generate_node import FiboGenerateNode
|
||||
from .fibo_refine_node import FiboRefineAndRegenerateNode
|
||||
from .generate_image_node_v2 import GenerateImageNodeV2, GenerateImageProNodeV2
|
||||
from .refine_image_node_v2 import RefineImageNodeV2, RefineImageProNodeV2
|
||||
@@ -9,9 +9,13 @@ from .common import (
|
||||
)
|
||||
|
||||
|
||||
class FiboGenerateNode:
|
||||
class _BaseGenerateImageNodeV2:
|
||||
"""Base class for image generation nodes (standard & pro)."""
|
||||
|
||||
api_url = None # Each subclass must define its API endpoint
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"api_token": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
@@ -25,7 +29,7 @@ class FiboGenerateNode:
|
||||
["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}),
|
||||
"steps_num": ("INT", {"default": 50, "min": 20, "max": 50}),
|
||||
"guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}),
|
||||
"seed": ("INT", {"default": 123456}),
|
||||
},
|
||||
@@ -36,8 +40,37 @@ class FiboGenerateNode:
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/image/generate"
|
||||
def _validate_token(self, api_token: str):
|
||||
if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API token.")
|
||||
|
||||
def _build_payload(
|
||||
self,
|
||||
prompt,
|
||||
mode,
|
||||
negative_prompt,
|
||||
aspect_ratio,
|
||||
steps_num,
|
||||
guidance_scale,
|
||||
seed,
|
||||
image=None,
|
||||
):
|
||||
payload = {
|
||||
"prompt": prompt,
|
||||
"mode": mode,
|
||||
"negative_prompt": negative_prompt,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"steps_num": steps_num,
|
||||
"guidance_scale": guidance_scale,
|
||||
"seed": seed,
|
||||
}
|
||||
|
||||
if image is not None:
|
||||
if isinstance(image, torch.Tensor):
|
||||
image = preprocess_image(image)
|
||||
payload["images"] = [image_to_base64(image)]
|
||||
|
||||
return payload
|
||||
|
||||
def execute(
|
||||
self,
|
||||
@@ -51,36 +84,26 @@ class FiboGenerateNode:
|
||||
seed,
|
||||
image=None,
|
||||
):
|
||||
|
||||
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
|
||||
self._validate_token(api_token)
|
||||
payload = self._build_payload(
|
||||
prompt,
|
||||
mode,
|
||||
negative_prompt,
|
||||
aspect_ratio,
|
||||
steps_num,
|
||||
guidance_scale,
|
||||
seed,
|
||||
image,
|
||||
)
|
||||
|
||||
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:
|
||||
if response.status_code in (200, 202):
|
||||
print(
|
||||
"Initial FIBO generate request successful, polling for completion..."
|
||||
f"Initial request successful to {self.api_url}, polling for completion..."
|
||||
)
|
||||
response_dict = response.json()
|
||||
status_url = response_dict.get("status_url")
|
||||
@@ -91,25 +114,33 @@ class FiboGenerateNode:
|
||||
|
||||
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")
|
||||
|
||||
# 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}"
|
||||
)
|
||||
|
||||
raise Exception(
|
||||
f"Error: API request failed with status code {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
|
||||
|
||||
class GenerateImageNodeV2(_BaseGenerateImageNodeV2):
|
||||
"""Standard Image Generation Node"""
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/image/generate"
|
||||
|
||||
|
||||
class GenerateImageProNodeV2(_BaseGenerateImageNodeV2):
|
||||
"""Pro Image Generation Node"""
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/image/generate/pro"
|
||||
@@ -1,11 +1,14 @@
|
||||
import requests
|
||||
|
||||
from .common import poll_status_until_completed
|
||||
|
||||
|
||||
class FiboRefineAndRegenerateNode:
|
||||
class _BaseRefineImageNodeV2:
|
||||
"""Base class for refine image nodes (standard & pro)."""
|
||||
|
||||
api_url = None # Must be overridden by subclasses
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"api_token": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
@@ -19,7 +22,7 @@ class FiboRefineAndRegenerateNode:
|
||||
["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}),
|
||||
"steps_num": ("INT", {"default": 50, "min": 20, "max": 50}),
|
||||
"guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}),
|
||||
"seed": ("INT", {"default": 123456}),
|
||||
},
|
||||
@@ -30,8 +33,31 @@ class FiboRefineAndRegenerateNode:
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate"
|
||||
def _validate_token(self, api_token: str):
|
||||
if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API token.")
|
||||
|
||||
def _build_payload(
|
||||
self,
|
||||
prompt,
|
||||
structured_prompt,
|
||||
mode,
|
||||
negative_prompt,
|
||||
aspect_ratio,
|
||||
steps_num,
|
||||
guidance_scale,
|
||||
seed,
|
||||
):
|
||||
return {
|
||||
"prompt": prompt,
|
||||
"mode": mode,
|
||||
"negative_prompt": negative_prompt,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"steps_num": steps_num,
|
||||
"guidance_scale": guidance_scale,
|
||||
"seed": seed,
|
||||
"structured_prompt": structured_prompt,
|
||||
}
|
||||
|
||||
def execute(
|
||||
self,
|
||||
@@ -45,31 +71,25 @@ class FiboRefineAndRegenerateNode:
|
||||
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,
|
||||
"structured_prompt":structured_prompt
|
||||
}
|
||||
self._validate_token(api_token)
|
||||
payload = self._build_payload(
|
||||
prompt,
|
||||
structured_prompt,
|
||||
mode,
|
||||
negative_prompt,
|
||||
aspect_ratio,
|
||||
steps_num,
|
||||
guidance_scale,
|
||||
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 FIBO Refine and Regenerate request successful, polling for completion..."
|
||||
)
|
||||
if response.status_code in (200, 202):
|
||||
print(f"Initial refine request successful to {self.api_url}, polling for completion...")
|
||||
response_dict = response.json()
|
||||
status_url = response_dict.get("status_url")
|
||||
request_id = response_dict.get("request_id")
|
||||
@@ -79,20 +99,29 @@ class FiboRefineAndRegenerateNode:
|
||||
|
||||
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)
|
||||
structured_prompt = result.get("structured_prompt", "")
|
||||
used_seed = result.get("seed", seed)
|
||||
|
||||
return (structured_prompt, used_seed)
|
||||
else:
|
||||
raise Exception(
|
||||
f"Error: API request failed with status code {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
raise Exception(
|
||||
f"Error: API request failed with status code {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
|
||||
|
||||
class RefineImageNodeV2(_BaseRefineImageNodeV2):
|
||||
"""Standard Refine Image Node"""
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate"
|
||||
|
||||
|
||||
class RefineImageProNodeV2(_BaseRefineImageNodeV2):
|
||||
"""Pro Refine Image Node"""
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate/pro"
|
||||
Reference in New Issue
Block a user