Merge pull request #3 from Bria-AI/t2i-comfy

tailored and some code cleaning
This commit is contained in:
BriaOr
2025-01-08 16:14:11 +02:00
committed by GitHub
10 changed files with 251 additions and 154 deletions
+5 -1
View File
@@ -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",
}
+2
View File
@@ -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
-96
View File
@@ -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}")
+92
View File
@@ -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}")
+6 -15
View File
@@ -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)
+9 -12
View File
@@ -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 = {
+9 -15
View File
@@ -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}")
+9 -15
View File
@@ -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}")
+84
View File
@@ -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}")
+35
View File
@@ -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}")