diff --git a/__init__.py b/__init__.py index 26a7828..e83ab09 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/nodes/__init__.py b/nodes/__init__.py index 2002920..dbaed9c 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -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 \ No newline at end of file +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 diff --git a/nodes/shot_by_image_automatic_aspect_ratio_node.py b/nodes/shot_by_image_automatic_aspect_ratio_node.py new file mode 100644 index 0000000..a32d075 --- /dev/null +++ b/nodes/shot_by_image_automatic_aspect_ratio_node.py @@ -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) diff --git a/nodes/shot_by_image_automatic_node.py b/nodes/shot_by_image_automatic_node.py new file mode 100644 index 0000000..b5cc3ea --- /dev/null +++ b/nodes/shot_by_image_automatic_node.py @@ -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) diff --git a/nodes/shot_by_image_custom_coordinates_node.py b/nodes/shot_by_image_custom_coordinates_node.py new file mode 100644 index 0000000..9e1b6e5 --- /dev/null +++ b/nodes/shot_by_image_custom_coordinates_node.py @@ -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) diff --git a/nodes/shot_by_image_manual_padding_node.py b/nodes/shot_by_image_manual_padding_node.py new file mode 100644 index 0000000..246236e --- /dev/null +++ b/nodes/shot_by_image_manual_padding_node.py @@ -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) diff --git a/nodes/shot_by_image_manual_placement_node.py b/nodes/shot_by_image_manual_placement_node.py new file mode 100644 index 0000000..963ed61 --- /dev/null +++ b/nodes/shot_by_image_manual_placement_node.py @@ -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) diff --git a/nodes/shot_by_image_node.py b/nodes/shot_by_image_node.py index 29dce19..04168b5 100644 --- a/nodes/shot_by_image_node.py +++ b/nodes/shot_by_image_node.py @@ -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) diff --git a/nodes/shot_by_text_automatic_aspect_ratio_node.py b/nodes/shot_by_text_automatic_aspect_ratio_node.py new file mode 100644 index 0000000..9658a70 --- /dev/null +++ b/nodes/shot_by_text_automatic_aspect_ratio_node.py @@ -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) diff --git a/nodes/shot_by_text_automatic_node.py b/nodes/shot_by_text_automatic_node.py new file mode 100644 index 0000000..d184847 --- /dev/null +++ b/nodes/shot_by_text_automatic_node.py @@ -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) diff --git a/nodes/shot_by_text_custom_coordinates_node.py b/nodes/shot_by_text_custom_coordinates_node.py new file mode 100644 index 0000000..a7f5fc4 --- /dev/null +++ b/nodes/shot_by_text_custom_coordinates_node.py @@ -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) diff --git a/nodes/shot_by_text_manual_padding_node.py b/nodes/shot_by_text_manual_padding_node.py new file mode 100644 index 0000000..5700e1f --- /dev/null +++ b/nodes/shot_by_text_manual_padding_node.py @@ -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) diff --git a/nodes/shot_by_text_manual_placement_node.py b/nodes/shot_by_text_manual_placement_node.py new file mode 100644 index 0000000..b35dd85 --- /dev/null +++ b/nodes/shot_by_text_manual_placement_node.py @@ -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) diff --git a/nodes/shot_by_text_node.py b/nodes/shot_by_text_node.py index 0998f1f..8020775 100644 --- a/nodes/shot_by_text_node.py +++ b/nodes/shot_by_text_node.py @@ -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) diff --git a/nodes/utils/shot_utils.py b/nodes/utils/shot_utils.py new file mode 100644 index 0000000..5a5287b --- /dev/null +++ b/nodes/utils/shot_utils.py @@ -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