Compare commits
23
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f9b21a5e93 | ||
|
|
fbee93b5b5 | ||
|
|
845b9d46c5 | ||
|
|
31572e6e45 | ||
|
|
b60c18d8a8 | ||
|
|
58c54acbce | ||
|
|
27580456ed | ||
|
|
a68c56134c | ||
|
|
cd9eb99568 | ||
|
|
ef774a511b | ||
|
|
66d4dcf54d | ||
|
|
a6d061c0eb | ||
|
|
34d3a8396e | ||
|
|
06a30a6f21 | ||
|
|
f4f486edb0 | ||
|
|
1f6f476679 | ||
|
|
1e561ac944 | ||
|
|
cf523888a7 | ||
|
|
93aa2cbc04 | ||
|
|
5be02175f3 | ||
|
|
4ff17aa6ef | ||
|
|
a6d29a2d4c | ||
|
|
4215edebf0 |
@@ -44,6 +44,11 @@ Custom nodes for using Flux models with fal API in ComfyUI with only one API Ke
|
||||
FAL_KEY = your_actual_api_key
|
||||
```
|
||||
|
||||
4. Alternatively, you can set the FAL_KEY environment variable:
|
||||
```bash
|
||||
export FAL_KEY=your_actual_api_key
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
After installation and configuration, restart ComfyUI. The new nodes will be available in the node browser under the "FAL" category.
|
||||
|
||||
+12
-8
@@ -1,12 +1,13 @@
|
||||
import importlib.util
|
||||
import importlib
|
||||
import importlib.util
|
||||
|
||||
node_list = [
|
||||
"image_node",
|
||||
"video_node",
|
||||
"llm_node",
|
||||
"vlm_node",
|
||||
"trainer_node",
|
||||
"image_node",
|
||||
"video_node",
|
||||
"llm_node",
|
||||
"vlm_node",
|
||||
"trainer_node",
|
||||
"upscaler_node",
|
||||
]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
@@ -16,7 +17,10 @@ for module_name in node_list:
|
||||
imported_module = importlib.import_module(f".nodes.{module_name}", __name__)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
**NODE_DISPLAY_NAME_MAPPINGS,
|
||||
**imported_module.NODE_DISPLAY_NAME_MAPPINGS,
|
||||
}
|
||||
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
import configparser
|
||||
import io
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
from fal_client.client import SyncClient
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class FalConfig:
|
||||
"""Singleton class to handle FAL configuration and client setup."""
|
||||
|
||||
_instance = None
|
||||
_client = None
|
||||
_key = None
|
||||
|
||||
def __new__(cls):
|
||||
if cls._instance is None:
|
||||
cls._instance = super(FalConfig, cls).__new__(cls)
|
||||
cls._instance._initialize()
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self):
|
||||
"""Initialize configuration and API key."""
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
self._key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
self._key = config["API"]["FAL_KEY"]
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = self._key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if self._key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
def get_client(self):
|
||||
"""Get or create the FAL client."""
|
||||
if self._client is None:
|
||||
self._client = SyncClient(key=self._key)
|
||||
return self._client
|
||||
|
||||
def get_key(self):
|
||||
"""Get the FAL API key."""
|
||||
return self._key
|
||||
|
||||
|
||||
class ImageUtils:
|
||||
"""Utility functions for image processing."""
|
||||
|
||||
@staticmethod
|
||||
def tensor_to_pil(image):
|
||||
"""Convert image tensor to PIL Image."""
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(
|
||||
image_np, (1, 2, 0)
|
||||
) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
return Image.fromarray(image_np)
|
||||
except Exception as e:
|
||||
print(f"Error converting tensor to PIL: {str(e)}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def upload_image(image):
|
||||
"""Upload image tensor to FAL and return URL."""
|
||||
try:
|
||||
pil_image = ImageUtils.tensor_to_pil(image)
|
||||
if not pil_image:
|
||||
return None
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
client = FalConfig().get_client()
|
||||
image_url = client.upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if "temp_file_path" in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
@staticmethod
|
||||
def mask_to_image(mask):
|
||||
"""Convert mask tensor to image tensor."""
|
||||
result = (
|
||||
mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1]))
|
||||
.movedim(1, -1)
|
||||
.expand(-1, -1, -1, 3)
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
class ResultProcessor:
|
||||
"""Utility functions for processing API results."""
|
||||
|
||||
@staticmethod
|
||||
def process_image_result(result):
|
||||
"""Process image generation result and return tensor."""
|
||||
try:
|
||||
images = []
|
||||
for img_info in result["images"]:
|
||||
img_url = img_info["url"]
|
||||
img_response = requests.get(img_url)
|
||||
img = Image.open(io.BytesIO(img_response.content))
|
||||
img_array = np.array(img).astype(np.float32) / 255.0
|
||||
images.append(img_array)
|
||||
|
||||
# Stack the images along a new first dimension
|
||||
stacked_images = np.stack(images, axis=0)
|
||||
|
||||
# Convert to PyTorch tensor
|
||||
img_tensor = torch.from_numpy(stacked_images)
|
||||
|
||||
return (img_tensor,)
|
||||
except Exception as e:
|
||||
print(f"Error processing image result: {str(e)}")
|
||||
return ResultProcessor.create_blank_image()
|
||||
|
||||
@staticmethod
|
||||
def process_single_image_result(result):
|
||||
"""Process single image result and return tensor."""
|
||||
try:
|
||||
img_url = result["image"]["url"]
|
||||
img_response = requests.get(img_url)
|
||||
img = Image.open(io.BytesIO(img_response.content))
|
||||
img_array = np.array(img).astype(np.float32) / 255.0
|
||||
|
||||
# Stack the images along a new first dimension
|
||||
stacked_images = np.stack([img_array], axis=0)
|
||||
|
||||
# Convert to PyTorch tensor
|
||||
img_tensor = torch.from_numpy(stacked_images)
|
||||
return (img_tensor,)
|
||||
except Exception as e:
|
||||
print(f"Error processing single image result: {str(e)}")
|
||||
return ResultProcessor.create_blank_image()
|
||||
|
||||
@staticmethod
|
||||
def create_blank_image():
|
||||
"""Create a blank black image tensor."""
|
||||
blank_img = Image.new("RGB", (512, 512), color="black")
|
||||
img_array = np.array(blank_img).astype(np.float32) / 255.0
|
||||
img_tensor = torch.from_numpy(img_array)[None,]
|
||||
return (img_tensor,)
|
||||
|
||||
|
||||
class ApiHandler:
|
||||
"""Utility functions for API interactions."""
|
||||
|
||||
@staticmethod
|
||||
def submit_and_get_result(endpoint, arguments):
|
||||
"""Submit job to FAL API and get result."""
|
||||
try:
|
||||
client = FalConfig().get_client()
|
||||
handler = client.submit(endpoint, arguments=arguments)
|
||||
return handler.get()
|
||||
except Exception as e:
|
||||
print(f"Error submitting to {endpoint}: {str(e)}")
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
def handle_video_generation_error(model_name, error):
|
||||
"""Handle video generation errors consistently."""
|
||||
print(f"Error generating video with {model_name}: {str(error)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
@staticmethod
|
||||
def handle_image_generation_error(model_name, error):
|
||||
"""Handle image generation errors consistently."""
|
||||
print(f"Error generating image with {model_name}: {str(error)}")
|
||||
return ResultProcessor.create_blank_image()
|
||||
|
||||
@staticmethod
|
||||
def handle_text_generation_error(model_name, error):
|
||||
"""Handle text generation errors consistently."""
|
||||
print(f"Error generating text with {model_name}: {str(error)}")
|
||||
return ("Error: Unable to generate text.",)
|
||||
+978
-359
File diff suppressed because it is too large
Load Diff
+28
-47
@@ -1,37 +1,8 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
from .fal_utils import ApiHandler, FalConfig
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
class LLMNode:
|
||||
@classmethod
|
||||
@@ -39,11 +10,22 @@ class LLMNode:
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"model": (["google/gemini-flash-1.5-8b", "anthropic/claude-3.5-sonnet", "anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5", "google/gemini-flash-1.5", "meta-llama/llama-3.2-1b-instruct",
|
||||
"meta-llama/llama-3.2-3b-instruct", "meta-llama/llama-3.1-8b-instruct",
|
||||
"meta-llama/llama-3.1-70b-instruct", "openai/gpt-4o-mini", "openai/gpt-4o"],
|
||||
{"default": "google/gemini-flash-1.5-8b"}),
|
||||
"model": (
|
||||
[
|
||||
"google/gemini-flash-1.5-8b",
|
||||
"anthropic/claude-3.5-sonnet",
|
||||
"anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5",
|
||||
"google/gemini-flash-1.5",
|
||||
"meta-llama/llama-3.2-1b-instruct",
|
||||
"meta-llama/llama-3.2-3b-instruct",
|
||||
"meta-llama/llama-3.1-8b-instruct",
|
||||
"meta-llama/llama-3.1-70b-instruct",
|
||||
"openai/gpt-4o-mini",
|
||||
"openai/gpt-4o",
|
||||
],
|
||||
{"default": "google/gemini-flash-1.5-8b"},
|
||||
),
|
||||
"system_prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
},
|
||||
}
|
||||
@@ -53,19 +35,18 @@ class LLMNode:
|
||||
CATEGORY = "FAL/LLM"
|
||||
|
||||
def generate_text(self, prompt, model, system_prompt):
|
||||
arguments = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
}
|
||||
|
||||
try:
|
||||
handler = fal_client.submit("fal-ai/any-llm", arguments=arguments)
|
||||
result = handler.get()
|
||||
arguments = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
}
|
||||
|
||||
result = ApiHandler.submit_and_get_result("fal-ai/any-llm", arguments)
|
||||
return (result["output"],)
|
||||
except Exception as e:
|
||||
print(f"Error generating text with LLM: {str(e)}")
|
||||
return ("Error: Unable to generate text.",)
|
||||
return ApiHandler.handle_text_generation_error(model, str(e))
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
+193
-125
@@ -1,68 +1,52 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
import tempfile
|
||||
import zipfile
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
from .fal_utils import ApiHandler, FalConfig
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
def create_zip_from_images(images):
|
||||
"""Create a zip file from a list of images."""
|
||||
with tempfile.NamedTemporaryFile(suffix='.zip', delete=False) as temp_zip:
|
||||
with zipfile.ZipFile(temp_zip, 'w') as zf:
|
||||
for idx, img_tensor in enumerate(images):
|
||||
# Convert tensor to PIL Image
|
||||
if isinstance(img_tensor, torch.Tensor):
|
||||
# Convert to numpy and scale to 0-255 range
|
||||
img_np = (img_tensor.cpu().numpy() * 255).astype('uint8')
|
||||
# Handle different tensor formats
|
||||
if img_np.shape[0] == 3: # If in format (C, H, W)
|
||||
img_np = img_np.transpose(1, 2, 0)
|
||||
img = Image.fromarray(img_np)
|
||||
else:
|
||||
img = img_tensor
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as temp_zip:
|
||||
with zipfile.ZipFile(temp_zip, "w") as zf:
|
||||
for idx, img_tensor in enumerate(images):
|
||||
# Convert tensor to PIL Image
|
||||
if isinstance(img_tensor, torch.Tensor):
|
||||
# Convert to numpy and scale to 0-255 range
|
||||
img_np = (img_tensor.cpu().numpy() * 255).astype("uint8")
|
||||
# Handle different tensor formats
|
||||
if img_np.shape[0] == 3: # If in format (C, H, W)
|
||||
img_np = img_np.transpose(1, 2, 0)
|
||||
img = Image.fromarray(img_np)
|
||||
else:
|
||||
img = img_tensor
|
||||
|
||||
# Save image to temporary file
|
||||
with tempfile.NamedTemporaryFile(
|
||||
suffix=".png", delete=False
|
||||
) as temp_img:
|
||||
img.save(temp_img, format="PNG")
|
||||
temp_img_path = temp_img.name
|
||||
|
||||
# Add to zip file
|
||||
zf.write(temp_img_path, f"image_{idx}.png")
|
||||
os.unlink(temp_img_path)
|
||||
|
||||
# Use fal_client.upload_file instead of ApiHandler.upload_file
|
||||
client = FalConfig().get_client()
|
||||
return client.upload_file(temp_zip.name)
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"flux-lora-fast-training", f"Failed to create zip file: {str(e)}"
|
||||
)
|
||||
|
||||
# Save image to temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix='.png', delete=False) as temp_img:
|
||||
img.save(temp_img, format='PNG')
|
||||
temp_img_path = temp_img.name
|
||||
|
||||
# Add to zip file
|
||||
zf.write(temp_img_path, f'image_{idx}.png')
|
||||
os.unlink(temp_img_path)
|
||||
|
||||
return fal_client.upload_file(temp_zip.name)
|
||||
|
||||
class FluxLoraTrainerNode:
|
||||
@classmethod
|
||||
@@ -70,7 +54,10 @@ class FluxLoraTrainerNode:
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 100}),
|
||||
"steps": (
|
||||
"INT",
|
||||
{"default": 1000, "min": 100, "max": 10000, "step": 100},
|
||||
),
|
||||
"create_masks": ("BOOLEAN", {"default": True}),
|
||||
"is_style": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
@@ -79,21 +66,34 @@ class FluxLoraTrainerNode:
|
||||
"images_zip_url": ("STRING", {"default": ""}),
|
||||
"is_input_format_already_preprocessed": ("BOOLEAN", {"default": False}),
|
||||
"data_archive_format": ("STRING", {"default": ""}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, images, steps, create_masks, is_style, trigger_word="", images_zip_url="",
|
||||
is_input_format_already_preprocessed=False, data_archive_format=""):
|
||||
def train_lora(
|
||||
self,
|
||||
images,
|
||||
steps,
|
||||
create_masks,
|
||||
is_style,
|
||||
trigger_word="",
|
||||
images_zip_url="",
|
||||
is_input_format_already_preprocessed=False,
|
||||
data_archive_format="",
|
||||
):
|
||||
try:
|
||||
# Use provided zip URL if available, otherwise create and upload zip file
|
||||
images_url = images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
images_url = (
|
||||
images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
)
|
||||
if not images_url:
|
||||
return ("Error: Unable to upload images.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"flux-lora-fast-training", "Failed to upload images"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
@@ -103,24 +103,25 @@ class FluxLoraTrainerNode:
|
||||
"is_style": is_style,
|
||||
"is_input_format_already_preprocessed": is_input_format_already_preprocessed,
|
||||
}
|
||||
|
||||
|
||||
if trigger_word:
|
||||
arguments["trigger_word"] = trigger_word
|
||||
|
||||
|
||||
if data_archive_format:
|
||||
arguments["data_archive_format"] = data_archive_format
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/flux-lora-fast-training", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/flux-lora-fast-training", arguments
|
||||
)
|
||||
lora_url = result["diffusers_lora_file"]["url"]
|
||||
|
||||
return (lora_url, )
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"flux-lora-fast-training", str(e)
|
||||
)
|
||||
|
||||
|
||||
class HunyuanVideoLoraTrainerNode:
|
||||
@classmethod
|
||||
@@ -128,55 +129,74 @@ class HunyuanVideoLoraTrainerNode:
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 100}),
|
||||
"steps": (
|
||||
"INT",
|
||||
{"default": 1000, "min": 100, "max": 10000, "step": 100},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_word": ("STRING", {"default": ""}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0001, "min": 0.00001, "max": 0.01}),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{"default": 0.0001, "min": 0.00001, "max": 0.01},
|
||||
),
|
||||
"do_caption": ("BOOLEAN", {"default": True}),
|
||||
"images_zip_url": ("STRING", {"default": ""}),
|
||||
"data_archive_format": ("STRING", {"default": ""}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, images, steps, trigger_word="", learning_rate=0.0001, do_caption=True,
|
||||
images_zip_url="", data_archive_format=""):
|
||||
def train_lora(
|
||||
self,
|
||||
images,
|
||||
steps,
|
||||
trigger_word="",
|
||||
learning_rate=0.0001,
|
||||
do_caption=True,
|
||||
images_zip_url="",
|
||||
data_archive_format="",
|
||||
):
|
||||
try:
|
||||
# Use provided zip URL if available, otherwise create and upload zip file
|
||||
images_url = images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
images_url = (
|
||||
images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
)
|
||||
if not images_url:
|
||||
return ("Error: Unable to upload images.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"hunyuan-video-lora-training", "Failed to upload images"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
"images_data_url": images_url,
|
||||
"steps": steps,
|
||||
"learning_rate": learning_rate,
|
||||
"do_caption": do_caption
|
||||
"do_caption": do_caption,
|
||||
}
|
||||
|
||||
|
||||
if trigger_word:
|
||||
arguments["trigger_word"] = trigger_word
|
||||
|
||||
|
||||
if data_archive_format:
|
||||
arguments["data_archive_format"] = data_archive_format
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/hunyuan-video-lora-training", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/hunyuan-video-lora-training", arguments
|
||||
)
|
||||
lora_url = result["diffusers_lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"hunyuan-video-lora-training", str(e)
|
||||
)
|
||||
|
||||
|
||||
class WanLoraTrainerNode:
|
||||
@classmethod
|
||||
@@ -184,47 +204,59 @@ class WanLoraTrainerNode:
|
||||
return {
|
||||
"required": {
|
||||
"training_data_url": ("STRING", {"default": ""}),
|
||||
"number_of_steps": ("INT", {"default": 400, "min": 5, "max": 10000, "step": 1}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0002, "min": 0.00001, "max": 0.01}),
|
||||
"number_of_steps": (
|
||||
"INT",
|
||||
{"default": 400, "min": 5, "max": 10000, "step": 1},
|
||||
),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{"default": 0.0002, "min": 0.00001, "max": 0.01},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_phrase": ("STRING", {"default": ""}),
|
||||
"auto_scale_input": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, training_data_url, number_of_steps, learning_rate, trigger_phrase="", auto_scale_input=True):
|
||||
def train_lora(
|
||||
self,
|
||||
training_data_url,
|
||||
number_of_steps,
|
||||
learning_rate,
|
||||
trigger_phrase="",
|
||||
auto_scale_input=True,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
return ("Error: No training data URL provided.",)
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"wan-trainer", "No training data URL provided"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
"training_data_url": training_data_url,
|
||||
"number_of_steps": number_of_steps,
|
||||
"learning_rate": learning_rate,
|
||||
"auto_scale_input": auto_scale_input
|
||||
"auto_scale_input": auto_scale_input,
|
||||
}
|
||||
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/wan-trainer", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result("fal-ai/wan-trainer", arguments)
|
||||
lora_url = result["lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.",)
|
||||
return ApiHandler.handle_text_generation_error("wan-trainer", str(e))
|
||||
|
||||
|
||||
class LtxVideoTrainerNode:
|
||||
@classmethod
|
||||
@@ -233,40 +265,77 @@ class LtxVideoTrainerNode:
|
||||
"required": {
|
||||
"training_data_url": ("STRING", {"default": ""}),
|
||||
"rank": (["8", "16", "32", "64", "128"], {"default": "128"}),
|
||||
"number_of_steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 1}),
|
||||
"number_of_steps": (
|
||||
"INT",
|
||||
{"default": 1000, "min": 100, "max": 10000, "step": 1},
|
||||
),
|
||||
"number_of_frames": ("INT", {"default": 81, "min": 1, "max": 1000}),
|
||||
"frame_rate": ("INT", {"default": 25, "min": 1, "max": 60}),
|
||||
"resolution": (["low", "medium", "high"], {"default": "medium"}),
|
||||
"aspect_ratio": (["16:9", "1:1", "9:16"], {"default": "1:1"}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0002, "min": 0.00001, "max": 0.01}),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{"default": 0.0002, "min": 0.00001, "max": 0.01},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_phrase": ("STRING", {"default": ""}),
|
||||
"auto_scale_input": ("BOOLEAN", {"default": False}),
|
||||
"split_input_into_scenes": ("BOOLEAN", {"default": True}),
|
||||
"split_input_duration_threshold": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 300.0}),
|
||||
"validation_negative_prompt": ("STRING", {"default": "blurry, low quality, bad quality, out of focus"}),
|
||||
"validation_number_of_frames": ("INT", {"default": 81, "min": 1, "max": 1000}),
|
||||
"validation_resolution": (["low", "medium", "high"], {"default": "high"}),
|
||||
"validation_aspect_ratio": (["16:9", "1:1", "9:16"], {"default": "1:1"}),
|
||||
"split_input_duration_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 30.0, "min": 1.0, "max": 300.0},
|
||||
),
|
||||
"validation_negative_prompt": (
|
||||
"STRING",
|
||||
{"default": "blurry, low quality, bad quality, out of focus"},
|
||||
),
|
||||
"validation_number_of_frames": (
|
||||
"INT",
|
||||
{"default": 81, "min": 1, "max": 1000},
|
||||
),
|
||||
"validation_resolution": (
|
||||
["low", "medium", "high"],
|
||||
{"default": "high"},
|
||||
),
|
||||
"validation_aspect_ratio": (
|
||||
["16:9", "1:1", "9:16"],
|
||||
{"default": "1:1"},
|
||||
),
|
||||
"validation_reverse": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, training_data_url, rank, number_of_steps, number_of_frames, frame_rate,
|
||||
resolution, aspect_ratio, learning_rate, trigger_phrase="", auto_scale_input=False,
|
||||
split_input_into_scenes=True, split_input_duration_threshold=30.0,
|
||||
validation_negative_prompt="blurry, low quality, bad quality, out of focus",
|
||||
validation_number_of_frames=81, validation_resolution="high",
|
||||
validation_aspect_ratio="1:1", validation_reverse=False):
|
||||
def train_lora(
|
||||
self,
|
||||
training_data_url,
|
||||
rank,
|
||||
number_of_steps,
|
||||
number_of_frames,
|
||||
frame_rate,
|
||||
resolution,
|
||||
aspect_ratio,
|
||||
learning_rate,
|
||||
trigger_phrase="",
|
||||
auto_scale_input=False,
|
||||
split_input_into_scenes=True,
|
||||
split_input_duration_threshold=30.0,
|
||||
validation_negative_prompt="blurry, low quality, bad quality, out of focus",
|
||||
validation_number_of_frames=81,
|
||||
validation_resolution="high",
|
||||
validation_aspect_ratio="1:1",
|
||||
validation_reverse=False,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
return ("Error: No training data URL provided.",)
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"ltx-video-trainer", "No training data URL provided"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
@@ -285,23 +354,22 @@ class LtxVideoTrainerNode:
|
||||
"validation_number_of_frames": validation_number_of_frames,
|
||||
"validation_resolution": validation_resolution,
|
||||
"validation_aspect_ratio": validation_aspect_ratio,
|
||||
"validation_reverse": validation_reverse
|
||||
"validation_reverse": validation_reverse,
|
||||
}
|
||||
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/ltx-video-trainer", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/ltx-video-trainer", arguments
|
||||
)
|
||||
lora_url = result["lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.",)
|
||||
return ApiHandler.handle_text_generation_error("ltx-video-trainer", str(e))
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -317,4 +385,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HunyuanVideoLoraTrainer_fal": "Hunyuan Video LoRA Trainer (fal)",
|
||||
"WanLoraTrainer_fal": "WAN LoRA Trainer (fal)",
|
||||
"LtxVideoTrainer_fal": "LTX Video LoRA Trainer (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
+67
-125
@@ -1,75 +1,8 @@
|
||||
import os
|
||||
import configparser
|
||||
import tempfile
|
||||
import requests
|
||||
from PIL import Image
|
||||
import io
|
||||
import numpy as np
|
||||
import torch
|
||||
from fal_client.client import SyncClient
|
||||
from .fal_utils import ApiHandler, FalConfig, ImageUtils, ResultProcessor
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
def upload_image(image):
|
||||
try:
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0)
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1)
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0))
|
||||
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
image_url = fal_client.upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
class UpscalerNode:
|
||||
@classmethod
|
||||
@@ -77,74 +10,83 @@ class UpscalerNode:
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"upscale_factor": ("FLOAT", {"default": 2.0, "min": 1.0, "max": 4.0, "step": 0.5}),
|
||||
"negative_prompt": ("STRING", {"default": "(worst quality, low quality, normal quality:2)", "multiline": True}),
|
||||
"creativity": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 1.0, "step": 0.05}),
|
||||
"resemblance": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.05}),
|
||||
"guidance_scale": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.5}),
|
||||
"upscale_factor": (
|
||||
"FLOAT",
|
||||
{"default": 2.0, "min": 1.0, "max": 4.0, "step": 0.5},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "(worst quality, low quality, normal quality:2)",
|
||||
"multiline": True,
|
||||
},
|
||||
),
|
||||
"creativity": (
|
||||
"FLOAT",
|
||||
{"default": 0.35, "min": 0.0, "max": 1.0, "step": 0.05},
|
||||
),
|
||||
"resemblance": (
|
||||
"FLOAT",
|
||||
{"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.05},
|
||||
),
|
||||
"guidance_scale": (
|
||||
"FLOAT",
|
||||
{"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.5},
|
||||
),
|
||||
"num_inference_steps": ("INT", {"default": 18, "min": 1, "max": 100}),
|
||||
"enable_safety_checker": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"seed": ("INT", {"default": -1}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate_upscaled_image"
|
||||
CATEGORY = "FAL/Image"
|
||||
|
||||
def generate_upscaled_image(self, image, upscale_factor, negative_prompt, creativity, resemblance, guidance_scale, num_inference_steps, enable_safety_checker, seed=-1):
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
print("Failed to upload image for upscaling.")
|
||||
return self.create_blank_image()
|
||||
|
||||
arguments = {
|
||||
"image_url": image_url,
|
||||
"prompt": "masterpiece, best quality, highres",
|
||||
"upscale_factor": upscale_factor,
|
||||
"negative_prompt": negative_prompt,
|
||||
"creativity": creativity,
|
||||
"resemblance": resemblance,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"enable_safety_checker": enable_safety_checker
|
||||
}
|
||||
|
||||
if seed != -1:
|
||||
arguments["seed"] = seed
|
||||
|
||||
def generate_upscaled_image(
|
||||
self,
|
||||
image,
|
||||
upscale_factor,
|
||||
negative_prompt,
|
||||
creativity,
|
||||
resemblance,
|
||||
guidance_scale,
|
||||
num_inference_steps,
|
||||
enable_safety_checker,
|
||||
seed=-1,
|
||||
):
|
||||
try:
|
||||
handler = fal_client.submit("fal-ai/clarity-upscaler", arguments=arguments)
|
||||
result = handler.get()
|
||||
return self.process_result(result)
|
||||
# Upload the image using ImageUtils
|
||||
image_url = ImageUtils.upload_image(image)
|
||||
if not image_url:
|
||||
return ApiHandler.handle_image_generation_error(
|
||||
"clarity-upscaler", "Failed to upload image for upscaling"
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"image_url": image_url,
|
||||
"prompt": "masterpiece, best quality, highres",
|
||||
"upscale_factor": upscale_factor,
|
||||
"negative_prompt": negative_prompt,
|
||||
"creativity": creativity,
|
||||
"resemblance": resemblance,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"enable_safety_checker": enable_safety_checker,
|
||||
}
|
||||
|
||||
if seed != -1:
|
||||
arguments["seed"] = seed
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/clarity-upscaler", arguments
|
||||
)
|
||||
return ResultProcessor.process_image_result(result)
|
||||
except Exception as e:
|
||||
print(f"Error generating upscaled image: {str(e)}")
|
||||
return self.create_blank_image()
|
||||
return ApiHandler.handle_image_generation_error("clarity-upscaler", str(e))
|
||||
|
||||
def process_result(self, result):
|
||||
try:
|
||||
img_url = result["image"]["url"]
|
||||
img_response = requests.get(img_url)
|
||||
img = Image.open(io.BytesIO(img_response.content))
|
||||
img_array = np.array(img).astype(np.float32) / 255.0
|
||||
|
||||
# Stack the images along a new first dimension
|
||||
stacked_images = np.stack([img_array], axis=0)
|
||||
|
||||
# Convert to PyTorch tensor
|
||||
img_tensor = torch.from_numpy(stacked_images)
|
||||
return (img_tensor,)
|
||||
except Exception as e:
|
||||
print(f"Error processing result: {str(e)}")
|
||||
return self.create_blank_image()
|
||||
|
||||
def create_blank_image(self):
|
||||
blank_img = Image.new('RGB', (512, 512), color='black')
|
||||
img_array = np.array(blank_img).astype(np.float32) / 255.0
|
||||
img_tensor = torch.from_numpy(img_array)[None,]
|
||||
return (img_tensor,)
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -154,4 +96,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Upscaler_fal": "Clarity Upscaler (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
+796
-493
File diff suppressed because it is too large
Load Diff
+26
-76
@@ -1,41 +1,8 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
import torch
|
||||
from PIL import Image
|
||||
import tempfile
|
||||
import numpy as np
|
||||
from .fal_utils import ApiHandler, FalConfig, ImageUtils
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
class VLMNode:
|
||||
@classmethod
|
||||
@@ -43,9 +10,17 @@ class VLMNode:
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"model": (["google/gemini-flash-1.5-8b", "anthropic/claude-3.5-sonnet", "anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5", "google/gemini-flash-1.5", "openai/gpt-4o"],
|
||||
{"default": "google/gemini-flash-1.5-8b"}),
|
||||
"model": (
|
||||
[
|
||||
"google/gemini-flash-1.5-8b",
|
||||
"anthropic/claude-3.5-sonnet",
|
||||
"anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5",
|
||||
"google/gemini-flash-1.5",
|
||||
"openai/gpt-4o",
|
||||
],
|
||||
{"default": "google/gemini-flash-1.5-8b"},
|
||||
),
|
||||
"system_prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
@@ -57,34 +32,12 @@ class VLMNode:
|
||||
|
||||
def generate_text(self, prompt, model, system_prompt, image):
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
image_url = fal_client.upload_file(temp_file_path)
|
||||
# Upload the image using ImageUtils
|
||||
image_url = ImageUtils.upload_image(image)
|
||||
if not image_url:
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
model, "Failed to upload image"
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"model": model,
|
||||
@@ -93,16 +46,13 @@ class VLMNode:
|
||||
"image_url": image_url,
|
||||
}
|
||||
|
||||
handler = fal_client.submit("fal-ai/any-llm/vision", arguments=arguments)
|
||||
result = handler.get()
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/any-llm/vision", arguments
|
||||
)
|
||||
return (result["output"],)
|
||||
except Exception as e:
|
||||
print(f"Error generating text with VLM: {str(e)}")
|
||||
return ("Error: Unable to generate text.",)
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
return ApiHandler.handle_text_generation_error(model, str(e))
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -112,4 +62,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLM_fal": "VLM (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "fal-api"
|
||||
description = "Custom nodes for using fal API. Video generation with Kling, Runway, Luma. Image generation with Flux. LLMs and VLMs OpenAI, Claude, Llama and Gemini."
|
||||
version = "1.0.1"
|
||||
version = "1.0.5"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["fal-client", "torch"]
|
||||
|
||||
|
||||
+2
-1
@@ -1,2 +1,3 @@
|
||||
fal-client
|
||||
torch
|
||||
torch
|
||||
opencv-python
|
||||
|
||||
Reference in New Issue
Block a user