diff --git a/__init__.py b/__init__.py index 41cf2b9..6f6c7d4 100644 --- a/__init__.py +++ b/__init__.py @@ -1,10 +1,12 @@ -from .nodes import EraserNode, GenFillNode, ShotByTextNode, ShotByImageNode +from .nodes import EraserNode, GenFillNode, ShotByTextNode, ShotByImageNode, TailoredGenNode, TailoredModelInfoNode # 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, + "BriaTailoredGen": TailoredGenNode, + "TailoredModelInfoNode": TailoredModelInfoNode, } # Map the node display name to the one shown in the ComfyUI node interface NODE_DISPLAY_NAME_MAPPINGS = { @@ -12,4 +14,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "BriaGenFill": "Bria GenFill", "ShotByTextNode": "Bria Shot By Text", "ShotByImageNode": "Bria Shot By Image", + "BriaTailoredGen": "Bria Tailored Gen", + "TailoredModelInfoNode": "Bria Tailored Model Info", } diff --git a/nodes/__init__.py b/nodes/__init__.py index e8046bf..6218cbe 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -2,3 +2,5 @@ from .eraser_node import EraserNode from .generative_fill_node import GenFillNode 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 \ No newline at end of file diff --git a/nodes/base_node.py b/nodes/base_node.py deleted file mode 100644 index 819b090..0000000 --- a/nodes/base_node.py +++ /dev/null @@ -1,96 +0,0 @@ -import numpy as np -import requests -from PIL import Image -import io -import base64 -from torchvision.transforms import ToPILImage, ToTensor -import torch - -# Base class for shared functionality between both nodes -class BriaAPINode: - def __init__(self, api_url): - self.api_url = api_url - - def preprocess_image(self, image): - if isinstance(image, torch.Tensor): - # Print image shape for debugging - if image.dim() == 4: # (batch_size, height, width, channels) - image = image.squeeze(0) # Remove the batch dimension (1) - # Convert to PIL after permuting to (height, width, channels) - image = ToPILImage()(image.permute(2, 0, 1)) # (height, width, channels) - else: - print("Unexpected image dimensions. Expected 4D tensor.") - return image - - def preprocess_mask(self, mask): - if isinstance(mask, torch.Tensor): - # Print mask shape for debugging - if mask.dim() == 3: # (batch_size, height, width) - mask = mask.squeeze(0) # Remove the batch dimension (1) - # Convert to PIL (grayscale mask) - mask = ToPILImage()(mask) # No permute needed for grayscale - else: - 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() - pil_image.save(buffered, format="PNG") # Save the image to the buffer in PNG format - buffered.seek(0) # Rewind the buffer to the beginning - return base64.b64encode(buffered.getvalue()).decode('utf-8') - - def process_request(self, image, mask, api_key): - 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(mask, torch.Tensor): - mask = self.preprocess_mask(mask) - - # Convert the image and mask directly to Base64 strings - image_base64 = self.image_to_base64(image) - mask_base64 = self.image_to_base64(mask) - - # Prepare the API request payload - payload = { - "file": f"{image_base64}", - "mask_file": f"{mask_base64}" - } - - 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_url']) - result_image = Image.open(io.BytesIO(image_response.content)) - result_image = result_image.convert("RGBA") - result_image = np.array(result_image).astype(np.float32) / 255.0 - result_image = torch.from_numpy(result_image)[None,] - # image_tensor = image_tensor = ToTensor()(output_image) - # image_tensor = image_tensor.permute(1, 2, 0) / 255.0 # Shape now becomes [1, 2200, 1548, 3] - # print(f"output tensor shape is: {image_tensor.shape}") - 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}") diff --git a/nodes/common.py b/nodes/common.py new file mode 100644 index 0000000..b57a492 --- /dev/null +++ b/nodes/common.py @@ -0,0 +1,92 @@ +import numpy as np +from PIL import Image +import io +import torch +import base64 +from torchvision.transforms import ToPILImage +import requests + +def postprocess_image(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(pil_image): + # Convert a PIL image to a base64-encoded string + buffered = io.BytesIO() + pil_image.save(buffered, format="PNG") # Save the image to the buffer in PNG format + buffered.seek(0) # Rewind the buffer to the beginning + return base64.b64encode(buffered.getvalue()).decode('utf-8') + +def preprocess_image(image): + if isinstance(image, torch.Tensor): + # Print image shape for debugging + if image.dim() == 4: # (batch_size, height, width, channels) + image = image.squeeze(0) # Remove the batch dimension (1) + # Convert to PIL after permuting to (height, width, channels) + image = ToPILImage()(image.permute(2, 0, 1)) # (height, width, channels) + else: + print("Unexpected image dimensions. Expected 4D tensor.") + return image + + +def preprocess_mask(mask): + if isinstance(mask, torch.Tensor): + # Print mask shape for debugging + if mask.dim() == 3: # (batch_size, height, width) + mask = mask.squeeze(0) # Remove the batch dimension (1) + # Convert to PIL (grayscale mask) + mask = ToPILImage()(mask) # No permute needed for grayscale + else: + print("Unexpected mask dimensions. Expected 3D tensor.") + return mask + + +def process_request(api_url, image, mask, api_key): + 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(mask, torch.Tensor): + mask = preprocess_mask(mask) + + # Convert the image and mask directly to Base64 strings + image_base64 = image_to_base64(image) + mask_base64 = image_to_base64(mask) + + # Prepare the API request payload + payload = { + "file": f"{image_base64}", + "mask_file": f"{mask_base64}" + } + + headers = { + "Content-Type": "application/json", + "api_token": f"{api_key}" + } + + try: + response = requests.post(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_url']) + result_image = Image.open(io.BytesIO(image_response.content)) + result_image = result_image.convert("RGBA") + result_image = np.array(result_image).astype(np.float32) / 255.0 + result_image = torch.from_numpy(result_image)[None,] + # image_tensor = image_tensor = ToTensor()(output_image) + # image_tensor = image_tensor.permute(1, 2, 0) / 255.0 # Shape now becomes [1, 2200, 1548, 3] + # print(f"output tensor shape is: {image_tensor.shape}") + 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}") diff --git a/nodes/eraser_node.py b/nodes/eraser_node.py index 8f81afd..b6a3c69 100644 --- a/nodes/eraser_node.py +++ b/nodes/eraser_node.py @@ -1,17 +1,8 @@ -import numpy as np -import requests -from PIL import Image -import io -import base64 -from torchvision.transforms import ToPILImage, ToTensor -import torch +from .common import process_request -from .base_node import BriaAPINode - -# Eraser Node -class EraserNode(BriaAPINode): - @staticmethod - def INPUT_TYPES(): +class EraserNode(): + @classmethod + def INPUT_TYPES(self): return { "required": { "image": ("IMAGE",), # Input image from another node @@ -26,9 +17,9 @@ class EraserNode(BriaAPINode): FUNCTION = "execute" # This is the method that will be executed def __init__(self): - super().__init__("https://engine.prod.bria-api.com/v1/eraser") # Eraser API URL + self.api_url = "https://engine.prod.bria-api.com/v1/eraser" # Eraser API URL # Define the execute method as expected by ComfyUI def execute(self, image, mask, api_key): - return self.process_request(image, mask, api_key) + return process_request(self.api_url, image, mask, api_key) \ No newline at end of file diff --git a/nodes/generative_fill_node.py b/nodes/generative_fill_node.py index fb321e6..b92221c 100644 --- a/nodes/generative_fill_node.py +++ b/nodes/generative_fill_node.py @@ -2,17 +2,14 @@ import numpy as np import requests from PIL import Image import io -import base64 -from torchvision.transforms import ToPILImage, ToTensor import torch -from .base_node import BriaAPINode +from .common import image_to_base64, preprocess_image, preprocess_mask -# Generative Fill Node -class GenFillNode(BriaAPINode): - @staticmethod - def INPUT_TYPES(): +class GenFillNode(): + @classmethod + def INPUT_TYPES(self): return { "required": { "image": ("IMAGE",), # Input image from another node @@ -28,7 +25,7 @@ class GenFillNode(BriaAPINode): FUNCTION = "execute" # This is the method that will be executed def __init__(self): - super().__init__("https://engine.prod.bria-api.com/v1/gen_fill") # Eraser API URL + self.api_url = "https://engine.prod.bria-api.com/v1/gen_fill" # Eraser API URL # Define the execute method as expected by ComfyUI def execute(self, image, mask, prompt, api_key): @@ -37,13 +34,13 @@ class GenFillNode(BriaAPINode): # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): - image = self.preprocess_image(image) + image = preprocess_image(image) if isinstance(mask, torch.Tensor): - mask = self.preprocess_mask(mask) + mask = preprocess_mask(mask) # Convert the image and mask directly to Base64 strings - image_base64 = self.image_to_base64(image) - mask_base64 = self.image_to_base64(mask) + image_base64 = image_to_base64(image) + mask_base64 = image_to_base64(mask) # Prepare the API request payload payload = { diff --git a/nodes/shot_by_image_node.py b/nodes/shot_by_image_node.py index e12515b..6ee1647 100644 --- a/nodes/shot_by_image_node.py +++ b/nodes/shot_by_image_node.py @@ -1,17 +1,11 @@ -import numpy as np import requests -from PIL import Image -import io -import base64 -from torchvision.transforms import ToPILImage, ToTensor import torch -from .base_node import BriaAPINode +from .common import postprocess_image, preprocess_image, image_to_base64 -# shot by image Node -class ShotByImageNode(BriaAPINode): - @staticmethod - def INPUT_TYPES(): +class ShotByImageNode(): + @classmethod + def INPUT_TYPES(self): return { "required": { "image": ("IMAGE",), # Input image from another node @@ -36,13 +30,13 @@ class ShotByImageNode(BriaAPINode): # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): - image = self.preprocess_image(image) + image = preprocess_image(image) if isinstance(ref_image, torch.Tensor): - ref_image = self.preprocess_image(ref_image) + ref_image = 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) + image_base64 = image_to_base64(image) + ref_image_base64 = image_to_base64(ref_image) enhance_ref_image = bool(enhance_ref_image) payload = { @@ -65,7 +59,7 @@ class ShotByImageNode(BriaAPINode): # 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) + result_image = postprocess_image(image_response.content) return (result_image,) else: raise Exception(f"Error: API request failed with status code {response.status_code}") diff --git a/nodes/shot_by_text_node.py b/nodes/shot_by_text_node.py index 4f30ec3..16516a3 100644 --- a/nodes/shot_by_text_node.py +++ b/nodes/shot_by_text_node.py @@ -1,17 +1,11 @@ -import numpy as np import requests -from PIL import Image -import io -import base64 -from torchvision.transforms import ToPILImage, ToTensor import torch -from .base_node import BriaAPINode +from .common import postprocess_image, preprocess_image, image_to_base64 -# shot by text Node -class ShotByTextNode(BriaAPINode): - @staticmethod - def INPUT_TYPES(): +class ShotByTextNode(): + @classmethod + def INPUT_TYPES(self): return { "required": { "image": ("IMAGE",), # Input image from another node @@ -27,8 +21,8 @@ class ShotByTextNode(BriaAPINode): 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 - + 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, optimize_description, ): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": @@ -36,10 +30,10 @@ class ShotByTextNode(BriaAPINode): # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): - image = self.preprocess_image(image) + image = preprocess_image(image) optimize_description = bool(optimize_description) - image_base64 = self.image_to_base64(image) + image_base64 = image_to_base64(image) payload = { "file": image_base64, "scene_description": scene_description, @@ -60,7 +54,7 @@ class ShotByTextNode(BriaAPINode): # 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) + result_image = postprocess_image(image_response.content) return (result_image,) else: raise Exception(f"Error: API request failed with status code {response.status_code}") diff --git a/nodes/tailored_gen_node.py b/nodes/tailored_gen_node.py new file mode 100644 index 0000000..c18af6b --- /dev/null +++ b/nodes/tailored_gen_node.py @@ -0,0 +1,84 @@ +import requests + +from .common import postprocess_image, preprocess_image, image_to_base64 + + +class TailoredGenNode(): + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "model_id": ("STRING",), + "api_key": ("STRING", ), + }, + "optional": { + "prompt": ("STRING",), + "generation_prefix": ("STRING",), # possibly get this from the tailored model info node + "aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], {"default": "4:3"}), + "seed": ("INT", {"default": -1}), + "model_influence": ("FLOAT", {"default": 1.0}), + "include_generation_prefix": ("INT", {"default": 0}), + "negative_prompt": ("STRING", {"default": ""}), + "fast": ("INT", {"default": 1}), # possibly get this from the tailored model info node + "steps_num": ("INT", {"default": 8}), # possibly get this from the tailored model info node + "guidance_method_1": (["controlnet_canny", "controlnet_depth", "controlnet_recoloring", "controlnet_color_grid"],), + "guidance_method_1_scale": ("FLOAT", {"default": 1.0}), + "guidance_method_1_image": ("IMAGE", ), + "guidance_method_2": (["controlnet_canny", "controlnet_depth", "controlnet_recoloring", "controlnet_color_grid"],), + "guidance_method_2_scale": ("FLOAT", {"default": 1.0}), + "guidance_method_2_image": ("IMAGE", ), + } + } + + 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/text-to-image/tailored/" #"http://0.0.0.0:5000/v1/text-to-image/tailored/" + + def execute( + self, model_id, api_key, prompt, generation_prefix, aspect_ratio, + seed, model_influence, include_generation_prefix, negative_prompt, fast, steps_num, + guidance_method_1=None, guidance_method_1_scale=None, guidance_method_1_image=None, + guidance_method_2=None, guidance_method_2_scale=None, guidance_method_2_image=None, + ): + include_generation_prefix = bool(include_generation_prefix) + fast = bool(fast) + payload = { + "prompt": generation_prefix + prompt, + "num_results": 1, + "aspect_ratio": aspect_ratio, + "sync": True, + "seed": seed, + "model_influence": model_influence, + "include_generation_prefix": include_generation_prefix, + "negative_prompt": negative_prompt, + "fast": fast, + "steps_num": steps_num, + } + if guidance_method_1_image is not None: + guidance_method_1_image = preprocess_image(guidance_method_1_image) + guidance_method_1_image = image_to_base64(guidance_method_1_image) + payload["guidance_method_1"] = guidance_method_1 + payload["guidance_method_1_scale"] = guidance_method_1_scale + payload["guidance_method_1_image_file"] = guidance_method_1_image + if guidance_method_2_image is not None: + guidance_method_2_image = preprocess_image(guidance_method_2_image) + guidance_method_2_image = image_to_base64(guidance_method_2_image) + payload["guidance_method_2"] = guidance_method_2 + payload["guidance_method_2_scale"] = guidance_method_2_scale + payload["guidance_method_2_image_file"] = guidance_method_2_image + response = requests.post( + self.api_url + model_id, + json=payload, + headers={"api_token": api_key} + ) + if response.status_code == 200: + response_dict = response.json() + image_response = requests.get(response_dict['result'][0]["urls"][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} and text {response.text}") diff --git a/nodes/tailored_model_info_node.py b/nodes/tailored_model_info_node.py new file mode 100644 index 0000000..aec2b24 --- /dev/null +++ b/nodes/tailored_model_info_node.py @@ -0,0 +1,35 @@ +import requests + + +class TailoredModelInfoNode(): + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "model_id": ("STRING",), + "api_key": ("STRING", ) + } + } + + RETURN_TYPES = ("STRING", "INT", "INT", ) + RETURN_NAMES = ("generation_prefix", "default_fast", "default_steps_num", ) + 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/tailored-gen/models/" + + # Define the execute method as expected by ComfyUI + def execute(self, model_id, api_key): + response = requests.get( + self.api_url + model_id, + headers={"api_token": api_key} + ) + if response.status_code == 200: + generation_prefix = response.json()["generation_prefix"] + training_version = response.json()["training_version"] + default_fast = 1 if training_version == "light" else 0 + default_steps_num = 8 if training_version == "light" else 30 + return (generation_prefix, default_fast, default_steps_num,) + else: + raise Exception(f"Error: API request failed with status code {response.status_code} and text {response.text}")