support load, preview videos nodes
This commit is contained in:
@@ -33,6 +33,8 @@ from .nodes import (
|
||||
VideoMaskByKeyPointsNode,
|
||||
VideoIncreaseResolutionNode,
|
||||
VideoEraseElementsNode,
|
||||
LoadVideoFramesNode,
|
||||
PreviewVideoFramesNode
|
||||
)
|
||||
|
||||
# Map the node class to a name used internally by ComfyUI
|
||||
@@ -71,6 +73,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"VideoMaskByKeyPointsNode":VideoMaskByKeyPointsNode,
|
||||
"VideoIncreaseResolutionNode":VideoIncreaseResolutionNode,
|
||||
"VideoEraseElementsNode":VideoEraseElementsNode,
|
||||
"LoadVideoFramesNode":LoadVideoFramesNode,
|
||||
"PreviewVideoFramesNode":PreviewVideoFramesNode
|
||||
}
|
||||
# Map the node display name to the one shown in the ComfyUI node interface
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -108,4 +112,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VideoMaskByKeyPointsNode":"Bria Video Mask By Key Points",
|
||||
"VideoIncreaseResolutionNode":"Bria Video Increase Resolution",
|
||||
"VideoEraseElementsNode":"Bria Video Erase Elements",
|
||||
"LoadVideoFramesNode":"Load Video Node",
|
||||
"PreviewVideoFramesNode":"Preview Video Node"
|
||||
}
|
||||
|
||||
@@ -34,3 +34,5 @@ from .video_nodes.video_solid_color_background_node import VideoSolidColorBackgr
|
||||
from .video_nodes.video_erase_elements_node import VideoEraseElementsNode
|
||||
from .video_nodes.video_mask_by_prompt_node import VideoMaskByPromptNode
|
||||
from .video_nodes.video_mask_by_key_points_node import VideoMaskByKeyPointsNode
|
||||
from .video_nodes.load_video import LoadVideoFramesNode
|
||||
from .video_nodes.preview_video import PreviewVideoFramesNode
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import folder_paths
|
||||
import av
|
||||
|
||||
class LoadVideoFramesNode:
|
||||
"""
|
||||
Bria Load Video Frames Node
|
||||
|
||||
This node loads a video file and extracts all frames as images.
|
||||
The user can upload a video and it will be converted into a batch of image tensors
|
||||
that can be used with other nodes in the pipeline.
|
||||
|
||||
Parameters:
|
||||
- video: Video file to load from input directory
|
||||
|
||||
Returns:
|
||||
- frames: Batch of images (IMAGE tensor format)
|
||||
- frame_count: Total number of frames extracted
|
||||
- fps: Frames per second of the video
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
|
||||
files = folder_paths.filter_files_content_types(files, ["video"])
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"video": (sorted(files), {"video_upload": True}),
|
||||
},
|
||||
"optional": {
|
||||
"max_frames": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 10000,
|
||||
"step": 1,
|
||||
"tooltip": "Maximum number of frames to extract (0 = all frames)"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "FLOAT", "STRING")
|
||||
RETURN_NAMES = ("frames", "frame_count", "fps", "video_format")
|
||||
FUNCTION = "load_video"
|
||||
CATEGORY = "API Nodes"
|
||||
DESCRIPTION = "Loads a video file and extracts frames as a batch of images."
|
||||
|
||||
def load_video(self, video, max_frames=0):
|
||||
video_path = folder_paths.get_annotated_filepath(video)
|
||||
video_format = os.path.splitext(video)[1].lower().lstrip('.')
|
||||
|
||||
if not os.path.exists(video_path):
|
||||
raise FileNotFoundError(f"Video file not found: {video_path}")
|
||||
|
||||
frames = []
|
||||
frame_index = 0
|
||||
extracted_count = 0
|
||||
|
||||
try:
|
||||
# Open video using PyAV
|
||||
container = av.open(video_path)
|
||||
|
||||
# Get video stream
|
||||
video_stream = None
|
||||
for stream in container.streams:
|
||||
if stream.type == 'video':
|
||||
video_stream = stream
|
||||
break
|
||||
|
||||
if video_stream is None:
|
||||
raise ValueError(f"No video stream found in {video}")
|
||||
|
||||
# Get FPS
|
||||
if video_stream.average_rate is not None:
|
||||
fps = float(video_stream.average_rate)
|
||||
|
||||
print(f"Loading video: {video}")
|
||||
print(f"Video FPS: {fps}")
|
||||
print(f"Video resolution: {video_stream.width}x{video_stream.height}")
|
||||
|
||||
# Decode frames
|
||||
for frame in container.decode(video_stream):
|
||||
# Convert frame to RGB PIL Image
|
||||
img = frame.to_image()
|
||||
|
||||
# Convert PIL Image to numpy array
|
||||
img_array = np.array(img.convert('RGB')).astype(np.float32) / 255.0
|
||||
|
||||
# Convert to torch tensor with shape [H, W, C]
|
||||
img_tensor = torch.from_numpy(img_array)
|
||||
|
||||
frames.append(img_tensor)
|
||||
extracted_count += 1
|
||||
frame_index += 1
|
||||
|
||||
# Check if we've reached max_frames
|
||||
if max_frames > 0 and extracted_count >= max_frames:
|
||||
break
|
||||
|
||||
container.close()
|
||||
|
||||
if len(frames) == 0:
|
||||
raise ValueError(f"No frames could be extracted from {video}")
|
||||
|
||||
# Stack frames into a single tensor with shape [B, H, W, C]
|
||||
frames_tensor = torch.stack(frames, dim=0)
|
||||
|
||||
print(f"Extracted {extracted_count} frames from video")
|
||||
print(f"Output tensor shape: {frames_tensor.shape}")
|
||||
print(f"Video format: {video_format}")
|
||||
|
||||
return (frames_tensor, extracted_count, fps, video_format)
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"Error loading video {video}: {str(e)}")
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, video, **kwargs):
|
||||
"""Force re-execution when video file changes"""
|
||||
video_path = folder_paths.get_annotated_filepath(video)
|
||||
if os.path.exists(video_path):
|
||||
return os.path.getmtime(video_path)
|
||||
return float("nan")
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, video, **kwargs):
|
||||
"""Validate that the video file exists"""
|
||||
if not folder_paths.exists_annotated_filepath(video):
|
||||
return f"Invalid video file: {video}"
|
||||
return True
|
||||
@@ -0,0 +1,284 @@
|
||||
import os
|
||||
import json
|
||||
import random
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import folder_paths
|
||||
from comfy.cli_args import args
|
||||
import av
|
||||
from fractions import Fraction
|
||||
|
||||
class PreviewVideoFramesNode:
|
||||
"""
|
||||
Bria Preview Video Frames Node
|
||||
|
||||
This node takes a batch of image frames and creates an actual video file
|
||||
that can be played as animation in the ComfyUI interface.
|
||||
|
||||
Parameters:
|
||||
- frames: Batch of image tensors to preview
|
||||
- fps: Frames per second for animation display
|
||||
- format: Video format (webp or mp4)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_temp_directory()
|
||||
self.type = "temp"
|
||||
self.prefix_append = "_video_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for _ in range(5))
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of image frames to preview"}),
|
||||
},
|
||||
"optional": {
|
||||
"fps": ("FLOAT", {
|
||||
"default": 30.0,
|
||||
"min": 1.0,
|
||||
"max": 120.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "Frames per second for animation playback"
|
||||
}),
|
||||
"format": ([
|
||||
"mp4_h264",
|
||||
"mp4_h265",
|
||||
"webm_vp9",
|
||||
"mov_h265",
|
||||
"mov_proresks",
|
||||
"mkv_h264",
|
||||
"mkv_h265",
|
||||
"mkv_vp9",
|
||||
"gif",
|
||||
"webp"
|
||||
], {
|
||||
"default": "mp4_h264",
|
||||
"tooltip": "Output video format - matches Bria API formats"
|
||||
}),
|
||||
"quality": (["high", "medium", "low"], {
|
||||
"default": "medium",
|
||||
"tooltip": "Video quality (affects file size)"
|
||||
}),
|
||||
"max_preview_frames": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 1000,
|
||||
"step": 1,
|
||||
"tooltip": "Maximum frames to preview (0 = all frames)"
|
||||
}),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO"
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "preview_frames"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "API Nodes"
|
||||
DESCRIPTION = "Previews video frames as an animated video in the ComfyUI interface."
|
||||
|
||||
def preview_frames(self, frames, fps=30.0, format="mp4_h264", quality="medium", max_preview_frames=0, prompt=None, extra_pnginfo=None):
|
||||
"""
|
||||
Preview video frames as animated video
|
||||
|
||||
Args:
|
||||
frames: Batch of image tensors [B, H, W, C]
|
||||
fps: Frames per second for display
|
||||
format: Output format (supports all Bria API formats)
|
||||
quality: Video quality (high, medium, low)
|
||||
max_preview_frames: Maximum number of frames to preview (0 = all)
|
||||
prompt: Hidden parameter for ComfyUI workflow
|
||||
extra_pnginfo: Hidden parameter for ComfyUI metadata
|
||||
|
||||
Returns:
|
||||
dict: UI output with video file for animation
|
||||
"""
|
||||
filename_prefix = "VideoPreview"
|
||||
filename_prefix += self.prefix_append
|
||||
|
||||
# Limit frames if specified
|
||||
total_frames = frames.shape[0]
|
||||
if max_preview_frames > 0 and total_frames > max_preview_frames:
|
||||
# Sample frames evenly across the video
|
||||
indices = torch.linspace(0, total_frames - 1, max_preview_frames).long()
|
||||
frames = frames[indices]
|
||||
print(f"Preview limited to {max_preview_frames} frames (sampled from {total_frames} total frames)")
|
||||
else:
|
||||
print(f"Previewing {total_frames} frames as animated video (format: {format})")
|
||||
|
||||
# Get save path
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix,
|
||||
self.output_dir,
|
||||
frames[0].shape[1], # width
|
||||
frames[0].shape[0] # height
|
||||
)
|
||||
|
||||
# Choose format and create video
|
||||
if format == "webp":
|
||||
result_file = self._save_as_webp(frames, full_output_folder, filename, counter, fps, quality, prompt, extra_pnginfo)
|
||||
elif format == "gif":
|
||||
result_file = self._save_as_gif(frames, full_output_folder, filename, counter, fps, quality, prompt, extra_pnginfo)
|
||||
else:
|
||||
# All video formats (mp4, webm, mov, mkv)
|
||||
result_file = self._save_as_video(frames, full_output_folder, filename, counter, fps, format, quality, prompt, extra_pnginfo)
|
||||
|
||||
print(f"Saved animated video: {result_file}")
|
||||
|
||||
# Return UI output with animation
|
||||
return {
|
||||
"ui": {
|
||||
"images": [{
|
||||
"filename": result_file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
}],
|
||||
"animated": (True,)
|
||||
}
|
||||
}
|
||||
|
||||
def _get_format_config(self, format_name):
|
||||
"""
|
||||
Get container, codec, and file extension for a given format
|
||||
|
||||
Returns: (container, codec, extension, pixel_format)
|
||||
"""
|
||||
format_map = {
|
||||
"mp4_h264": ("mp4", "libx264", "mp4", "yuv420p"),
|
||||
"mp4_h265": ("mp4", "libx265", "mp4", "yuv420p"),
|
||||
"webm_vp9": ("webm", "libvpx-vp9", "webm", "yuv420p"),
|
||||
"mov_h265": ("mov", "libx265", "mov", "yuv420p"),
|
||||
"mov_proresks": ("mov", "prores_ks", "mov", "yuv422p10le"),
|
||||
"mkv_h264": ("matroska", "libx264", "mkv", "yuv420p"),
|
||||
"mkv_h265": ("matroska", "libx265", "mkv", "yuv420p"),
|
||||
"mkv_vp9": ("matroska", "libvpx-vp9", "mkv", "yuv420p"),
|
||||
}
|
||||
return format_map.get(format_name, ("mp4", "libx264", "mp4", "yuv420p"))
|
||||
|
||||
|
||||
def _save_as_webp(self, frames, output_folder, filename, counter, fps, quality, prompt, extra_pnginfo):
|
||||
"""Save frames as animated WEBP"""
|
||||
file = f"{filename}_{counter:05}_.webp"
|
||||
|
||||
# Convert frames to PIL images
|
||||
pil_images = []
|
||||
for frame in frames:
|
||||
frame_array = 255.0 * frame.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(frame_array, 0, 255).astype(np.uint8))
|
||||
pil_images.append(img)
|
||||
|
||||
# Add metadata to first frame
|
||||
metadata = pil_images[0].getexif()
|
||||
if not args.disable_metadata:
|
||||
if prompt is not None:
|
||||
metadata[0x0110] = "prompt:{}".format(json.dumps(prompt))
|
||||
if extra_pnginfo is not None:
|
||||
inital_exif = 0x010f
|
||||
for x in extra_pnginfo:
|
||||
metadata[inital_exif] = "{}:{}".format(x, json.dumps(extra_pnginfo[x]))
|
||||
inital_exif -= 1
|
||||
|
||||
# Quality mapping
|
||||
quality_values = {"high": 95, "medium": 85, "low": 70}
|
||||
webp_quality = quality_values.get(quality, 85)
|
||||
|
||||
# Save as animated WEBP
|
||||
duration = int(1000.0 / fps) # duration in milliseconds
|
||||
pil_images[0].save(
|
||||
os.path.join(output_folder, file),
|
||||
save_all=True,
|
||||
duration=duration,
|
||||
append_images=pil_images[1:],
|
||||
exif=metadata,
|
||||
lossless=False,
|
||||
quality=webp_quality,
|
||||
method=4
|
||||
)
|
||||
|
||||
return file
|
||||
|
||||
def _save_as_gif(self, frames, output_folder, filename, counter, fps, quality, prompt, extra_pnginfo):
|
||||
"""Save frames as animated GIF"""
|
||||
file = f"{filename}_{counter:05}_.gif"
|
||||
|
||||
# Convert frames to PIL images
|
||||
pil_images = []
|
||||
for frame in frames:
|
||||
frame_array = 255.0 * frame.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(frame_array, 0, 255).astype(np.uint8))
|
||||
pil_images.append(img)
|
||||
|
||||
# Save as animated GIF
|
||||
duration = int(1000.0 / fps) # duration in milliseconds
|
||||
optimize = quality == "low" # Optimize for low quality to reduce size
|
||||
pil_images[0].save(
|
||||
os.path.join(output_folder, file),
|
||||
save_all=True,
|
||||
duration=duration,
|
||||
append_images=pil_images[1:],
|
||||
loop=0,
|
||||
optimize=optimize
|
||||
)
|
||||
|
||||
return file
|
||||
|
||||
def _save_as_video(self, frames, output_folder, filename, counter, fps, format_name, quality, prompt, extra_pnginfo):
|
||||
"""
|
||||
Save frames as video using PyAV with support for all Bria API formats
|
||||
|
||||
Supports: mp4_h264, mp4_h265, webm_vp9, mov_h265, mov_proresks, mkv_h264, mkv_h265, mkv_vp9
|
||||
"""
|
||||
# Get format configuration
|
||||
container_format, codec, extension, pix_fmt = self._get_format_config(format_name)
|
||||
|
||||
file = f"{filename}_{counter:05}_.{extension}"
|
||||
filepath = os.path.join(output_folder, file)
|
||||
|
||||
print(f"Encoding video: container={container_format}, codec={codec}, quality={quality}")
|
||||
|
||||
# Open video container
|
||||
container = av.open(filepath, mode="w", format=container_format)
|
||||
|
||||
# Add metadata if enabled
|
||||
if not args.disable_metadata:
|
||||
if prompt is not None:
|
||||
container.metadata["prompt"] = json.dumps(prompt)
|
||||
if extra_pnginfo is not None:
|
||||
for x in extra_pnginfo:
|
||||
container.metadata[x] = json.dumps(extra_pnginfo[x])
|
||||
|
||||
# Create video stream
|
||||
stream = container.add_stream(codec, rate=Fraction(round(fps * 1000), 1000))
|
||||
stream.width = frames.shape[2] # width
|
||||
stream.height = frames.shape[1] # height
|
||||
stream.pix_fmt = pix_fmt
|
||||
|
||||
# Encode frames
|
||||
for i, frame in enumerate(frames):
|
||||
# Convert tensor to numpy array and scale to 0-255
|
||||
frame_array = torch.clamp(frame * 255, min=0, max=255)
|
||||
frame_np = frame_array.to(device=torch.device("cpu"), dtype=torch.uint8).numpy()
|
||||
|
||||
# Create video frame
|
||||
video_frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
|
||||
|
||||
# Encode frame
|
||||
for packet in stream.encode(video_frame):
|
||||
container.mux(packet)
|
||||
|
||||
# Progress update for long videos
|
||||
if (i + 1) % 100 == 0:
|
||||
print(f"Encoded {i + 1}/{frames.shape[0]} frames...")
|
||||
|
||||
# Flush remaining packets
|
||||
for packet in stream.encode():
|
||||
container.mux(packet)
|
||||
|
||||
container.close()
|
||||
|
||||
print(f"Video encoding complete: {file}")
|
||||
|
||||
return file
|
||||
@@ -1,18 +1,20 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
|
||||
class RemoveVideoBackgroundNode():
|
||||
"""
|
||||
Bria Remove Video Background Node
|
||||
|
||||
This node removes the background from videos using the Bria API.
|
||||
It accepts a publicly accessible video URL and returns a processed video URL
|
||||
with the background removed.
|
||||
It accepts video frames from the Load Video node, uploads to S3,
|
||||
processes via API, and returns the processed frames for preview.
|
||||
|
||||
Supported input resolution: up to 16000x16000 (16K)
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- output_container_and_codec: Output video format and codec (default: webm_vp9)
|
||||
"""
|
||||
@@ -20,11 +22,11 @@ class RemoveVideoBackgroundNode():
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
"optional": {
|
||||
"preserve_audio": ("BOOLEAN", {"default":True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"output_container_and_codec": ([
|
||||
"mp4_h264",
|
||||
"mp4_h265",
|
||||
@@ -36,36 +38,56 @@ class RemoveVideoBackgroundNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "webm_vp9"}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "INT", "FLOAT")
|
||||
RETURN_NAMES = ("frames", "frame_count", "fps")
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/edit/remove_background" # Video RMBG API URL
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/edit/remove_background"
|
||||
|
||||
# Define the execute method as expected by ComfyUI
|
||||
def execute(self, video_url, api_key, preserve_audio, output_container_and_codec="webm_vp9"):
|
||||
def execute(self, frames, api_key, fbs, video_format, preserve_audio=False, output_container_and_codec="webm_vp9"):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"preserve_audio": preserve_audio,
|
||||
"output_container_and_codec": output_container_and_codec
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for background removal...")
|
||||
|
||||
# Step 1: Convert frames to video
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
|
||||
try:
|
||||
# Step 2: Upload video to S3
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename,api_key)
|
||||
|
||||
# Step 3: Call Bria API for background removal
|
||||
print("Step 3: Calling Bria API for background removal...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"preserve_audio": preserve_audio,
|
||||
"output_container_and_codec": output_container_and_codec
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -86,12 +108,23 @@ class RemoveVideoBackgroundNode():
|
||||
|
||||
print(f"Video processing completed. Result URL: {result_video_url}")
|
||||
|
||||
# Step 4: Download and convert processed video to frames
|
||||
print("Step 4: Downloading and converting processed video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_video_url)
|
||||
|
||||
|
||||
return (result_video_url,)
|
||||
print(f"Background removal complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_frame_count, result_fps)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
# Clean up temporary video file
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
@@ -1,31 +1,37 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
|
||||
class VideoEraseElementsNode():
|
||||
"""
|
||||
Bria Video Erase Elements Node
|
||||
|
||||
This node erases specific elements from videos based on a text prompt using the Bria API.
|
||||
It accepts a publicly accessible video URL and a text instruction describing the object
|
||||
to be masked and removed.
|
||||
It accepts video frames from the Load Video node, uploads to S3, processes via API,
|
||||
and returns the processed frames for preview.
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- mask_url: Publicly accessible URL of the mask video (optional)
|
||||
- prompt: Text instruction describing the object to be masked
|
||||
- output_container_and_codec: Output video format and codec (default: mp4_h264)
|
||||
- preserve_audio: Audio preservation (default: True)
|
||||
- preserve_audio: Audio preservation (default: False)
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"mask_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"prompt": ("STRING", {"default": ""}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
"optional": {
|
||||
"mask_url": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": "URL of mask video (optional)"
|
||||
}),
|
||||
"output_container_and_codec": ([
|
||||
"mp4_h264",
|
||||
"mp4_h265",
|
||||
@@ -37,38 +43,59 @@ class VideoEraseElementsNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "mp4_h264"}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "INT", "FLOAT")
|
||||
RETURN_NAMES = ("frames", "frame_count", "fps")
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/edit/erase"
|
||||
|
||||
def execute(self, video_url, mask_url, prompt, api_key, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
def execute(self, frames, prompt, api_key, fbs, video_format, mask_url="", output_container_and_codec="mp4_h264",
|
||||
preserve_audio=False):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"mask":mask_url,
|
||||
"prompt": prompt,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for element erasure...")
|
||||
|
||||
# Step 1: Convert frames to video
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
|
||||
try:
|
||||
# Step 2: Upload video to S3
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename, api_key)
|
||||
|
||||
# Step 3: Call Bria API for element erasure
|
||||
print("Step 3: Calling Bria API for element erasure...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"mask": mask_url,
|
||||
"prompt": prompt,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -89,9 +116,22 @@ class VideoEraseElementsNode():
|
||||
|
||||
print(f"Video processing completed. Result URL: {result_video_url}")
|
||||
|
||||
return (result_video_url,)
|
||||
# Step 4: Download and convert processed video to frames
|
||||
print("Step 4: Downloading and converting processed video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_video_url)
|
||||
|
||||
print(f"Element erasure complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_frame_count, result_fps)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
# Clean up temporary video file
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
@@ -1,26 +1,28 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
|
||||
class VideoIncreaseResolutionNode():
|
||||
"""
|
||||
Bria Video Increase Resolution Node
|
||||
|
||||
This node increases the resolution of videos using the Bria API.
|
||||
It accepts a publicly accessible video URL and returns a processed video URL
|
||||
with increased resolution.
|
||||
It accepts video frames from the Load Video node, uploads to S3,
|
||||
processes via API, and returns the processed frames for preview.
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- desired_increase: Integer scale factor for upscaling (2 or 4)
|
||||
- output_container_and_codec: Output video format and codec (default: mp4_h264)
|
||||
- preserve_audio: Audio preservation (default: True)
|
||||
- preserve_audio: Audio preservation (default: False)
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
"optional": {
|
||||
@@ -36,37 +38,58 @@ class VideoIncreaseResolutionNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "mp4_h264"}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "INT", "FLOAT")
|
||||
RETURN_NAMES = ("frames", "frame_count", "fps")
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/edit/increase_resolution"
|
||||
|
||||
def execute(self, video_url, api_key, desired_increase, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
def execute(self, frames, api_key, video_format, fbs, desired_increase='2', output_container_and_codec="mp4_h264",
|
||||
preserve_audio=False):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"desired_increase": desired_increase,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for resolution increase...")
|
||||
|
||||
# Step 1: Convert frames to video
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
|
||||
try:
|
||||
# Step 2: Upload video to S3
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename, api_key)
|
||||
|
||||
# Step 3: Call Bria API for resolution increase
|
||||
print("Step 3: Calling Bria API for resolution increase...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"desired_increase": desired_increase,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -87,9 +110,22 @@ class VideoIncreaseResolutionNode():
|
||||
|
||||
print(f"Video processing completed. Result URL: {result_video_url}")
|
||||
|
||||
return (result_video_url,)
|
||||
# Step 4: Download and convert processed video to frames
|
||||
print("Step 4: Downloading and converting processed video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_video_url)
|
||||
|
||||
print(f"Resolution increase complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_frame_count, result_fps)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
# Clean up temporary video file
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
@@ -1,5 +1,7 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
import json
|
||||
|
||||
class VideoMaskByKeyPointsNode():
|
||||
@@ -7,21 +9,21 @@ class VideoMaskByKeyPointsNode():
|
||||
Bria Video Mask by Key Points Node
|
||||
|
||||
This node generates a mask for specific elements in videos based on key point coordinates
|
||||
using the Bria API. It accepts a publicly accessible video URL and an array of coordinate
|
||||
objects defining the mask hints.
|
||||
using the Bria API. It accepts video frames from the Load Video node, uploads to S3,
|
||||
processes via API, and returns the processed mask frames for preview.
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- key_points: Array of coordinate objects defining the mask hints (JSON format)
|
||||
- output_container_and_codec: Output video format and codec (default: mp4_h264)
|
||||
- preserve_audio: Audio preservation (default: True)
|
||||
- preserve_audio: Audio preservation (default: False)
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"key_points": ("STRING", {"default": "[]", "multiline": True}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
@@ -37,19 +39,28 @@ class VideoMaskByKeyPointsNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "mp4_h264"}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("mask_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "INT","STRING", "FLOAT",)
|
||||
RETURN_NAMES = ("mask_frames",'mask_url', "frame_count", "fps",)
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/segment/mask_by_key_points"
|
||||
|
||||
def execute(self, video_url, key_points, api_key, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
def execute(self, frames, key_points, api_key,video_format, fbs, output_container_and_codec="mp4_h264",
|
||||
preserve_audio=False):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
@@ -59,20 +70,32 @@ class VideoMaskByKeyPointsNode():
|
||||
except json.JSONDecodeError as e:
|
||||
raise Exception(f"Invalid JSON format for key_points: {e}")
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"key_points": key_points_array,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for video mask generation by key points...")
|
||||
|
||||
# Step 1: Convert frames to video
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
|
||||
try:
|
||||
# Step 2: Upload video to S3
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename, api_key)
|
||||
|
||||
# Step 3: Call Bria API for mask generation
|
||||
print("Step 3: Calling Bria API for video mask generation by key points...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"key_points": key_points_array,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -93,9 +116,22 @@ class VideoMaskByKeyPointsNode():
|
||||
|
||||
print(f"Video mask processing completed. Result URL: {result_mask_url}")
|
||||
|
||||
return (result_mask_url,)
|
||||
# Step 4: Download and convert processed mask video to frames
|
||||
print("Step 4: Downloading and converting processed mask video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_mask_url)
|
||||
|
||||
print(f"Video mask generation complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_mask_url, result_frame_count, result_fps,)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
# Clean up temporary video file
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
@@ -1,26 +1,28 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
|
||||
class VideoMaskByPromptNode():
|
||||
"""
|
||||
Bria Video Mask by Prompt Node
|
||||
|
||||
This node generates a mask for specific elements in videos based on a text prompt
|
||||
using the Bria API. It accepts a publicly accessible video URL and a text instruction
|
||||
describing the object to be masked.
|
||||
using the Bria API. It accepts video frames from the Load Video node, uploads to S3,
|
||||
processes via API, and returns the processed mask frames for preview.
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- prompt: Text instruction describing the object to be masked
|
||||
- output_container_and_codec: Output video format and codec (default: mp4_h264)
|
||||
- preserve_audio: Audio preservation (default: True)
|
||||
- preserve_audio: Audio preservation (default: False)
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"prompt": ("STRING", {"default": ""}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
@@ -36,37 +38,58 @@ class VideoMaskByPromptNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "mp4_h264"}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("mask_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "STRING", "INT", "FLOAT", )
|
||||
RETURN_NAMES = ("mask_frames", "mask_url", "frame_count", "fps")
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/segment/mask_by_prompt"
|
||||
|
||||
def execute(self, video_url, prompt, api_key, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
def execute(self, frames, prompt, api_key, video_format, fbs, output_container_and_codec="mp4_h264",
|
||||
preserve_audio=False):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"prompt": prompt,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for video mask generation...")
|
||||
|
||||
# Step 1: Convert frames to video
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
|
||||
try:
|
||||
# Step 2: Upload video to S3
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename, api_key)
|
||||
|
||||
# Step 3: Call Bria API for mask generation
|
||||
print("Step 3: Calling Bria API for video mask generation...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"prompt": prompt,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -87,9 +110,22 @@ class VideoMaskByPromptNode():
|
||||
|
||||
print(f"Video mask processing completed. Result URL: {result_mask_url}")
|
||||
|
||||
return (result_mask_url,)
|
||||
# Step 4: Download and convert processed mask video to frames
|
||||
print("Step 4: Downloading and converting processed mask video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_mask_url)
|
||||
|
||||
print(f"Video mask generation complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_mask_url, result_frame_count, result_fps,)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
# Clean up temporary video file
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
@@ -1,26 +1,28 @@
|
||||
import os
|
||||
import requests
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from .video_utils import frames_to_video, video_to_frames, upload_video_to_s3
|
||||
|
||||
class VideoSolidColorBackgroundNode():
|
||||
"""
|
||||
Bria Video Solid Color Background Node
|
||||
|
||||
This node removes the background from videos and replaces it with a solid color
|
||||
using the Bria API. It accepts a publicly accessible video URL and returns a
|
||||
processed video URL with the specified background color.
|
||||
using the Bria API. It accepts video frames from the Load Video node, uploads to S3,
|
||||
processes via API, and returns the processed frames for preview.
|
||||
|
||||
Parameters:
|
||||
- video_url: Publicly accessible URL of the input video
|
||||
- frames: Batch of image frames from Load Video node
|
||||
- api_key: Your Bria API token
|
||||
- background_color: Predefined color string (default: Transparent)
|
||||
- output_container_and_codec: Output video format and codec (default: mp4_h264)
|
||||
- preserve_audio: Audio preservation (default: True)
|
||||
- preserve_audio: Audio preservation (default: False)
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"video_url": ("STRING", {"default": ""}),
|
||||
"frames": ("IMAGE", {"tooltip": "Batch of video frames"}),
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}),
|
||||
},
|
||||
"optional": {
|
||||
@@ -48,37 +50,55 @@ class VideoSolidColorBackgroundNode():
|
||||
"mkv_vp9",
|
||||
"gif"
|
||||
], {"default": "mp4_h264"}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": True}),
|
||||
"preserve_audio": ("BOOLEAN", {"default": False}),
|
||||
"video_format": ("STRING", {
|
||||
"default": "mp4",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
"fbs": ("FLOAT", {
|
||||
"default": "30",
|
||||
"tooltip": "Original video format from Load Video node"
|
||||
}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url_response",)
|
||||
RETURN_TYPES = ("IMAGE", "INT", "FLOAT")
|
||||
RETURN_NAMES = ("frames", "frame_count", "fps")
|
||||
CATEGORY = "API Nodes"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v2/video/edit/remove_background"
|
||||
|
||||
def execute(self, video_url, api_key, background_color="Transparent", output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
def execute(self, frames, api_key, video_format,fbs, background_color="Transparent", output_container_and_codec="mp4_h264",
|
||||
preserve_audio=False):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"background_color": background_color,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
print(f"Processing {frames.shape[0]} frames for solid color background...")
|
||||
|
||||
print("Step 1: Converting frames to video...")
|
||||
video_path = frames_to_video(frames, fbs, video_format=video_format)
|
||||
try:
|
||||
print("Step 2: Uploading video to S3...")
|
||||
filename = f"input_video_{os.path.basename(video_path)}"
|
||||
video_url = upload_video_to_s3(video_path, filename, api_key)
|
||||
|
||||
print("Step 3: Calling Bria API for solid color background...")
|
||||
payload = {
|
||||
"video": video_url,
|
||||
"background_color": background_color,
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200 or response.status_code == 202:
|
||||
@@ -99,9 +119,20 @@ class VideoSolidColorBackgroundNode():
|
||||
|
||||
print(f"Video processing completed. Result URL: {result_video_url}")
|
||||
|
||||
return (result_video_url,)
|
||||
print("Step 4: Downloading and converting processed video to frames...")
|
||||
result_frames, result_frame_count, result_fps = video_to_frames(result_video_url)
|
||||
|
||||
print(f"Solid color background processing complete! Processed {result_frame_count} frames.")
|
||||
|
||||
return (result_frames, result_frame_count, result_fps)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
raise Exception(f"{e}")
|
||||
finally:
|
||||
try:
|
||||
if os.path.exists(video_path):
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
@@ -0,0 +1,230 @@
|
||||
import mimetypes
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import av
|
||||
import requests
|
||||
import tempfile
|
||||
from fractions import Fraction
|
||||
|
||||
def get_video_config(ext):
|
||||
ext = ext.lower().replace(".", "")
|
||||
|
||||
configs = {
|
||||
"mp4": ("mp4", "libx264", "yuv420p"),
|
||||
"mov": ("mov", "libx264", "yuv420p"),
|
||||
"mkv": ("matroska", "libx264", "yuv420p"),
|
||||
"webm": ("webm", "libvpx-vp9", "yuv420p"),
|
||||
"gif": ("gif", "gif", "pal8"),
|
||||
}
|
||||
|
||||
if ext in configs:
|
||||
return configs[ext]
|
||||
|
||||
print(f"[WARNING] Unknown video format '{ext}', falling back to mp4/h264.")
|
||||
return ("mp4", "libx264", "yuv420p")
|
||||
|
||||
|
||||
|
||||
def frames_to_video(frames, fps, video_format="mp4"):
|
||||
"""
|
||||
Convert a batch of image frames to a video file.
|
||||
Automatically selects correct container/codec/pixel-format based on extension.
|
||||
"""
|
||||
# Pick container, codec and pix_fmt
|
||||
container_format, codec, pix_fmt = get_video_config(video_format)
|
||||
|
||||
# Create temp file
|
||||
temp_file = tempfile.NamedTemporaryFile(suffix=f".{video_format}", delete=False)
|
||||
filepath = temp_file.name
|
||||
temp_file.close()
|
||||
|
||||
print(f"Converting {frames.shape[0]} frames → {video_format}")
|
||||
print(f"Container={container_format}, Codec={codec}, PixFmt={pix_fmt}")
|
||||
|
||||
# Open container with correct format
|
||||
container = av.open(filepath, mode="w", format=container_format)
|
||||
|
||||
# Create video stream
|
||||
stream = container.add_stream(codec, rate=Fraction(round(fps * 1000), 1000))
|
||||
stream.width = frames.shape[2]
|
||||
stream.height = frames.shape[1]
|
||||
stream.pix_fmt = pix_fmt
|
||||
|
||||
# Encode frames
|
||||
for i, frame in enumerate(frames):
|
||||
frame_array = torch.clamp(frame * 255, 0, 255)
|
||||
frame_np = frame_array.to(dtype=torch.uint8, device="cpu").numpy()
|
||||
|
||||
|
||||
video_frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
|
||||
|
||||
for packet in stream.encode(video_frame):
|
||||
container.mux(packet)
|
||||
|
||||
if (i + 1) % 100 == 0:
|
||||
print(f"Encoded {i + 1}/{frames.shape[0]} frames...")
|
||||
|
||||
# Flush remaining packets
|
||||
for packet in stream.encode():
|
||||
container.mux(packet)
|
||||
|
||||
container.close()
|
||||
return filepath
|
||||
|
||||
|
||||
def video_to_frames(video_path):
|
||||
"""
|
||||
Convert a video file to a batch of image frames
|
||||
|
||||
Args:
|
||||
video_path: Path to video file (local or URL)
|
||||
|
||||
Returns:
|
||||
tuple: (frames_tensor, frame_count, fps)
|
||||
"""
|
||||
# Download if it's a URL
|
||||
if video_path.startswith("http://") or video_path.startswith("https://"):
|
||||
print(f"Downloading video from: {video_path}")
|
||||
response = requests.get(video_path, stream=True)
|
||||
if response.status_code != 200:
|
||||
raise Exception(f"Failed to download video: {response.status_code}")
|
||||
|
||||
# Save to temporary file
|
||||
temp_file = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False)
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file.close()
|
||||
video_path = temp_file.name
|
||||
print(f"Video downloaded to: {video_path}")
|
||||
|
||||
if not os.path.exists(video_path):
|
||||
raise FileNotFoundError(f"Video file not found: {video_path}")
|
||||
|
||||
frames = []
|
||||
fps = 24.0 # Default FPS
|
||||
|
||||
try:
|
||||
# Open video using PyAV
|
||||
container = av.open(video_path)
|
||||
|
||||
# Get video stream
|
||||
video_stream = None
|
||||
for stream in container.streams:
|
||||
if stream.type == 'video':
|
||||
video_stream = stream
|
||||
break
|
||||
|
||||
if video_stream is None:
|
||||
raise ValueError(f"No video stream found in video")
|
||||
|
||||
# Get FPS
|
||||
if video_stream.average_rate is not None:
|
||||
fps = float(video_stream.average_rate)
|
||||
|
||||
print(f"Loading video: {video_path}")
|
||||
print(f"Video FPS: {fps}")
|
||||
print(f"Video resolution: {video_stream.width}x{video_stream.height}")
|
||||
|
||||
# Decode frames
|
||||
for frame in container.decode(video_stream):
|
||||
# Convert frame to RGB PIL Image
|
||||
img = frame.to_image()
|
||||
|
||||
# Convert PIL Image to numpy array
|
||||
img_array = np.array(img.convert('RGB')).astype(np.float32) / 255.0
|
||||
|
||||
# Convert to torch tensor with shape [H, W, C]
|
||||
img_tensor = torch.from_numpy(img_array)
|
||||
|
||||
frames.append(img_tensor)
|
||||
|
||||
container.close()
|
||||
|
||||
if len(frames) == 0:
|
||||
raise ValueError(f"No frames could be extracted from video")
|
||||
|
||||
# Stack frames into a single tensor with shape [B, H, W, C]
|
||||
frames_tensor = torch.stack(frames, dim=0)
|
||||
|
||||
print(f"Extracted {len(frames)} frames from video")
|
||||
print(f"Output tensor shape: {frames_tensor.shape}")
|
||||
|
||||
# Clean up temporary file if downloaded
|
||||
if video_path.startswith(tempfile.gettempdir()):
|
||||
try:
|
||||
os.unlink(video_path)
|
||||
except:
|
||||
pass
|
||||
|
||||
return (frames_tensor, len(frames), fps)
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"Error loading video: {str(e)}")
|
||||
|
||||
|
||||
def upload_video_to_s3(video_path, filename, api_token):
|
||||
api_url = "https://platform.prod.bria-api.com/upload-video/anonymous/presigned-url"
|
||||
headers = {
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
extension = os.path.splitext(filename)[1].lower()
|
||||
content_type_map = {
|
||||
'.mp4': 'video/mp4',
|
||||
'.webm': 'video/webm',
|
||||
'.mov': 'video/quicktime',
|
||||
'.mkv': 'video/x-matroska',
|
||||
'.avi': 'video/x-msvideo',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp'
|
||||
}
|
||||
content_type = content_type_map.get(extension, 'video/mp4')
|
||||
if api_token:
|
||||
headers["api_token"] = api_token
|
||||
|
||||
payload = {
|
||||
"file_name": filename,
|
||||
"content_type":content_type
|
||||
}
|
||||
|
||||
print(f"Requesting presigned URL for: {filename}")
|
||||
|
||||
try:
|
||||
response = requests.post(api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(f"Failed to get presigned URL: {response.status_code} {response.text}")
|
||||
|
||||
response_data = response.json()
|
||||
video_url = response_data.get("video_url")
|
||||
upload_url = response_data.get("upload_url")
|
||||
|
||||
if not video_url or not upload_url:
|
||||
raise Exception(f"Invalid response from presigned URL API: {response_data}")
|
||||
|
||||
print(f"Received presigned URL")
|
||||
print(f"Video URL: {video_url}")
|
||||
|
||||
# Step 2: Upload video to presigned URL
|
||||
print(f"Uploading video to S3...")
|
||||
|
||||
with open(video_path, 'rb') as f:
|
||||
video_data = f.read()
|
||||
|
||||
# Determine content type based on file extension
|
||||
upload_headers = {
|
||||
"Content-Type": content_type
|
||||
}
|
||||
|
||||
upload_response = requests.put(upload_url, data=video_data, headers=upload_headers)
|
||||
|
||||
if upload_response.status_code not in [200, 204]:
|
||||
raise Exception(f"Failed to upload video to S3: {upload_response.status_code}")
|
||||
|
||||
print(f"Video uploaded successfully to S3")
|
||||
|
||||
return video_url
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"Error uploading video to S3: {str(e)}")
|
||||
|
||||
Reference in New Issue
Block a user