WAI-4049
This commit is contained in:
+52
-7
@@ -1,6 +1,31 @@
|
||||
from .nodes import (EraserNode, GenFillNode, ImageExpansionNode, ReplaceBgNode, RmbgNode, RemoveForegroundNode, ShotByTextNode, ShotByImageNode, TailoredGenNode,
|
||||
TailoredModelInfoNode, Text2ImageBaseNode, Text2ImageFastNode, Text2ImageHDNode, TailoredPortraitNode,
|
||||
ReimagineNode)
|
||||
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,
|
||||
)
|
||||
|
||||
# Map the node class to a name used internally by ComfyUI
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"BriaEraser": EraserNode, # Return the class, not an instance
|
||||
@@ -9,8 +34,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,
|
||||
@@ -27,8 +62,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
@@ -4,12 +4,24 @@ 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
|
||||
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 .reimagine_node import ReimagineNode
|
||||
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
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
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,
|
||||
"automatic_aspect_ratio",
|
||||
aspect_ratio=aspect_ratio,
|
||||
sku=sku,
|
||||
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,45 @@
|
||||
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url
|
||||
|
||||
|
||||
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",)
|
||||
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,
|
||||
api_key,
|
||||
sku="",
|
||||
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,
|
||||
"automatic",
|
||||
shot_size=shot_size,
|
||||
sku=sku,
|
||||
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,57 @@
|
||||
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
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,
|
||||
"custom_coordinates",
|
||||
shot_size=shot_size,
|
||||
foreground_image_size=foreground_image_size,
|
||||
foreground_image_location=foreground_image_location,
|
||||
sku=sku,
|
||||
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,46 @@
|
||||
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
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,
|
||||
"manual_padding",
|
||||
padding_values=padding_values,
|
||||
sku=sku,
|
||||
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,61 @@
|
||||
from .utils.shot_utils import get_image_input_types, create_image_payload, make_api_request, shot_by_image_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
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,
|
||||
"manual_placement",
|
||||
shot_size=shot_size,
|
||||
manual_placement_selection=manual_placement_selection,
|
||||
sku=sku,
|
||||
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)
|
||||
+43
-72
@@ -1,72 +1,43 @@
|
||||
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
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
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,
|
||||
"original",
|
||||
original_quality=True,
|
||||
sku=sku,
|
||||
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,50 @@
|
||||
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
sync=False,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"automatic_aspect_ratio",
|
||||
aspect_ratio=aspect_ratio,
|
||||
sku=sku,
|
||||
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,47 @@
|
||||
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url
|
||||
|
||||
|
||||
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",)
|
||||
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,
|
||||
api_key,
|
||||
sku="",
|
||||
sync=False,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"automatic",
|
||||
shot_size=shot_size,
|
||||
sku=sku,
|
||||
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,59 @@
|
||||
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
sync=False,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"custom_coordinates",
|
||||
shot_size=shot_size,
|
||||
foreground_image_size=foreground_image_size,
|
||||
foreground_image_location=foreground_image_location,
|
||||
sku=sku,
|
||||
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,47 @@
|
||||
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
sync=False,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"manual_padding",
|
||||
padding_values=padding_values,
|
||||
sku=sku,
|
||||
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,64 @@
|
||||
from .utils.shot_utils import get_text_input_types, create_text_payload, make_api_request, shot_by_text_api_url
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
sync=False,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"manual_placement",
|
||||
shot_size=shot_size,
|
||||
manual_placement_selection=manual_placement_selection,
|
||||
sku=sku,
|
||||
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)
|
||||
+44
-68
@@ -1,68 +1,44 @@
|
||||
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
|
||||
|
||||
|
||||
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,
|
||||
sku="",
|
||||
sync=True,
|
||||
optimize_description=True,
|
||||
exclude_elements="",
|
||||
force_rmbg=False,
|
||||
content_moderation=False,
|
||||
):
|
||||
payload = create_text_payload(
|
||||
image,
|
||||
api_key,
|
||||
scene_description,
|
||||
mode,
|
||||
"original",
|
||||
original_quality=True,
|
||||
sku=sku,
|
||||
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,183 @@
|
||||
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"
|
||||
)
|
||||
|
||||
|
||||
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 == "automatic":
|
||||
payload["shot_size"] = [
|
||||
int(x.strip()) for x in kwargs.get("shot_size").split(",")
|
||||
]
|
||||
elif placement_type == "manual_placement":
|
||||
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 == "custom_coordinates":
|
||||
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 == "manual_padding":
|
||||
payload["padding_values"] = [
|
||||
int(x.strip()) for x in kwargs.get("padding_values").split(",")
|
||||
]
|
||||
|
||||
elif placement_type == "automatic_aspect_ratio":
|
||||
payload["aspect_ratio"] = kwargs.get("aspect_ratio", "1:1")
|
||||
elif placement_type == "original":
|
||||
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("sku", "").strip():
|
||||
payload["sku"] = kwargs["sku"]
|
||||
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),
|
||||
}
|
||||
|
||||
if kwargs.get("sku", "").strip():
|
||||
payload["sku"] = kwargs["sku"]
|
||||
payload = update_payload_for_placement(placement_type, payload, **kwargs)
|
||||
|
||||
return payload
|
||||
|
||||
|
||||
def make_api_request(api_url, payload, api_key):
|
||||
"""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()
|
||||
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": {
|
||||
"sku": ("STRING", {"default": ""}),
|
||||
"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
|
||||
Reference in New Issue
Block a user