diff --git a/__init__.py b/__init__.py index 48ade2d..e10459f 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/nodes/__init__.py b/nodes/__init__.py index 3440ffc..f99d7ce 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -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 \ No newline at end of file +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 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..b786f5f --- /dev/null +++ b/nodes/shot_by_image_automatic_aspect_ratio_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, 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) diff --git a/nodes/shot_by_image_automatic_node.py b/nodes/shot_by_image_automatic_node.py new file mode 100644 index 0000000..4c405a5 --- /dev/null +++ b/nodes/shot_by_image_automatic_node.py @@ -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) 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..7743973 --- /dev/null +++ b/nodes/shot_by_image_custom_coordinates_node.py @@ -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) 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..9cfee3e --- /dev/null +++ b/nodes/shot_by_image_manual_padding_node.py @@ -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) 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..6c7709f --- /dev/null +++ b/nodes/shot_by_image_manual_placement_node.py @@ -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) diff --git a/nodes/shot_by_image_node.py b/nodes/shot_by_image_node.py index 29dce19..1a0007a 100644 --- a/nodes/shot_by_image_node.py +++ b/nodes/shot_by_image_node.py @@ -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) 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..4679e71 --- /dev/null +++ b/nodes/shot_by_text_automatic_aspect_ratio_node.py @@ -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) diff --git a/nodes/shot_by_text_automatic_node.py b/nodes/shot_by_text_automatic_node.py new file mode 100644 index 0000000..20371f0 --- /dev/null +++ b/nodes/shot_by_text_automatic_node.py @@ -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) 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..f0d0725 --- /dev/null +++ b/nodes/shot_by_text_custom_coordinates_node.py @@ -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) 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..9bc9d70 --- /dev/null +++ b/nodes/shot_by_text_manual_padding_node.py @@ -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) 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..5b28048 --- /dev/null +++ b/nodes/shot_by_text_manual_placement_node.py @@ -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) diff --git a/nodes/shot_by_text_node.py b/nodes/shot_by_text_node.py index 0998f1f..e54bdce 100644 --- a/nodes/shot_by_text_node.py +++ b/nodes/shot_by_text_node.py @@ -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) diff --git a/nodes/utils/shot_utils.py b/nodes/utils/shot_utils.py new file mode 100644 index 0000000..bcdda1a --- /dev/null +++ b/nodes/utils/shot_utils.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index cf86390..20d4b5e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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] diff --git a/workflows/product shot generation_workflow.json b/workflows/product_original_ shot_generation_workflow.json similarity index 60% rename from workflows/product shot generation_workflow.json rename to workflows/product_original_ shot_generation_workflow.json index 41eba15..7478c80 100644 --- a/workflows/product shot generation_workflow.json +++ b/workflows/product_original_ shot_generation_workflow.json @@ -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 -} +} \ No newline at end of file