Merge pull request #26 from Bria-AI/WAI-4049

WAI-4049
This commit is contained in:
Yazan Numoor
2025-10-23 09:48:24 +03:00
committed by GitHub
17 changed files with 1020 additions and 287 deletions
+53 -7
View File
@@ -1,6 +1,32 @@
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,
ShotByTextOriginalNode,
ShotByImageOriginalNode,
TailoredGenNode,
TailoredModelInfoNode,
Text2ImageBaseNode,
Text2ImageFastNode,
Text2ImageHDNode,
TailoredPortraitNode,
ReimagineNode,
ShotByTextAutomaticNode,
ShotByImageManualPaddingNode,
ShotByImageAutomaticAspectRatioNode,
ShotByImageCustomCoordinatesNode,
ShotByImageManualPlacementNode,
ShotByImageAutomaticNode,
ShotByTextAutomaticAspectRatioNode,
ShotByTextManualPlacementNode,
ShotByTextManualPaddingNode,
ShotByTextCustomCoordinatesNode,
AttributionByImageNode
)
# Map the node class to a name used internally by ComfyUI
NODE_CLASS_MAPPINGS = {
"BriaEraser": EraserNode, # Return the class, not an instance
@@ -9,8 +35,18 @@ NODE_CLASS_MAPPINGS = {
"ReplaceBgNode": ReplaceBgNode,
"RmbgNode": RmbgNode,
"RemoveForegroundNode": RemoveForegroundNode,
"ShotByTextNode": ShotByTextNode,
"ShotByImageNode": ShotByImageNode,
"ShotByTextOriginal": ShotByTextOriginalNode,
"ShotByImageOriginal": ShotByImageOriginalNode,
"ShotByTextAutomatic": ShotByTextAutomaticNode,
"ShotByTextManualPlacement": ShotByTextManualPlacementNode,
"ShotByTextCustomCoordinates": ShotByTextCustomCoordinatesNode,
"ShotByTextManualPadding": ShotByTextManualPaddingNode,
"ShotByTextAutomaticAspectRatio": ShotByTextAutomaticAspectRatioNode,
"ShotByImageAutomatic": ShotByImageAutomaticNode,
"ShotByImageManualPlacement": ShotByImageManualPlacementNode,
"ShotByImageCustomCoordinates": ShotByImageCustomCoordinatesNode,
"ShotByImageManualPadding": ShotByImageManualPaddingNode,
"ShotByImageAutomaticAspectRatio": ShotByImageAutomaticAspectRatioNode,
"BriaTailoredGen": TailoredGenNode,
"TailoredModelInfoNode": TailoredModelInfoNode,
"TailoredPortraitNode": TailoredPortraitNode,
@@ -28,8 +64,18 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ReplaceBgNode": "Bria Replace Background",
"RmbgNode": "Bria RMBG",
"RemoveForegroundNode": "Bria Remove Foreground",
"ShotByTextNode": "Bria Shot By Text",
"ShotByImageNode": "Bria Shot By Image",
"ShotByTextOriginal": "Shot by Text - Original",
"ShotByImageOriginal": "Shot by Image - Original",
"ShotByTextAutomatic": "Shot by Text - Automatic",
"ShotByTextManualPlacement": "Shot by Text - Manual Placement",
"ShotByTextCustomCoordinates": "Shot by Text - Custom Coordinates",
"ShotByTextManualPadding": "Shot by Text - Manual Padding",
"ShotByTextAutomaticAspectRatio": "Shot by Text - Automatic Aspect Ratio",
"ShotByImageAutomatic": "Shot by Image - Automatic",
"ShotByImageManualPlacement": "Shot by Image - Manual Placement",
"ShotByImageCustomCoordinates": "Shot by Image - Custom Coordinates",
"ShotByImageManualPadding": "Shot by Image - Manual Padding",
"ShotByImageAutomaticAspectRatio": "Shot by Image - Automatic Aspect Ratio",
"BriaTailoredGen": "Bria Tailored Gen",
"TailoredModelInfoNode": "Bria Tailored Model Info",
"TailoredPortraitNode": "Bria Restyle Portrait",
+15 -3
View File
@@ -4,8 +4,6 @@ from .image_expansion_node import ImageExpansionNode
from .replace_bg_node import ReplaceBgNode
from .rmbg_node import RmbgNode
from .remove_foreground_node import RemoveForegroundNode
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
from .tailored_portrait_node import TailoredPortraitNode
@@ -13,4 +11,18 @@ 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 .shot_by_text_node import ShotByTextOriginalNode
from .shot_by_text_automatic_aspect_ratio_node import ShotByTextAutomaticAspectRatioNode
from .shot_by_text_automatic_node import ShotByTextAutomaticNode
from .shot_by_text_custom_coordinates_node import ShotByTextCustomCoordinatesNode
from .shot_by_text_manual_placement_node import ShotByTextManualPlacementNode
from .shot_by_text_manual_padding_node import ShotByTextManualPaddingNode
from .shot_by_image_automatic_aspect_ratio_node import (
ShotByImageAutomaticAspectRatioNode,
)
from .shot_by_image_automatic_node import ShotByImageAutomaticNode
from .shot_by_image_custom_coordinates_node import ShotByImageCustomCoordinatesNode
from .shot_by_image_node import ShotByImageOriginalNode
from .shot_by_image_manual_placement_node import ShotByImageManualPlacementNode
from .shot_by_image_manual_padding_node import ShotByImageManualPaddingNode
from .attribution_by_image_node import AttributionByImageNode
@@ -0,0 +1,46 @@
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageAutomaticAspectRatioNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
input_types["required"]["aspect_ratio"] = (
["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"],
{"default": "1:1"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
aspect_ratio,
api_key,
sync=False,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.AUTOMATIC_ASPECT_RATIO.value,
aspect_ratio=aspect_ratio,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+51
View File
@@ -0,0 +1,51 @@
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageAutomaticNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
return input_types
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE")
RETURN_NAMES = (
"output_image_1",
"output_image_2",
"output_image_3",
"output_image_4",
"output_image_5",
"output_image_6",
"output_image_7",
)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
shot_size,
api_key,
sync=False,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.AUTOMATIC.value,
shot_size=shot_size,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key, Placement_type = PlacementType.AUTOMATIC.value)
@@ -0,0 +1,55 @@
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageCustomCoordinatesNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
input_types["required"]["foreground_image_size"] = (
"STRING",
{"default": "500,500"},
)
input_types["required"]["foreground_image_location"] = (
"STRING",
{"default": "0, 0"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
shot_size,
foreground_image_size,
foreground_image_location,
api_key,
sync=False,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.CUSTOM_COORDINATES.value,
shot_size=shot_size,
foreground_image_size=foreground_image_size,
foreground_image_location=foreground_image_location,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
@@ -0,0 +1,44 @@
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageManualPaddingNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
input_types["required"]["padding_values"] = ("STRING", {"default": "0,0,0,0"})
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
padding_values,
api_key,
sync=False,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.MANUAL_PADDING.value,
padding_values=padding_values,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
@@ -0,0 +1,59 @@
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageManualPlacementNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
input_types["required"]["manual_placement_selection"] = (
[
"upper_left",
"upper_right",
"bottom_left",
"bottom_right",
"right_center",
"left_center",
"upper_center",
"bottom_center",
"center_vertical",
"center_horizontal",
],
{"default": "upper_left"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
shot_size,
manual_placement_selection,
api_key,
sync=False,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.MANUAL_PLACEMENT.value,
shot_size=shot_size,
manual_placement_selection=manual_placement_selection,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+41 -72
View File
@@ -1,72 +1,41 @@
import requests
import torch
from .common import postprocess_image, preprocess_image, image_to_base64
class ShotByImageNode():
@classmethod
def INPUT_TYPES(self):
return {
"required": {
"image": ("IMAGE",), # Input image from another node
"ref_image": ("IMAGE",), # ref image from another node
"enhance_ref_image": ("INT", {"default": 1}),
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value
},
"optional": {
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
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/product/lifestyle_shot_by_image" # Eraser API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, ref_image, api_key, enhance_ref_image, content_moderation):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
# Check if image and mask are tensors, if so, convert to NumPy arrays
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
if isinstance(ref_image, torch.Tensor):
ref_image = preprocess_image(ref_image)
# Convert the image and mask directly to Base64 strings
image_base64 = image_to_base64(image)
ref_image_base64 = image_to_base64(ref_image)
enhance_ref_image = bool(enhance_ref_image)
payload = {
"file": image_base64,
"ref_image_file": ref_image_base64,
"enhance_ref_image": enhance_ref_image,
"placement_type": "original",
"original_quality": True,
"sync": True,
"content_moderation": content_moderation
}
headers = {
"Content-Type": "application/json",
"api_token": f"{api_key}"
}
try:
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
image_response = requests.get(response_dict['result'][0][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}")
except Exception as e:
raise Exception(f"{e}")
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url, PlacementType
class ShotByImageOriginalNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_image_input_types()
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_image_api_url
def execute(
self,
image,
ref_image,
api_key,
sync=True,
enhance_ref_image=True,
ref_image_influence=1.0,
force_rmbg=False,
content_moderation=False,
):
payload = create_image_payload(
image,
ref_image,
api_key,
PlacementType.ORIGINAL.value,
original_quality=True,
sync=sync,
enhance_ref_image=enhance_ref_image,
ref_image_influence=ref_image_influence,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
@@ -0,0 +1,48 @@
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextAutomaticAspectRatioNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
input_types["required"]["aspect_ratio"] = (
["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"],
{"default": "1:1"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
aspect_ratio,
api_key,
sync=False,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.AUTOMATIC_ASPECT_RATIO.value,
aspect_ratio=aspect_ratio,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+53
View File
@@ -0,0 +1,53 @@
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextAutomaticNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
return input_types
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE", "IMAGE")
RETURN_NAMES = (
"output_image_1",
"output_image_2",
"output_image_3",
"output_image_4",
"output_image_5",
"output_image_6",
"output_image_7",
)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
shot_size,
api_key,
sync=False,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.AUTOMATIC.value,
shot_size=shot_size,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key, Placement_type= PlacementType.AUTOMATIC.value)
@@ -0,0 +1,57 @@
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextCustomCoordinatesNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
input_types["required"]["foreground_image_size"] = (
"STRING",
{"default": "500,500"},
)
input_types["required"]["foreground_image_location"] = (
"STRING",
{"default": "0, 0"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
shot_size,
foreground_image_size,
foreground_image_location,
api_key,
sync=False,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.CUSTOM_COORDINATES.value,
shot_size=shot_size,
foreground_image_size=foreground_image_size,
foreground_image_location=foreground_image_location,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+45
View File
@@ -0,0 +1,45 @@
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextManualPaddingNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
input_types["required"]["padding_values"] = ("STRING", {"default": "0,0,0,0"})
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
padding_values,
api_key,
sync=False,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.MANUAL_PADDING.value,
padding_values=padding_values,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
@@ -0,0 +1,62 @@
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextManualPlacementNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
input_types["required"]["shot_size"] = ("STRING", {"default": "1000, 1000"})
input_types["required"]["manual_placement_selection"] = (
[
"upper_left",
"upper_right",
"bottom_left",
"bottom_right",
"right_center",
"left_center",
"upper_center",
"bottom_center",
"center_vertical",
"center_horizontal",
],
{"default": "upper_left"},
)
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
shot_size,
manual_placement_selection,
api_key,
sync=False,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.MANUAL_PLACEMENT.value,
shot_size=shot_size,
manual_placement_selection=manual_placement_selection,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+42 -68
View File
@@ -1,68 +1,42 @@
import requests
import torch
from .common import postprocess_image, preprocess_image, image_to_base64
class ShotByTextNode():
@classmethod
def INPUT_TYPES(self):
return {
"required": {
"image": ("IMAGE",), # Input image from another node
"scene_description": ("STRING",),
"mode": (["base", "fast", "high_control"], {"default": "high_control"}),
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value
},
"optional": {
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
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/product/lifestyle_shot_by_text" # Eraser API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, api_key, scene_description, mode, content_moderation):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
# Check if image and mask are tensors, if so, convert to NumPy arrays
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
image_base64 = image_to_base64(image)
payload = {
"file": image_base64,
"scene_description": scene_description,
"mode": mode,
"placement_type": "original",
"original_quality": True,
"sync": True,
"content_moderation": content_moderation
}
headers = {
"Content-Type": "application/json",
"api_token": f"{api_key}"
}
try:
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
image_response = requests.get(response_dict['result'][0][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}")
except Exception as e:
raise Exception(f"{e}")
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url, PlacementType
class ShotByTextOriginalNode:
@classmethod
def INPUT_TYPES(self):
input_types = get_text_input_types()
return input_types
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("output_image",)
CATEGORY = "API Nodes"
FUNCTION = "execute"
def __init__(self):
self.api_url = shot_by_text_api_url
def execute(
self,
image,
scene_description,
mode,
api_key,
sync=True,
optimize_description=True,
exclude_elements="",
force_rmbg=False,
content_moderation=False,
):
payload = create_text_payload(
image,
api_key,
scene_description,
mode,
PlacementType.ORIGINAL.value,
original_quality=True,
sync=sync,
optimize_description=optimize_description,
exclude_elements=exclude_elements,
force_rmbg=force_rmbg,
content_moderation=content_moderation,
)
return make_api_request(self.api_url, payload, api_key)
+204
View File
@@ -0,0 +1,204 @@
import requests
import torch
from ..common import postprocess_image, preprocess_image, image_to_base64
shot_by_text_api_url = (
"https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_text"
)
shot_by_image_api_url = (
"https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_image"
)
from enum import Enum
class PlacementType(str, Enum):
ORIGINAL = "original"
AUTOMATIC = "automatic"
MANUAL_PLACEMENT = "manual_placement"
MANUAL_PADDING = "manual_padding"
CUSTOM_COORDINATES = "custom_coordinates"
AUTOMATIC_ASPECT_RATIO = "automatic_aspect_ratio"
def validate_api_key(api_key):
"""Validate API key input"""
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
def update_payload_for_placement(placement_type, payload, **kwargs):
if placement_type == PlacementType.AUTOMATIC.value:
payload["shot_size"] = [
int(x.strip()) for x in kwargs.get("shot_size").split(",")
]
elif placement_type == PlacementType.MANUAL_PLACEMENT.value:
payload["shot_size"] = [
int(x.strip()) for x in kwargs.get("shot_size").split(",")
]
payload["manual_placement_selection"] = [
kwargs.get("manual_placement_selection", "upper_left")
]
elif placement_type == PlacementType.CUSTOM_COORDINATES.value:
payload["shot_size"] = [
int(x.strip()) for x in kwargs.get("shot_size").split(",")
]
payload["foreground_image_size"] = [
int(x.strip()) for x in kwargs.get("foreground_image_size").split(",")
]
payload["foreground_image_location"] = [
int(x.strip()) for x in kwargs.get("foreground_image_location").split(",")
]
elif placement_type == PlacementType.MANUAL_PADDING.value:
payload["padding_values"] = [
int(x.strip()) for x in kwargs.get("padding_values").split(",")
]
elif placement_type == PlacementType.AUTOMATIC_ASPECT_RATIO.value:
payload["aspect_ratio"] = kwargs.get("aspect_ratio", "1:1")
elif placement_type == PlacementType.ORIGINAL.value:
payload["original_quality"] = kwargs.get("original_quality", True)
return payload
def create_text_payload(
image, api_key, scene_description, mode, placement_type, **kwargs
):
validate_api_key(api_key)
# Process image
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
image_base64 = image_to_base64(image)
payload = {
"file": image_base64,
"placement_type": placement_type,
"sync": True,
"num_results": 1,
"force_rmbg": kwargs.get("force_rmbg", False),
"content_moderation": kwargs.get("content_moderation", False),
"scene_description": scene_description,
"mode": mode,
"optimize_description": kwargs.get("optimize_description", True),
}
if kwargs.get("exclude_elements", "").strip():
payload["exclude_elements"] = kwargs["exclude_elements"]
payload = update_payload_for_placement(placement_type, payload, **kwargs)
return payload
def create_image_payload(image, ref_image, api_key, placement_type, **kwargs):
"""Create payload for image-based shot nodes"""
validate_api_key(api_key)
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
if isinstance(ref_image, torch.Tensor):
ref_image = preprocess_image(ref_image)
image_base64 = image_to_base64(image)
ref_image_base64 = image_to_base64(ref_image)
# Base payload
payload = {
"file": image_base64,
"ref_image_file": ref_image_base64,
"enhance_ref_image": kwargs.get("enhance_ref_image", True),
"ref_image_influence": kwargs.get("ref_image_influence", 1.0),
"placement_type": placement_type,
"sync": True,
"num_results": 1,
"force_rmbg": kwargs.get("force_rmbg", False),
"content_moderation": kwargs.get("content_moderation", False),
}
payload = update_payload_for_placement(placement_type, payload, **kwargs)
return payload
def make_api_request(api_url, payload, api_key, Placement_type = None):
"""Make API request and return processed image"""
headers = {"Content-Type": "application/json", "api_token": f"{api_key}"}
try:
response = requests.post(api_url, json=payload, headers=headers)
if response.status_code == 200:
print("response is 200")
response_dict = response.json()
if Placement_type == PlacementType.AUTOMATIC.value:
result_images = []
for i, result in enumerate(response_dict.get("result", [])[:7]):
image_url = result[0]
image_response = requests.get(image_url)
processed = postprocess_image(image_response.content)
result_images.append(processed)
# If less than 7 images, pad with None to match ComfyUI return structure
while len(result_images) < 7:
result_images.append(None)
print(result_images)
return tuple(result_images)
image_response = requests.get(response_dict["result"][0][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}{response.text}"
)
except Exception as e:
raise Exception(f"{e}")
def get_common_input_types():
"""Get common input types for all nodes"""
return {
"required": {"api_key": ("STRING", {"default": "BRIA_API_TOKEN"})},
"optional": {
"force_rmbg": ("BOOLEAN", {"default": False}),
"content_moderation": ("BOOLEAN", {"default": False}),
},
}
def get_text_input_types():
"""Get text-specific input types"""
common = get_common_input_types()
common["required"].update(
{
"image": ("IMAGE",),
"scene_description": ("STRING",),
"mode": (["base", "fast", "high_control"], {"default": "fast"}),
}
)
common["optional"].update(
{
"optimize_description": ("BOOLEAN", {"default": True}),
"exclude_elements": ("STRING", {"default": ""}),
}
)
return common
def get_image_input_types():
"""Get image-specific input types"""
common = get_common_input_types()
common["required"].update({"image": ("IMAGE",), "ref_image": ("IMAGE",)})
common["optional"].update(
{
"enhance_ref_image": ("BOOLEAN", {"default": True}),
"ref_image_influence": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
}
)
return common
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-bria-api"
description = "Custom nodes for ComfyUI using BRIA's API."
version = "2.1.2"
version = "2.1.3"
license = {file = "LICENSE"}
[project.urls]
@@ -1,18 +1,47 @@
{
"last_node_id": 42,
"last_link_id": 65,
"id": "1cdd7d4c-58b5-4047-947b-1977ad36d364",
"revision": 0,
"last_node_id": 14,
"last_link_id": 11,
"nodes": [
{
"id": 42,
"id": 6,
"type": "PreviewImage",
"pos": [
1351.83154296875,
26.696861267089844
],
"size": [
399.811279296875,
246
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 1,
"type": "LoadImage",
"pos": {
"0": 591,
"1": 593
},
"size": {
"0": 315,
"1": 314
},
"pos": [
383.92852783203125,
38.40964889526367
],
"size": [
397.5969543457031,
314
],
"flags": {},
"order": 0,
"mode": 0,
@@ -22,10 +51,9 @@
"name": "IMAGE",
"type": "IMAGE",
"links": [
64,
65
],
"slot_index": 0
1,
3
]
},
{
"name": "MASK",
@@ -37,21 +65,21 @@
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"A_bottle_of_perfume.png",
"CAR.png",
"image"
]
},
{
"id": 39,
"id": 2,
"type": "LoadImage",
"pos": {
"0": 600,
"1": 988
},
"size": {
"0": 315,
"1": 314
},
"pos": [
369.0213623046875,
415.34893798828125
],
"size": [
450.83685302734375,
314.0000305175781
],
"flags": {},
"order": 1,
"mode": 0,
@@ -61,9 +89,8 @@
"name": "IMAGE",
"type": "IMAGE",
"links": [
59
],
"slot_index": 0
4
]
},
{
"name": "MASK",
@@ -75,21 +102,21 @@
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"A_red_studio_with_a_shelf__close_up.png",
"seed_713360865.png",
"image"
]
},
{
"id": 15,
"id": 5,
"type": "Note",
"pos": {
"0": 995.4524536132812,
"1": 601.5353393554688
},
"size": {
"0": 306.28387451171875,
"1": 58
},
"pos": [
951.4669189453125,
5.085720062255859
],
"size": [
210,
88
],
"flags": {},
"order": 2,
"mode": 0,
@@ -103,40 +130,14 @@
"bgcolor": "#653"
},
{
"id": 40,
"id": 7,
"type": "PreviewImage",
"pos": {
"0": 1408,
"1": 623
},
"size": [
210,
246
"pos": [
1368.5450439453125,
340.4889221191406
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 62
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 41,
"type": "PreviewImage",
"pos": {
"0": 1408,
"1": 951
},
"size": [
210,
387.0335693359375,
246
],
"flags": {},
@@ -146,25 +147,26 @@
{
"name": "images",
"type": "IMAGE",
"link": 63
"link": 6
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
"widgets_values": []
},
{
"id": 36,
"type": "ShotByTextNode",
"pos": {
"0": 996,
"1": 736
},
"size": {
"0": 315,
"1": 106
},
"id": 3,
"type": "ShotByTextOriginal",
"pos": [
941.8839111328125,
149.14889526367188
],
"size": [
273.388671875,
202
],
"flags": {},
"order": 3,
"mode": 0,
@@ -172,7 +174,7 @@
{
"name": "image",
"type": "IMAGE",
"link": 64
"link": 1
}
],
"outputs": [
@@ -180,31 +182,34 @@
"name": "output_image",
"type": "IMAGE",
"links": [
62
],
"slot_index": 0
5
]
}
],
"properties": {
"Node name for S&R": "ShotByTextNode"
"Node name for S&R": "ShotByTextOriginal"
},
"widgets_values": [
"a beautiful sunset",
1,
"BRIA_API_TOKEN",
"sea",
"fast",
false,
false,
true,
""
]
},
{
"id": 37,
"type": "ShotByImageNode",
"pos": {
"0": 999,
"1": 932
},
"size": {
"0": 315,
"1": 102
},
"id": 4,
"type": "ShotByImageOriginal",
"pos": [
945.0781860351562,
441.9688415527344
],
"size": [
272.0703125,
174
],
"flags": {},
"order": 4,
"mode": 0,
@@ -212,12 +217,12 @@
{
"name": "image",
"type": "IMAGE",
"link": 65
"link": 3
},
{
"name": "ref_image",
"type": "IMAGE",
"link": 59
"link": 4
}
],
"outputs": [
@@ -225,58 +230,60 @@
"name": "output_image",
"type": "IMAGE",
"links": [
63
],
"slot_index": 0
6
]
}
],
"properties": {
"Node name for S&R": "ShotByImageNode"
"Node name for S&R": "ShotByImageOriginal"
},
"widgets_values": [
0,
""
"BRIA_API_TOKEN",
false,
false,
true,
1
]
}
],
"links": [
[
59,
39,
1,
1,
0,
37,
3,
0,
"IMAGE"
],
[
3,
1,
0,
4,
0,
"IMAGE"
],
[
4,
2,
0,
4,
1,
"IMAGE"
],
[
62,
36,
5,
3,
0,
40,
6,
0,
"IMAGE"
],
[
63,
37,
6,
4,
0,
41,
0,
"IMAGE"
],
[
64,
42,
0,
36,
0,
"IMAGE"
],
[
65,
42,
0,
37,
7,
0,
"IMAGE"
]
@@ -285,12 +292,13 @@
"config": {},
"extra": {
"ds": {
"scale": 0.9849732675807669,
"scale": 0.7513148009015777,
"offset": [
-339.6686422794803,
-496.4354678014682
11.112206386364164,
66.47311795454547
]
}
},
"frontendVersion": "1.25.11"
},
"version": 0.4
}
}