From 07827ef34ff30d4bf5187cfe82b62cd38493c019 Mon Sep 17 00:00:00 2001 From: xenia-kra <81739495+xenia-kra@users.noreply.github.com> Date: Mon, 10 Mar 2025 12:45:44 +0400 Subject: [PATCH] tailored portrait (#18) --- __init__.py | 2 + nodes/__init__.py | 1 + nodes/tailored_portrait_node.py | 72 +++++++++++++++++++++++++++++++++ 3 files changed, 75 insertions(+) create mode 100644 nodes/tailored_portrait_node.py diff --git a/__init__.py b/__init__.py index 77c0b2e..d578aa7 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,4 @@ +from nodes.tailored_portrait_node import TailoredPortraitNode from .nodes import (EraserNode, GenFillNode, ImageExpansionNode, ReplaceBgNode, RmbgNode, RemoveForegroundNode, ShotByTextNode, ShotByImageNode, TailoredGenNode, TailoredModelInfoNode, Text2ImageBaseNode, Text2ImageFastNode, Text2ImageHDNode, ReimagineNode) @@ -13,6 +14,7 @@ NODE_CLASS_MAPPINGS = { "ShotByImageNode": ShotByImageNode, "BriaTailoredGen": TailoredGenNode, "TailoredModelInfoNode": TailoredModelInfoNode, + "TailoredPortraitNode": TailoredPortraitNode, "Text2ImageBaseNode": Text2ImageBaseNode, "Text2ImageFastNode": Text2ImageFastNode, "Text2ImageHDNode": Text2ImageHDNode, diff --git a/nodes/__init__.py b/nodes/__init__.py index b4621a1..2002920 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -8,6 +8,7 @@ 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 diff --git a/nodes/tailored_portrait_node.py b/nodes/tailored_portrait_node.py new file mode 100644 index 0000000..bb5eb6a --- /dev/null +++ b/nodes/tailored_portrait_node.py @@ -0,0 +1,72 @@ +import numpy as np +import requests +from PIL import Image +import io +import torch + +from .common import image_to_base64 + +class TailoredPortraitNode(): + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "image": ("IMAGE",), # Input image from another node + "tailored_model_id": ("INT",), + "api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value + }, + "optional": { + "seed": ("INT", {"default": 123456}), + "tailored_model_influence": ("FLOAT", {"default": 0.9}), + "id_strength": ("FLOAT", {"default": 0.7}), + } + } + + 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/tailored-gen/restyle_portrait" # Eraser API URL + + # Define the execute method as expected by ComfyUI + def execute(self, image, tailored_model_id, api_key, seed, tailored_model_influence, id_strength): + if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": + raise Exception("Please insert a valid API key.") + + # Convert the image and mask directly to Base64 strings + image_base64 = image_to_base64(image) + + # Prepare the API request payload + payload = { + "id_image_file": f"{image_base64}", + "tailored_model_id": tailored_model_id, + "tailored_model_influence": tailored_model_influence, + "id_strength": id_strength, + "seed": seed + } + + 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['image_res']) + result_image = Image.open(io.BytesIO(image_response.content)) + 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,) + else: + raise Exception(f"Error: API request failed with status code {response.status_code}") + + except Exception as e: + raise Exception(f"{e}")