From 0b935dcffcb9f899e9999898c5da19e3ec57d332 Mon Sep 17 00:00:00 2001 From: ori-liberman Date: Wed, 18 Dec 2024 15:11:57 +0000 Subject: [PATCH] Add ShotByTextNode and ShotByImageNode classes with API integration --- __init__.py | 6 ++- bria_api_node.py | 132 +++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 137 insertions(+), 1 deletion(-) diff --git a/__init__.py b/__init__.py index c752fac..16febbd 100644 --- a/__init__.py +++ b/__init__.py @@ -1,11 +1,15 @@ -from .bria_api_node import EraserNode, GenFillNode +from .bria_api_node import EraserNode, GenFillNode, ShotByTextNode, ShotByImageNode # Map the node class to a name used internally by ComfyUI NODE_CLASS_MAPPINGS = { "BriaEraser": EraserNode, # Return the class, not an instance "BriaGenFill": GenFillNode, + "ShotByTextNode": ShotByTextNode, + "ShotByImageNode": ShotByImageNode, } # Map the node display name to the one shown in the ComfyUI node interface NODE_DISPLAY_NAME_MAPPINGS = { "BriaEraser": "Bria Eraser", "BriaGenFill": "Bria GenFill", + "ShotByTextNode": "Bria Shot By Text", + "ShotByImageNode": "Bria Shot By Image", } diff --git a/bria_api_node.py b/bria_api_node.py index 269e089..db1a089 100644 --- a/bria_api_node.py +++ b/bria_api_node.py @@ -33,6 +33,14 @@ class BriaAPINode: print("Unexpected mask dimensions. Expected 3D tensor.") return mask + def postprocess_image(self, image): + result_image = Image.open(io.BytesIO(image)) + result_image = result_image.convert("RGB") + result_image = np.array(result_image).astype(np.float32) / 255.0 + result_image = torch.from_numpy(result_image)[None,] + return result_image + + def image_to_base64(self, pil_image): # Convert a PIL image to a base64-encoded string buffered = io.BytesIO() @@ -110,6 +118,130 @@ class EraserNode(BriaAPINode): # Define the execute method as expected by ComfyUI def execute(self, image, mask, api_key): return self.process_request(image, mask, api_key) + +# shot by text Node +class ShotByTextNode(BriaAPINode): + @staticmethod + def INPUT_TYPES(): + return { + "required": { + "image": ("IMAGE",), # Input image from another node + "scene_description": ("STRING",), + "optimize_description": ("BOOLEAN", {"default": "True"}), + "api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("output_image",) + CATEGORY = "API Nodes" + FUNCTION = "execute" # This is the method that will be executed + + def __init__(self): + super().__init__("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, optimize_description, ): + 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 = self.preprocess_image(image) + + image_base64 = self.image_to_base64(image) + payload = { + "file": image_base64, + "scene_description": scene_description, + "optimize_description": optimize_description, + "placement_type": "original", + "original_quality": True, + "sync": True + } + 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 = self.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}") + + +# shot by text Node +class ShotByImageNode(BriaAPINode): + @staticmethod + def INPUT_TYPES(): + return { + "required": { + "image": ("IMAGE",), # Input image from another node + "ref_image": ("IMAGE",), # ref image from another node + "enhance_ref_image": ("BOOLEAN", {"default": "True"}), + "api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("output_image",) + CATEGORY = "API Nodes" + FUNCTION = "execute" # This is the method that will be executed + + def __init__(self): + super().__init__("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, ): + 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 = self.preprocess_image(image) + if isinstance(ref_image, torch.Tensor): + ref_image = self.preprocess_image(ref_image) + + # Convert the image and mask directly to Base64 strings + image_base64 = self.image_to_base64(image) + ref_image_base64 = self.image_to_base64(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 + } + 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 = self.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}") # Generative Fill Node