Compare commits

...
Author SHA1 Message Date
gökay aydoğan f9b21a5e93 change optional image number 2025-09-11 16:12:08 +03:00
Gökay Aydoğan fbee93b5b5 Update pyproject.toml 2025-09-11 15:55:31 +03:00
Gökay Aydoğan 845b9d46c5 Update requirements.txt 2025-09-05 11:34:43 +03:00
Gökay Aydoğan 31572e6e45 Merge pull request #39 from PierrunoYT/feature/add-qwen-image-edit
feat: add Qwen Image Edit node with parallel CFG support
2025-09-05 11:32:45 +03:00
Gökay Aydoğan b60c18d8a8 Merge pull request #40 from gokayfem/nano-banana
nano banana edit
2025-09-05 11:32:19 +03:00
gokayfem 58c54acbce nano banana edit 2025-09-05 11:31:40 +03:00
PierrunoYTandClaude 27580456ed feat: add Qwen Image Edit node with parallel CFG support
- Added QwenImageEdit class to nodes/image_node.py
- Supports image editing with text prompts using fal-ai/qwen-image-edit endpoint
- Features flexible image sizing, inference control, and acceleration options
- Includes safety checker, output format selection, and seed control
- Added to NODE_CLASS_MAPPINGS and NODE_DISPLAY_NAME_MAPPINGS
- Follows existing code patterns for consistency and error handling

🤖 Generated with [Claude Code](https://claude.ai/code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-08-28 17:30:08 +02:00
Gökay Aydoğan a68c56134c Merge pull request #33 from jimlee2048/main
feat: support SeedEdit 3.0
2025-07-20 18:40:57 +03:00
Jim Lee cd9eb99568 feat: support SeedEdit 3.0 2025-07-20 22:55:01 +08:00
Gökay Aydoğan ef774a511b Update pyproject.toml 2025-06-20 13:38:43 +03:00
Gökay Aydoğan 66d4dcf54d Merge pull request #28 from KhDu/feature/adding_new_nodes
Added new fal model nodes (Veo3 +Seedance + Imagen4)
2025-06-20 13:38:28 +03:00
Gökay Aydoğan a6d061c0eb Update image_node.py 2025-06-20 13:37:07 +03:00
KhDu 34d3a8396e Reverted Flux Kontext Max to being a boolean switch 2025-06-19 21:04:41 +03:00
KhDu 06a30a6f21 Added Veo3 model 2025-06-19 18:29:55 +03:00
KhDu f4f486edb0 Added Kontext Multi, seperated Kontext Max into its own node.
Added Seedance Video model.

Added Imagen4 Image model.
2025-06-19 18:21:20 +03:00
KhDu 1f6f476679 added imagen4 text-to-image, and seedance image-to-video 2025-06-16 21:21:37 +03:00
Gökay Aydoğan 1e561ac944 Update pyproject.toml 2025-06-02 20:39:18 +03:00
Gökay Aydoğan cf523888a7 Merge pull request #26 from gokayfem/big-refactor-cleaning
fix: big refactor and cleaning
2025-06-02 20:38:59 +03:00
gokayfem 93aa2cbc04 fix: big refactor and cleaning 2025-06-02 19:03:45 +03:00
Gökay Aydoğan 5be02175f3 Update pyproject.toml 2025-06-01 14:35:30 +03:00
gokayfem 4ff17aa6ef Merge branch 'main' of https://github.com/gokayfem/ComfyUI-FLUX-fal-API 2025-06-01 14:23:34 +03:00
gokayfem a6d29a2d4c readme 2025-06-01 14:23:13 +03:00
Gökay Aydoğan 4215edebf0 Merge pull request #23 from gokayfem/api-key-setup
feat: api key setup
2025-06-01 14:15:50 +03:00
11 changed files with 2328 additions and 1235 deletions
+5
View File
@@ -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
View File
@@ -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"]
+220
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+28 -47
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+26 -76
View File
@@ -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
View File
@@ -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
View File
@@ -1,2 +1,3 @@
fal-client
torch
torch
opencv-python