Add ShotByTextNode and ShotByImageNode classes with API integration

This commit is contained in:
ori-liberman
2024-12-18 15:11:57 +00:00
parent 99155fb898
commit 0b935dcffc
2 changed files with 137 additions and 1 deletions
+132
View File
@@ -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