Compare commits

..
Author SHA1 Message Date
Your Name fe10d6990e fix image preprocess 2025-03-10 17:27:35 +04:00
Your Name d5b7d2550c fix import 2025-03-10 15:40:25 +04:00
Xenia 80db13454c Merge remote-tracking branch 'origin/main' into comfy_tailored_portrait 2025-03-10 08:53:57 +00:00
Xenia ae89bb1aa0 tailored portrait 2025-03-10 08:52:49 +00:00
xenia-kra 07827ef34f tailored portrait (#18) 2025-03-10 12:45:44 +04:00
Xenia da131c5a49 tailored portrait 2025-03-10 08:26:15 +00:00
14 changed files with 133 additions and 305 deletions
+3 -5
View File
@@ -7,19 +7,17 @@ on:
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Bria-AI' }}
# if this is a forked repository. Skipping the workflow.
if: github.event.repository.fork == false
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+1 -2
View File
@@ -1,2 +1 @@
*.pyc
.idea
*.pyc
+2 -3
View File
@@ -13,7 +13,7 @@ An API token is required to use the nodes in your workflows. Get started quickly
<img src="https://img.shields.io/badge/GET%20YOUR%20TOKEN-1000%20Free%20Calls-blue?style=flat-square" alt="Get Your Token" height="20">
</a>.
for direct API endpoint use, you can find our APIs through partners like [**fal.ai**](https://fal.ai/models?keywords=bria).
For direct API Endpoint use, look for the endpoint in our of our API partners like: [**fal.ai**](https://fal.ai/models?keywords=bria).
For source code and weigths access, go to our [**Hugging Face**](https://huggingface.co/briaai) space.
To load a workflow, import the compatible workflow.json files from this [folder](workflows).
@@ -45,8 +45,7 @@ These nodes use pre-trained tailored models to generate images that faithfully r
| Node | Description |
|------------------------|--------------------------------------------------------------------|
| **Tailored Gen** | Generates images using a trained tailored model, reproducing specific visual IP elements or guidelines. Use the Tailored Model Info node to load the model's default settings. |
| **Tailored Model Info**| Retrieves the default settings and prompt prefix of a trained tailored model, which can be used to configure the Tailored Gen node. |
| **Restyle Portrait** | Transforms the style of a portrait while preserving the person's facial features. |
| **Tailored Model Info** | Retrieves the default settings and prompt prefix of a trained tailored model, which can be used to configure the Tailored Gen node. |
## Image Editing Nodes
These nodes modify specific parts of images, enabling adjustments while maintaining the integrity of the rest of the image.
-1
View File
@@ -31,7 +31,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ShotByImageNode": "Bria Shot By Image",
"BriaTailoredGen": "Bria Tailored Gen",
"TailoredModelInfoNode": "Bria Tailored Model Info",
"TailoredPortraitNode": "Bria Restyle Portrait",
"Text2ImageBaseNode": "Bria Text2Image Base",
"Text2ImageFastNode": "Bria Text2Image Fast",
"Text2ImageHDNode": "Bria Text2Image HD",
+12 -67
View File
@@ -5,7 +5,6 @@ import torch
import base64
from torchvision.transforms import ToPILImage
import requests
import time
def postprocess_image(image):
result_image = Image.open(io.BytesIO(image))
@@ -45,7 +44,7 @@ def preprocess_mask(mask):
return mask
def process_request(api_url, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation):
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.")
@@ -59,12 +58,10 @@ def process_request(api_url, image, mask, api_key, visual_input_content_moderati
image_base64 = image_to_base64(image)
mask_base64 = image_to_base64(mask)
# Prepare the API request payload for v2 API
# Prepare the API request payload
payload = {
"image": image_base64,
"mask": mask_base64,
"visual_input_content_moderation":visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation
"file": f"{image_base64}",
"mask_file": f"{mask_base64}"
}
headers = {
@@ -73,23 +70,13 @@ def process_request(api_url, image, mask, api_key, visual_input_content_moderati
}
try:
response = requests.post(api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial request successful, polling for completion...')
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()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
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
@@ -97,51 +84,9 @@ def process_request(api_url, image, mask, api_key, visual_input_content_moderati
# 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,)
return (result_image,)
else:
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
def poll_status_until_completed(status_url, api_key, timeout=360, check_interval=2):
"""
Poll a status URL until the status is COMPLETED or timeout is reached.
Args:
status_url (str): The status URL to poll
api_key (str): API token for authentication
timeout (int): Maximum time to wait in seconds (default: 360)
check_interval (int): Time between checks in seconds (default: 2)
Returns:
dict: The final response containing the result
Raises:
Exception: If timeout is reached or API request fails
"""
start_time = time.time()
headers = {"api_token": api_key}
while time.time() - start_time < timeout:
try:
response = requests.get(status_url, headers=headers)
if response.status_code == 200 or response.status_code == 202:
response_dict = response.json()
status = response_dict.get("status", "").upper()
if status == "COMPLETED":
return response_dict
elif status == "ERROR":
raise Exception(f"Request failed: {response_dict}")
else:
print(f"Status: {status}, waiting...")
time.sleep(check_interval)
else:
raise Exception(f"Status check failed with status code {response.status_code}")
except requests.exceptions.RequestException as e:
raise Exception(f"Error checking status: {e}")
raise Exception(f"Timeout reached after {timeout} seconds")
+3 -7
View File
@@ -8,10 +8,6 @@ class EraserNode():
"image": ("IMAGE",), # Input image from another node
"mask": ("MASK",), # Binary mask input
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value
},
"optional": {
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -21,9 +17,9 @@ class EraserNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/erase" # 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, visual_input_content_moderation, visual_output_content_moderation):
return process_request(self.api_url, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation)
def execute(self, image, mask, api_key):
return process_request(self.api_url, image, mask, api_key)
+13 -30
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed
from .common import image_to_base64, preprocess_image, preprocess_mask
class GenFillNode():
@@ -18,12 +18,7 @@ class GenFillNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"seed": ("INT", {"default": 123456}),
"prompt_content_moderation": ("BOOLEAN", {"default": True}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"seed": ("INT", {"default": 123456})
}
}
@@ -33,10 +28,10 @@ class GenFillNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/gen_fill"
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, seed, prompt_content_moderation, visual_input_content_moderation, visual_output_content_moderation):
def execute(self, image, mask, prompt, api_key, seed):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -52,14 +47,12 @@ class GenFillNode():
# Prepare the API request payload
payload = {
"image": image_base64,
"mask": mask_base64,
"file": f"{image_base64}",
"mask_file": f"{mask_base64}",
"prompt": prompt,
"negative_prompt": "blurry",
"sync": True,
"seed": seed,
"prompt_content_moderation":prompt_content_moderation,
"visual_input_content_moderation":visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation
}
headers = {
@@ -68,30 +61,20 @@ class GenFillNode():
}
try:
# Send initial request to get status URL
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial genfill request successful, polling for completion...')
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
result_image_url = final_response['result']['image_url']
image_response = requests.get(result_image_url)
image_response = requests.get(response_dict['urls'][0])
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} {response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
+26 -60
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import image_to_base64, preprocess_image, poll_status_until_completed
from .common import image_to_base64, preprocess_image
class ImageExpansionNode():
@@ -13,21 +13,17 @@ class ImageExpansionNode():
return {
"required": {
"image": ("IMAGE",), # Input image from another node
"original_image_size": ("STRING",),
"original_image_location": ("STRING",),
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"original_image_size": ("STRING",),
"original_image_location": ("STRING",),
"canvas_size": ("STRING", {"default": "1000, 1000"}),
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9","None"], {"default": "None"}),
"prompt": ("STRING", {"default": ""}),
"seed": ("INT", {"default": 681794}),
"negative_prompt": ("STRING", {"default": "Ugly, mutated"}),
"prompt_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -37,28 +33,27 @@ class ImageExpansionNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/expand" # Image Expansion API URL
self.api_url = "https://engine.prod.bria-api.com/v1/image_expansion" # Image Expansion API URL
# Define the execute method as expected by ComfyUI
def execute(self, image,
original_image_size,
original_image_location,
canvas_size,
aspect_ratio,
prompt,
seed,
negative_prompt,
prompt_content_moderation,
preserve_alpha,
visual_input_content_moderation,
visual_output_content_moderation,
content_moderation,
api_key):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
original_image_size = [int(x.strip()) for x in original_image_size.split(",")] if original_image_size else ()
original_image_location = [int(x.strip()) for x in original_image_location.split(",")] if original_image_location else ()
canvas_size = [int(x.strip()) for x in canvas_size.split(",")] if canvas_size else ()
original_image_size = [int(x.strip()) for x in original_image_size.split(",")]
original_image_location = [int(x.strip()) for x in original_image_location.split(",")]
canvas_size = [int(x.strip()) for x in canvas_size.split(",")]
if prompt == "":
prompt = None
if negative_prompt == "":
negative_prompt = " " # hack to avoid error in triton which expects non-empty string
@@ -68,32 +63,18 @@ class ImageExpansionNode():
# Convert the image directly to Base64 string
image_base64 = image_to_base64(image)
if aspect_ratio and aspect_ratio != "None":
payload = {
"image": image_base64,
"aspect_ratio": aspect_ratio,
# Prepare the API request payload
payload = {
"file": f"{image_base64}",
"original_image_size": original_image_size,
"original_image_location": original_image_location,
"canvas_size": canvas_size,
"prompt": prompt,
"negative_prompt": negative_prompt,
"seed": seed,
"prompt_content_moderation": prompt_content_moderation,
"preserve_alpha": preserve_alpha,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation": visual_output_content_moderation
"content_moderation": content_moderation
}
else:
payload = {
"image": image_base64,
"original_image_size": original_image_size,
"original_image_location": original_image_location,
"canvas_size": canvas_size,
"prompt": prompt,
"negative_prompt": negative_prompt,
"seed": seed,
"prompt_content_moderation": prompt_content_moderation,
"preserve_alpha": preserve_alpha,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation": visual_output_content_moderation
}
headers = {
"Content-Type": "application/json",
@@ -102,34 +83,19 @@ class ImageExpansionNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial image expansion request successful, polling for completion...')
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
image_response = requests.get(response_dict['result_url'])
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}: {response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
+11 -31
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import preprocess_image, image_to_base64, poll_status_until_completed
from .common import preprocess_image, image_to_base64
class RemoveForegroundNode():
@@ -16,9 +16,7 @@ class RemoveForegroundNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -28,10 +26,10 @@ class RemoveForegroundNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/erase_foreground" # remove foreground API URL
self.api_url = "https://engine.internal.prod.bria-api.com/v1/erase_foreground" # remove foreground API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key):
def execute(self, image, content_moderation, api_key):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -46,12 +44,7 @@ class RemoveForegroundNode():
# files=[('file',('temp_img.jpeg', open(temp_img_path, 'rb'),'image/jpeg'))
# ]
payload = {
"image": image_to_base64(image),
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"preserve_alpha": preserve_alpha
}
payload = {"file": image_to_base64(image), "content_moderation": content_moderation}
headers = {
"Content-Type": "application/json",
@@ -60,31 +53,18 @@ class RemoveForegroundNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial request successful, polling for completion...')
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
image_response = requests.get(result_image_url)
image_response = requests.get(response_dict['result_url'])
result_image = Image.open(io.BytesIO(image_response.content))
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} {response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
+35 -53
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed
from .common import image_to_base64, preprocess_image, preprocess_mask
class ReplaceBgNode():
@@ -16,17 +16,16 @@ class ReplaceBgNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"mode": (["base", "fast", "high_control"], {"default": "base"}),
"prompt": ("STRING",),
"ref_images": ("IMAGE",),
"fast": ("BOOLEAN", {"default": True}),
"bg_prompt": ("STRING",),
"ref_image": ("IMAGE",), # Input ref image from another node
"refine_prompt": ("BOOLEAN", {"default": True}),
"enhance_ref_images": ("BOOLEAN", {"default": True}),
"enhance_ref_image": ("BOOLEAN", {"default": True}),
"original_quality": ("BOOLEAN", {"default": False}),
"force_rmbg": ("BOOLEAN", {"default": False}),
"negative_prompt": ("STRING", {"default": None}),
"seed": ("INT", {"default": 681794}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"prompt_content_moderation": ("BOOLEAN", {"default": False}),
"force_background_detection": ("BOOLEAN", {"default": False}),
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -36,21 +35,20 @@ class ReplaceBgNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/replace_background" # Replace BG API URL
self.api_url = "https://engine.prod.bria-api.com/v1/background/replace" # Replace BG API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, mode,
def execute(self, image, fast,
refine_prompt,
enhance_ref_image,
original_quality,
force_rmbg,
negative_prompt,
seed,
api_key,
visual_output_content_moderation,
prompt_content_moderation,
enhance_ref_images,
force_background_detection,
prompt=None,
ref_images=None,):
content_moderation,
bg_prompt=None,
ref_image=None,):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -58,29 +56,28 @@ class ReplaceBgNode():
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
# Convert the image to Base64 string
# Convert the image and mask directly to Base64 strings
image_base64 = image_to_base64(image)
if ref_images is not None:
ref_images = preprocess_image(ref_images)
ref_images = [image_to_base64(ref_images)]
else:
ref_images=[]
ref_image_file = None # initialization, will be updated if it is supplied
if ref_image is not None:
ref_image = preprocess_image(ref_image)
ref_image_file = image_to_base64(ref_image)
# Prepare the API request payload for v2 API
# Prepare the API request payload
payload = {
"image": image_base64,
"mode": mode,
"prompt": prompt,
"ref_images":ref_images,
"file": f"{image_base64}",
"fast": fast,
"bg_prompt": bg_prompt,
"ref_image_file": ref_image_file,
"refine_prompt": refine_prompt,
"enhance_ref_image": enhance_ref_image,
"original_quality": original_quality,
"force_rmbg": force_rmbg,
"negative_prompt": negative_prompt,
"seed": seed,
"prompt_content_moderation": prompt_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"enhance_ref_images":enhance_ref_images,
"force_background_detection": force_background_detection
"sync": True,
"num_results": 1,
"content_moderation": content_moderation
}
headers = {
@@ -90,34 +87,19 @@ class ReplaceBgNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial replace background request successful, polling for completion...')
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
image_response = requests.get(response_dict['result'][0][0]) # first indexing for batched, second for url
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}{response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
+21 -41
View File
@@ -4,7 +4,8 @@ from PIL import Image
import io
import torch
from .common import preprocess_image, image_to_base64, poll_status_until_completed
from .common import preprocess_image
from io import BytesIO
class RmbgNode():
@classmethod
@@ -15,10 +16,7 @@ class RmbgNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
"content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -28,10 +26,10 @@ class RmbgNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/remove_background" # RMBG API URL
self.api_url = "https://engine.prod.bria-api.com/v1/background/remove" # RMBG API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key):
def execute(self, image, content_moderation, api_key):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -39,49 +37,31 @@ class RmbgNode():
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
# Convert image to base64 for the new API format
image_base64 = image_to_base64(image)
payload = {
"image": image_base64,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"preserve_alpha":preserve_alpha
}
# Prepare the API request payload
image_buffer = BytesIO()
image.save(image_buffer, format="JPEG")
headers = {
"Content-Type": "application/json",
"api_token": f"{api_key}"
}
# Get binary data from buffer
image_buffer.seek(0) # Move cursor to the start of the buffer
binary_data = image_buffer.read()
files=[('file',('temp_img.jpeg', BytesIO(binary_data),'image/jpeg'))]
payload = {"content_moderation": content_moderation}
try:
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial RMBG request successful, polling for completion...')
response = requests.post(self.api_url, data=payload, headers={"api_token": api_key}, files=files)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response_dict = response.json()
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
image_response = requests.get(response_dict['result_url'])
result_image = Image.open(io.BytesIO(image_response.content))
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} {response.text}")
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
+4 -3
View File
@@ -10,7 +10,7 @@ class ShotByTextNode():
"required": {
"image": ("IMAGE",), # Input image from another node
"scene_description": ("STRING",),
"mode": (["base", "fast", "high_control"], {"default": "high_control"}),
"optimize_description": ("INT", {"default": 1}),
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value
},
"optional": {
@@ -27,7 +27,7 @@ class ShotByTextNode():
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, mode, content_moderation):
def execute(self, image, api_key, scene_description, optimize_description, content_moderation):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -35,11 +35,12 @@ class ShotByTextNode():
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
optimize_description = bool(optimize_description)
image_base64 = image_to_base64(image)
payload = {
"file": image_base64,
"scene_description": scene_description,
"mode": mode,
"optimize_description": optimize_description,
"placement_type": "original",
"original_quality": True,
"sync": True,
+1 -1
View File
@@ -38,7 +38,7 @@ class Text2ImageBaseNode():
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/base/3.2"
self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/base/2.3" #"http://0.0.0.0:5000/v1/text-to-image/base/2.3"
def execute(
self, api_key, prompt, aspect_ratio, seed, negative_prompt,
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-bria-api"
description = "Custom nodes for ComfyUI using BRIA's API."
version = "2.1.0"
version = "2.0.3"
license = {file = "LICENSE"}
[project.urls]