Add files via upload

This commit is contained in:
pvlprk
2025-10-07 12:51:52 +02:00
committed by GitHub
parent 7ab3d7f8a4
commit d4462d6022
7 changed files with 1991 additions and 1119 deletions
+3
View File
@@ -38,6 +38,7 @@ from .pvl_any2string import PVL_Any2String
from .pvl_fal_remove_bg_v2 import PVL_fal_RemoveBackground_API
from .pvl_fal_depth_anything_v2 import PVL_fal_DepthAnythingV2_API
from .pvl_google_nano_banana_mandatory_img import PVL_Google_NanoBanana_API_mandatory_IMG
from .pvl_fal_flux_with_lora_pulid import PVL_fal_FluxWithLoraPulID_API
NODE_CLASS_MAPPINGS = {
"PVL Call OpenAI Assistant": CallAssistantNode,
@@ -80,6 +81,7 @@ NODE_CLASS_MAPPINGS = {
"PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API,
"PVL_fal_DepthAnythingV2_API": PVL_fal_DepthAnythingV2_API,
"PVL_Google_NanoBanana_API_mandatory_IMG": PVL_Google_NanoBanana_API_mandatory_IMG,
"PVL_fal_FluxWithLoraPulID_API": PVL_fal_FluxWithLoraPulID_API,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -122,4 +124,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai)",
"PVL_fal_DepthAnythingV2_API": "PVL Depth Anything V2 (fal.ai)",
"PVL_Google_NanoBanana_API_mandatory_IMG": "PVL Google Nano-Banana API mandatory IMG",
"PVL_fal_FluxWithLoraPulID_API": "PVL Flux Lora PulID (fal.ai)",
}
+255 -47
View File
@@ -1,8 +1,14 @@
import os
import re
import torch
import time
import requests
import io
from concurrent.futures import ThreadPoolExecutor, as_completed
from .fal_utils import FalConfig, ImageUtils, ResultProcessor, ApiHandler
class PVL_fal_FluxWithLora_API:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -19,77 +25,279 @@ class PVL_fal_FluxWithLora_API:
"sync_mode": ("BOOLEAN", {"default": False}),
},
"optional": {
"delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Delimiter for splitting prompts (e.g., [*], \\n, |)"}),
"lora1_name": ("STRING", {"default": ""}),
"lora1_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}),
"lora2_name": ("STRING", {"default": ""}),
"lora2_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}),
"lora3_name": ("STRING", {"default": ""}),
"lora3_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}),
}
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "generate_image"
CATEGORY = "PVL_tools_FAL"
def generate_image(self, prompt, width, height, steps, CFG, seed,
def _build_call_prompts(self, base_prompts, num_images):
"""
Maps prompts to calls according to the rule:
- If len(prompts) >= num_images → take first num_images
- If len(prompts) < num_images → Use each prompt in order. For remaining calls, reuse the last prompt.
"""
N = max(1, int(num_images))
if not base_prompts:
return []
if len(base_prompts) >= N:
call_prompts = base_prompts[:N]
else:
print(f"[PVL WARNING] Provided {len(base_prompts)} prompts but num_images={N}. "
f"Reusing the last prompt for remaining calls.")
call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts))
return call_prompts
# -------- FAL Queue API - TWO PHASE EXECUTION --------
def _fal_submit_only(self, prompt_text, width, height, steps, CFG, seed,
enable_safety_checker, output_format, sync_mode,
lora1_name, lora1_scale, lora2_name, lora2_scale,
lora3_name, lora3_scale):
"""
Phase 1: Submit request to FAL and return request info immediately.
Does NOT wait for completion.
"""
arguments = {
"prompt": prompt_text,
"num_inference_steps": steps,
"guidance_scale": CFG,
"num_images": 1, # Each call generates 1 image
"enable_safety_checker": enable_safety_checker,
"output_format": output_format,
"sync_mode": sync_mode,
"image_size": {
"width": width,
"height": height
}
}
if seed != -1:
arguments["seed"] = seed
# Handle LoRAs
loras = []
if lora1_name.strip():
loras.append({"path": lora1_name.strip(), "scale": lora1_scale})
if lora2_name.strip():
loras.append({"path": lora2_name.strip(), "scale": lora2_scale})
if lora3_name.strip():
loras.append({"path": lora3_name.strip(), "scale": lora3_scale})
if loras:
arguments["loras"] = loras
# Check if ApiHandler supports async submission
if hasattr(ApiHandler, 'submit_only'):
return ApiHandler.submit_only("fal-ai/flux-lora", arguments)
else:
# Fallback to direct FAL queue API
return self._direct_fal_submit("fal-ai/flux-lora", arguments)
def _direct_fal_submit(self, endpoint, arguments):
"""Direct FAL queue API submission when ApiHandler doesn't support async."""
fal_key = os.getenv("FAL_KEY", "")
if not fal_key:
raise RuntimeError("FAL_KEY environment variable not set")
base = "https://queue.fal.run"
submit_url = f"{base}/{endpoint}"
headers = {"Authorization": f"Key {fal_key}"}
r = requests.post(submit_url, headers=headers, json=arguments, timeout=120)
if not r.ok:
raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}")
sub = r.json()
req_id = sub.get("request_id")
if not req_id:
raise RuntimeError("FAL did not return a request_id")
status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
}
def _fal_poll_and_fetch(self, request_info, timeout=120):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns image tensor.
"""
# Check if ApiHandler supports async polling
if hasattr(ApiHandler, 'poll_and_get_result'):
result = ApiHandler.poll_and_get_result(request_info, timeout)
else:
# Fallback to direct polling
fal_key = os.getenv("FAL_KEY", "")
headers = {"Authorization": f"Key {fal_key}"}
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
# Poll for completion
deadline = time.time() + timeout
completed = False
while time.time() < deadline:
try:
sr = requests.get(status_url, headers=headers, timeout=10)
if sr.ok and sr.json().get("status") == "COMPLETED":
completed = True
break
except Exception:
pass
time.sleep(0.6)
if not completed:
raise RuntimeError(f"FAL request timed out after {timeout}s")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
result = rr.json().get("response", rr.json())
# Process result using ResultProcessor
return ResultProcessor.process_image_result(result)
def generate_image(self, prompt, width, height, steps, CFG, seed,
num_images, enable_safety_checker, output_format, sync_mode,
delimiter="[*]",
lora1_name="", lora1_scale=1.0,
lora2_name="", lora2_scale=1.0,
lora3_name="", lora3_scale=1.0):
_t0 = time.time()
try:
# Prepare the arguments for the API call
arguments = {
"prompt": prompt,
"num_inference_steps": steps, # Using the renamed parameter
"guidance_scale": CFG, # Using the renamed parameter
"num_images": num_images,
"enable_safety_checker": enable_safety_checker,
"output_format": output_format,
"sync_mode": sync_mode,
"image_size": { # Always use custom dimensions now
"width": width,
"height": height
# Split prompts using delimiter with regex support
try:
base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()]
except re.error:
print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.")
base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()]
if not base_prompts:
raise RuntimeError("No valid prompts provided.")
# Map prompts to num_images calls
call_prompts = self._build_call_prompts(base_prompts, num_images)
print(f"[PVL INFO] Processing {len(call_prompts)} prompts")
# Single call: process directly (less overhead)
if len(call_prompts) == 1:
req_info = self._fal_submit_only(
call_prompts[0], width, height, steps, CFG, seed,
enable_safety_checker, output_format, sync_mode,
lora1_name, lora1_scale, lora2_name, lora2_scale,
lora3_name, lora3_scale
)
result = self._fal_poll_and_fetch(req_info)
img_tensor = result[0] if isinstance(result, tuple) else result
if img_tensor.ndim == 3:
img_tensor = img_tensor.unsqueeze(0)
_t1 = time.time()
print(f"[PVL INFO] Successfully generated 1 image in {(_t1 - _t0):.2f}s")
return (img_tensor,)
# Multiple calls: TRUE PARALLEL execution
print(f"[PVL INFO] Submitting {len(call_prompts)} requests in parallel...")
# PHASE 1: Submit all requests in parallel
submit_results = []
max_workers = min(len(call_prompts), 6) # Increased from 4 to 6
with ThreadPoolExecutor(max_workers=max_workers) as executor:
submit_futs = {
executor.submit(
self._fal_submit_only,
call_prompts[i], width, height, steps, CFG, seed,
enable_safety_checker, output_format, sync_mode,
lora1_name, lora1_scale, lora2_name, lora2_scale,
lora3_name, lora3_scale
): i
for i in range(len(call_prompts))
}
}
for fut in as_completed(submit_futs):
idx = submit_futs[fut]
try:
req_info = fut.result()
submit_results.append((idx, req_info))
except Exception as e:
print(f"[PVL ERROR] Submit failed for prompt {idx}: {e}")
# Add seed if provided (not -1)
if seed != -1:
arguments["seed"] = seed
if not submit_results:
raise RuntimeError("All FAL submission requests failed")
# Handle LoRAs
loras = []
print(f"[PVL INFO] {len(submit_results)} requests submitted. Polling for results...")
# Add LoRA 1 if provided
if lora1_name.strip():
loras.append({
"path": lora1_name.strip(),
"scale": lora1_scale
})
# PHASE 2: Poll all requests in parallel
results = {}
failed_count = 0
# Add LoRA 2 if provided
if lora2_name.strip():
loras.append({
"path": lora2_name.strip(),
"scale": lora2_scale
})
with ThreadPoolExecutor(max_workers=max_workers) as executor:
poll_futs = {
executor.submit(self._fal_poll_and_fetch, req_info): idx
for idx, req_info in submit_results
}
for fut in as_completed(poll_futs):
idx = poll_futs[fut]
try:
result = fut.result()
results[idx] = result
except Exception as e:
failed_count += 1
print(f"[PVL ERROR] Poll failed for prompt {idx}: {e}")
# Add LoRA 3 if provided
if lora3_name.strip():
loras.append({
"path": lora3_name.strip(),
"scale": lora3_scale
})
if not results:
raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)")
if loras:
arguments["loras"] = loras
if failed_count > 0:
print(f"[PVL WARNING] {failed_count}/{len(call_prompts)} requests failed, continuing with {len(results)} successful results")
# Submit the request and get the result
result = ApiHandler.submit_and_get_result("fal-ai/flux-lora", arguments)
# Combine all image tensors in order
all_images = []
for i in range(len(call_prompts)):
if i in results:
result = results[i]
img_tensor = result[0] if isinstance(result, tuple) else result
if torch.is_tensor(img_tensor):
# Handle both 3D (H,W,C) and 4D (B,H,W,C) tensors
if img_tensor.ndim == 3:
img_tensor = img_tensor.unsqueeze(0)
all_images.append(img_tensor)
# Process the result and return the image tensor
return ResultProcessor.process_image_result(result)
if not all_images:
raise RuntimeError("No images were generated from API calls")
# Stack all images into single batch
final_tensor = torch.cat(all_images, dim=0)
_t1 = time.time()
print(f"[PVL INFO] Successfully generated {final_tensor.shape[0]} images in {(_t1 - _t0):.2f}s")
return (final_tensor,)
except Exception as e:
print(f"Error generating image with FLUX: {str(e)}")
return ApiHandler.handle_image_generation_error("FLUX", e)
return ApiHandler.handle_image_generation_error("FLUX", e)
NODE_CLASS_MAPPINGS = {"PVL_fal_FluxWithLora_API": PVL_fal_FluxWithLora_API}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_fal_FluxWithLora_API": "PVL FAL Flux with LoRA"}
+314
View File
@@ -0,0 +1,314 @@
import os
import re
import torch
import time
import requests
import io
import base64
from concurrent.futures import ThreadPoolExecutor, as_completed
from .fal_utils import FalConfig, ImageUtils, ResultProcessor, ApiHandler
class PVL_fal_FluxWithLoraPulID_API:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "a woman holding sign with glowing green text 'PuLID for FLUX'"}),
"reference_image": ("IMAGE",),
"num_images": ("INT", {"default": 1, "min": 1, "max": 4}),
"num_inference_steps": ("INT", {"default": 20, "min": 1, "max": 100}),
"guidance_scale": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.1}),
"seed": ("INT", {"default": -1, "min": -1, "max": 4294967295}),
"enable_safety_checker": ("BOOLEAN", {"default": True}),
"sync_mode": ("BOOLEAN", {"default": False}),
},
"optional": {
"delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Delimiter for splitting prompts (e.g., [*], \\n, |)"}),
"image_size": (["square_hd", "square", "portrait_4_3", "portrait_16_9", "landscape_4_3", "landscape_16_9"], {"default": "landscape_4_3"}),
"custom_width": ("INT", {"default": 0, "min": 0, "max": 2048, "step": 64}),
"custom_height": ("INT", {"default": 0, "min": 0, "max": 2048, "step": 64}),
"negative_prompt": ("STRING", {"multiline": True, "default": "bad quality, worst quality, text, signature, watermark, extra limbs"}),
"lora_path": ("STRING", {"default": "", "placeholder": "Optional LoRA path"}),
"lora_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1}),
"start_step": ("INT", {"default": 0, "min": 0, "max": 100}),
"true_cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1}),
"id_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}),
"max_sequence_length": (["128", "256", "512"], {"default": "128"}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "generate_image"
CATEGORY = "PVL_tools_FAL"
def _build_call_prompts(self, base_prompts, num_images):
"""
Maps prompts to calls according to the rule:
- If len(prompts) >= num_images → take first num_images
- If len(prompts) < num_images → Use each prompt in order. For remaining calls, reuse the last prompt.
"""
N = max(1, int(num_images))
if not base_prompts:
return []
if len(base_prompts) >= N:
call_prompts = base_prompts[:N]
else:
print(f"[PVL WARNING] Provided {len(base_prompts)} prompts but num_images={N}. "
f"Reusing the last prompt for remaining calls.")
call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts))
return call_prompts
# -------- FAL Queue API - TWO PHASE EXECUTION --------
def _fal_submit_only(self, prompt_text, reference_image_url, image_size, custom_width, custom_height,
num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode,
enable_safety_checker, lora_path, lora_strength, start_step, true_cfg,
id_weight, max_sequence_length):
"""
Phase 1: Submit request to FAL and return request info immediately.
Does NOT wait for completion.
"""
# Build image_size argument
if custom_width > 0 and custom_height > 0:
image_size_arg = {"width": custom_width, "height": custom_height}
else:
image_size_arg = image_size
arguments = {
"prompt": prompt_text,
"reference_image_url": reference_image_url,
"image_size": image_size_arg,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"negative_prompt": negative_prompt,
"sync_mode": sync_mode,
"enable_safety_checker": enable_safety_checker,
"true_cfg": true_cfg,
"id_weight": id_weight,
"max_sequence_length": max_sequence_length,
}
if seed != -1:
arguments["seed"] = seed
if lora_path.strip():
arguments["lora_path"] = lora_path.strip()
arguments["lora_strength"] = lora_strength
if start_step > 0:
arguments["start_step"] = start_step
# Check if ApiHandler supports async submission
if hasattr(ApiHandler, 'submit_only'):
return ApiHandler.submit_only("fal-ai/flux-pulid-lora", arguments)
else:
# Fallback to direct FAL queue API
return self._direct_fal_submit("fal-ai/flux-pulid-lora", arguments)
def _direct_fal_submit(self, endpoint, arguments):
"""Direct FAL queue API submission when ApiHandler doesn't support async."""
fal_key = os.getenv("FAL_KEY", "")
if not fal_key:
raise RuntimeError("FAL_KEY environment variable not set")
base = "https://queue.fal.run"
submit_url = f"{base}/{endpoint}"
headers = {"Authorization": f"Key {fal_key}"}
r = requests.post(submit_url, headers=headers, json=arguments, timeout=120)
if not r.ok:
raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}")
sub = r.json()
req_id = sub.get("request_id")
if not req_id:
raise RuntimeError("FAL did not return a request_id")
status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
}
def _fal_poll_and_fetch(self, request_info, timeout=120):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns image tensor.
"""
# Check if ApiHandler supports async polling
if hasattr(ApiHandler, 'poll_and_get_result'):
result = ApiHandler.poll_and_get_result(request_info, timeout)
else:
# Fallback to direct polling
fal_key = os.getenv("FAL_KEY", "")
headers = {"Authorization": f"Key {fal_key}"}
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
# Poll for completion
deadline = time.time() + timeout
completed = False
while time.time() < deadline:
try:
sr = requests.get(status_url, headers=headers, timeout=10)
if sr.ok and sr.json().get("status") == "COMPLETED":
completed = True
break
except Exception:
pass
time.sleep(0.6)
if not completed:
raise RuntimeError(f"FAL request timed out after {timeout}s")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
result = rr.json().get("response", rr.json())
# Process result using ResultProcessor
return ResultProcessor.process_image_result(result)
def generate_image(self, prompt, reference_image, num_images, num_inference_steps,
guidance_scale, seed, enable_safety_checker, sync_mode,
delimiter="[*]",
image_size="landscape_4_3", custom_width=0, custom_height=0,
negative_prompt="bad quality, worst quality, text, signature, watermark, extra limbs",
lora_path="", lora_strength=1.0, start_step=0, true_cfg=1.0,
id_weight=1.0, max_sequence_length="128"):
_t0 = time.time()
try:
# Upload reference image once (shared for all calls)
print("[PVL INFO] Uploading reference image to fal.ai storage...")
reference_image_url = ImageUtils.upload_image(reference_image)
print(f"[PVL INFO] Reference image uploaded: {reference_image_url}")
# Split prompts using delimiter with regex support
try:
base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()]
except re.error:
print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.")
base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()]
if not base_prompts:
raise RuntimeError("No valid prompts provided.")
# Map prompts to num_images calls
call_prompts = self._build_call_prompts(base_prompts, num_images)
print(f"[PVL INFO] Processing {len(call_prompts)} prompts")
# Single call: process directly (less overhead)
if len(call_prompts) == 1:
req_info = self._fal_submit_only(
call_prompts[0], reference_image_url, image_size, custom_width, custom_height,
num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode,
enable_safety_checker, lora_path, lora_strength, start_step, true_cfg,
id_weight, max_sequence_length
)
result = self._fal_poll_and_fetch(req_info)
img_tensor = result[0] if isinstance(result, tuple) else result
if img_tensor.ndim == 3:
img_tensor = img_tensor.unsqueeze(0)
_t1 = time.time()
print(f"[PVL INFO] Successfully generated 1 image in {(_t1 - _t0):.2f}s")
return (img_tensor,)
# Multiple calls: TRUE PARALLEL execution
print(f"[PVL INFO] Submitting {len(call_prompts)} requests in parallel...")
# PHASE 1: Submit all requests in parallel
submit_results = []
max_workers = min(len(call_prompts), 6)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
submit_futs = {
executor.submit(
self._fal_submit_only,
call_prompts[i], reference_image_url, image_size, custom_width, custom_height,
num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode,
enable_safety_checker, lora_path, lora_strength, start_step, true_cfg,
id_weight, max_sequence_length
): i
for i in range(len(call_prompts))
}
for fut in as_completed(submit_futs):
idx = submit_futs[fut]
try:
req_info = fut.result()
submit_results.append((idx, req_info))
except Exception as e:
print(f"[PVL ERROR] Submit failed for prompt {idx}: {e}")
if not submit_results:
raise RuntimeError("All FAL submission requests failed")
print(f"[PVL INFO] {len(submit_results)} requests submitted. Polling for results...")
# PHASE 2: Poll all requests in parallel
results = {}
failed_count = 0
with ThreadPoolExecutor(max_workers=max_workers) as executor:
poll_futs = {
executor.submit(self._fal_poll_and_fetch, req_info): idx
for idx, req_info in submit_results
}
for fut in as_completed(poll_futs):
idx = poll_futs[fut]
try:
result = fut.result()
results[idx] = result
except Exception as e:
failed_count += 1
print(f"[PVL ERROR] Poll failed for prompt {idx}: {e}")
if not results:
raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)")
if failed_count > 0:
print(f"[PVL WARNING] {failed_count}/{len(call_prompts)} requests failed, continuing with {len(results)} successful results")
# Combine all image tensors in order
all_images = []
for i in range(len(call_prompts)):
if i in results:
result = results[i]
img_tensor = result[0] if isinstance(result, tuple) else result
if torch.is_tensor(img_tensor):
# Handle both 3D (H,W,C) and 4D (B,H,W,C) tensors
if img_tensor.ndim == 3:
img_tensor = img_tensor.unsqueeze(0)
all_images.append(img_tensor)
if not all_images:
raise RuntimeError("No images were generated from API calls")
# Stack all images into single batch
final_tensor = torch.cat(all_images, dim=0)
_t1 = time.time()
print(f"[PVL INFO] Successfully generated {final_tensor.shape[0]} images in {(_t1 - _t0):.2f}s")
return (final_tensor,)
except Exception as e:
print(f"Error generating image with FLUX PuLID LoRA: {str(e)}")
# Fallback error handling - return empty tensor
empty_tensor = torch.zeros((1, 64, 64, 3), dtype=torch.float32)
return (empty_tensor,)
NODE_CLASS_MAPPINGS = {"PVL_fal_FluxWithLoraPulID_API": PVL_fal_FluxWithLoraPulID_API}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_fal_FluxWithLoraPulID_API": "PVL Flux Lora PulID (fal.ai)"}
+376 -168
View File
@@ -1,168 +1,376 @@
import torch
import numpy as np
from .fal_utils import ImageUtils, ApiHandler
class PVL_fal_RemoveBackground_API:
"""
ComfyUI node for FAL 'fal-ai/birefnet/v2' — Remove Background V2.
Inputs:
- image (IMAGE): One input image (batch NOT supported).
- model (CHOICE): Which model variant to use.
- operating_resolution (CHOICE): Resolution for inference ("1024x1024" or "2048x2048").
- output_format (CHOICE): "png" or "webp".
- output_mask (BOOLEAN): Whether to also return the mask.
- refine_foreground (BOOLEAN): Whether to refine the foreground (default: True).
Outputs:
- IMAGE: Foreground with background removed (RGB or RGBA if provided).
- MASK: Optional mask (1-channel, float, shape [1,H,W]) if output_mask=True.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"model": (
[
"General Use (Light)",
"General Use (Light 2K)",
"General Use (Heavy)",
"Matting",
"Portrait",
],
{"default": "General Use (Light)"},
),
"operating_resolution": (
["1024x1024", "2048x2048"],
{"default": "1024x1024"},
),
"output_format": (
["png", "webp"],
{"default": "png"},
),
"output_mask": ("BOOLEAN", {"default": False}),
"refine_foreground": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("foreground", "mask")
FUNCTION = "remove_background"
CATEGORY = "PVL_tools"
# ------------------------- helpers -------------------------
def _raise(self, msg: str):
raise RuntimeError(msg)
def _split_image_batch(self, image):
import numpy as np
if isinstance(image, torch.Tensor):
if image.ndim == 4:
return [image[i] for i in range(image.shape[0])]
elif image.ndim == 3:
return [image]
else:
self._raise("FAL: unsupported image tensor dimensionality.")
elif isinstance(image, np.ndarray):
t = torch.from_numpy(image)
return self._split_image_batch(t)
else:
self._raise("FAL: unsupported image type (expected torch Tensor or numpy.ndarray).")
def _download_pil(self, url, mode=None):
import requests, io
from PIL import Image
resp = requests.get(url, timeout=120)
resp.raise_for_status()
pil = Image.open(io.BytesIO(resp.content))
if mode is not None:
pil = pil.convert(mode)
return pil
# ------------------------- main -------------------------
def remove_background(
self,
image,
model,
operating_resolution,
output_format,
output_mask,
refine_foreground,
):
# Split batch into frames
frames = self._split_image_batch(image)
if not frames:
self._raise("FAL: no input image frames provided.")
if len(frames) > 1:
self._raise("FAL: batch >1 not supported for Remove Background API.")
# Inline image as base64 data URI (avoids storage upload)
image_url = ImageUtils.image_to_data_uri(frames[0])
if not image_url:
self._raise("FAL: failed to convert input image.")
arguments = {
"image_url": image_url,
"model": model,
"operating_resolution": operating_resolution,
"output_format": output_format,
"output_mask": bool(output_mask),
"refine_foreground": bool(refine_foreground),
}
# Submit request
result = ApiHandler.submit_and_get_result("fal-ai/birefnet/v2", arguments)
if not isinstance(result, dict):
self._raise("FAL: unexpected response type (expected dict).")
# ---- Foreground image (keep alpha if present) ----
if "image" not in result or not isinstance(result["image"], dict):
self._raise("FAL: response missing foreground image.")
fg_url = result["image"].get("url")
if not fg_url:
self._raise("FAL: foreground image has no URL.")
pil_fg = self._download_pil(fg_url, mode=None) # keep native mode
fg_arr = np.array(pil_fg).astype(np.float32) / 255.0 # (H,W,3) or (H,W,4)
if fg_arr.ndim == 2: # grayscale fallback
fg_arr = np.expand_dims(fg_arr, axis=-1)
fg_tensor = torch.from_numpy(fg_arr).unsqueeze(0) # (1,H,W,C)
# ---- Mask (shape [1,H,W], float 0..1) ----
mask_tensor = torch.zeros((1, fg_arr.shape[0], fg_arr.shape[1]), dtype=torch.float32)
if bool(output_mask):
mask_url = None
if isinstance(result.get("mask_image"), dict):
mask_url = result["mask_image"].get("url")
if mask_url:
pil_mask_rgba = self._download_pil(mask_url, mode=None) # keep native
if "A" in pil_mask_rgba.getbands():
alpha = pil_mask_rgba.getchannel("A")
mask_arr = np.array(alpha).astype(np.float32) / 255.0 # (H,W)
else:
pil_mask_L = pil_mask_rgba.convert("L")
mask_arr = np.array(pil_mask_L).astype(np.float32) / 255.0
mask_tensor = torch.from_numpy(mask_arr).unsqueeze(0) # (1,H,W)
else:
# fallback: try to read alpha channel from foreground if present
if "A" in pil_fg.getbands():
alpha = pil_fg.getchannel("A")
mask_arr = np.array(alpha).astype(np.float32) / 255.0
mask_tensor = torch.from_numpy(mask_arr).unsqueeze(0)
return (fg_tensor, mask_tensor)
# ---- ComfyUI discovery ----
NODE_CLASS_MAPPINGS = {
"PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai)",
}
import torch
import numpy as np
import time
import requests
import io
from concurrent.futures import ThreadPoolExecutor, as_completed
from .fal_utils import ImageUtils, ApiHandler
class PVL_fal_RemoveBackground_API:
"""
ComfyUI node for FAL 'fal-ai/birefnet/v2' — Remove Background V2.
Features:
- TRUE PARALLEL execution: submits all requests first, then polls all in parallel
- Error handling for individual requests with partial results support
- Batch processing with worker cap to prevent thread explosion
- Optional sync_mode toggle
Inputs:
- image (IMAGE): Input image(s). Batch processing is supported with parallel API calls.
- model (CHOICE): Which model variant to use.
- operating_resolution (CHOICE): Resolution for inference ("1024x1024" or "2048x2048").
- output_format (CHOICE): "png" or "webp".
- output_mask (BOOLEAN): Whether to also return the mask.
- refine_foreground (BOOLEAN): Whether to refine the foreground (default: True).
- sync_mode (BOOLEAN): Use synchronous mode for FAL API (default: False).
Outputs:
- IMAGE: Foreground with background removed (RGB or RGBA if provided).
- MASK: Optional mask (1-channel, float, shape [B,H,W]) if output_mask=True.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"model": (
[
"General Use (Light)",
"General Use (Light 2K)",
"General Use (Heavy)",
"Matting",
"Portrait",
],
{"default": "General Use (Light)"},
),
"operating_resolution": (
["1024x1024", "2048x2048"],
{"default": "1024x1024"},
),
"output_format": (
["png", "webp"],
{"default": "png"},
),
"output_mask": ("BOOLEAN", {"default": False}),
"refine_foreground": ("BOOLEAN", {"default": True}),
"sync_mode": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("foreground", "mask")
FUNCTION = "remove_background"
CATEGORY = "PVL_tools"
# ------------------------- helpers -------------------------
def _raise(self, msg: str):
raise RuntimeError(msg)
def _split_image_batch(self, image):
if isinstance(image, torch.Tensor):
if image.ndim == 4:
return [image[i] for i in range(image.shape[0])]
elif image.ndim == 3:
return [image]
else:
self._raise("FAL: unsupported image tensor dimensionality.")
elif isinstance(image, np.ndarray):
t = torch.from_numpy(image)
return self._split_image_batch(t)
else:
self._raise("FAL: unsupported image type (expected torch Tensor or numpy.ndarray).")
def _download_pil(self, url, mode=None):
from PIL import Image
resp = requests.get(url, timeout=120)
resp.raise_for_status()
pil = Image.open(io.BytesIO(resp.content))
if mode is not None:
pil = pil.convert(mode)
return pil
# -------- FAL Queue API - TWO PHASE EXECUTION --------
def _fal_submit_only(self, frame, model, operating_resolution, output_format,
output_mask, refine_foreground, sync_mode):
"""
Phase 1: Submit request to FAL and return request info immediately.
Does NOT wait for completion.
"""
# Convert image to data URI
image_url = ImageUtils.image_to_data_uri(frame)
if not image_url:
raise RuntimeError("FAL: failed to convert input image.")
arguments = {
"image_url": image_url,
"model": model,
"operating_resolution": operating_resolution,
"output_format": output_format,
"output_mask": bool(output_mask),
"refine_foreground": bool(refine_foreground),
}
# Check if ApiHandler supports async submission
# If it has submit_only method, use it; otherwise use direct API call
if hasattr(ApiHandler, 'submit_only'):
return ApiHandler.submit_only("fal-ai/birefnet/v2", arguments, sync_mode)
else:
# Fallback to direct FAL queue API
return self._direct_fal_submit("fal-ai/birefnet/v2", arguments, sync_mode)
def _direct_fal_submit(self, endpoint, arguments, sync_mode):
"""Direct FAL queue API submission when ApiHandler doesn't support async."""
import os
fal_key = os.getenv("FAL_KEY", "")
if not fal_key:
raise RuntimeError("FAL_KEY environment variable not set")
base = "https://queue.fal.run"
submit_url = f"{base}/{endpoint}"
headers = {"Authorization": f"Key {fal_key}"}
payload = dict(arguments)
payload["sync_mode"] = sync_mode
r = requests.post(submit_url, headers=headers, json=payload, timeout=120)
if not r.ok:
raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}")
sub = r.json()
req_id = sub.get("request_id")
if not req_id:
raise RuntimeError("FAL did not return a request_id")
status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
}
def _fal_poll_and_fetch(self, request_info, timeout=120):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns (fg_tensor, mask_tensor).
"""
import os
fal_key = os.getenv("FAL_KEY", "")
headers = {"Authorization": f"Key {fal_key}"}
# Check if ApiHandler supports async polling
if hasattr(ApiHandler, 'poll_and_get_result'):
result = ApiHandler.poll_and_get_result(request_info, timeout)
else:
# Fallback to direct polling
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
# Poll for completion
deadline = time.time() + timeout
completed = False
while time.time() < deadline:
try:
sr = requests.get(status_url, headers=headers, timeout=10)
if sr.ok and sr.json().get("status") == "COMPLETED":
completed = True
break
except Exception:
pass
time.sleep(0.6)
if not completed:
raise RuntimeError(f"FAL request timed out after {timeout}s")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
rdata = rr.json()
result = rdata.get("response", rdata)
# Process result
return self._process_result(result)
def _process_result(self, result):
"""Process FAL API result and return tensors."""
if not isinstance(result, dict):
raise RuntimeError("FAL: unexpected response type (expected dict).")
# ---- Foreground image (keep alpha if present) ----
if "image" not in result or not isinstance(result["image"], dict):
raise RuntimeError("FAL: response missing foreground image.")
fg_url = result["image"].get("url")
if not fg_url:
raise RuntimeError("FAL: foreground image has no URL.")
pil_fg = self._download_pil(fg_url, mode=None) # keep native mode
fg_arr = np.array(pil_fg).astype(np.float32) / 255.0 # (H,W,3) or (H,W,4)
if fg_arr.ndim == 2: # grayscale fallback
fg_arr = np.expand_dims(fg_arr, axis=-1)
fg_tensor = torch.from_numpy(fg_arr) # (H,W,C)
# ---- Mask (shape [H,W], float 0..1) ----
mask_tensor = torch.zeros((fg_arr.shape[0], fg_arr.shape[1]), dtype=torch.float32)
# Try to get mask from response
mask_url = None
if isinstance(result.get("mask_image"), dict):
mask_url = result["mask_image"].get("url")
if mask_url:
pil_mask_rgba = self._download_pil(mask_url, mode=None)
if "A" in pil_mask_rgba.getbands():
alpha = pil_mask_rgba.getchannel("A")
mask_arr = np.array(alpha).astype(np.float32) / 255.0
else:
pil_mask_L = pil_mask_rgba.convert("L")
mask_arr = np.array(pil_mask_L).astype(np.float32) / 255.0
mask_tensor = torch.from_numpy(mask_arr)
else:
# Fallback: extract alpha channel from foreground if present
if "A" in pil_fg.getbands():
alpha = pil_fg.getchannel("A")
mask_arr = np.array(alpha).astype(np.float32) / 255.0
mask_tensor = torch.from_numpy(mask_arr)
return fg_tensor, mask_tensor
# ------------------------- main -------------------------
def remove_background(
self,
image,
model,
operating_resolution,
output_format,
output_mask,
refine_foreground,
sync_mode=False,
):
_t0 = time.time()
# Split batch into frames
frames = self._split_image_batch(image)
if not frames:
self._raise("FAL: no input image frames provided.")
batch_size = len(frames)
print(f"[FAL RemoveBG] Processing {batch_size} images...")
# Single image: process directly (less overhead)
if batch_size == 1:
try:
req_info = self._fal_submit_only(
frames[0], model, operating_resolution, output_format,
output_mask, refine_foreground, sync_mode
)
fg_tensor, mask_tensor = self._fal_poll_and_fetch(req_info)
_t1 = time.time()
print(f"[FAL RemoveBG] Completed in {(_t1 - _t0):.2f}s")
return (fg_tensor.unsqueeze(0), mask_tensor.unsqueeze(0))
except Exception as e:
self._raise(f"FAL: processing failed: {e}")
# Multiple images: TRUE PARALLEL execution
print(f"[FAL RemoveBG] Submitting {batch_size} requests in parallel...")
# PHASE 1: Submit all requests in parallel
submit_results = []
max_workers = min(batch_size, 6) # Cap at 6 workers
with ThreadPoolExecutor(max_workers=max_workers) as executor:
submit_futs = {
executor.submit(
self._fal_submit_only,
frames[i],
model,
operating_resolution,
output_format,
output_mask,
refine_foreground,
sync_mode
): i
for i in range(batch_size)
}
for fut in as_completed(submit_futs):
idx = submit_futs[fut]
try:
req_info = fut.result()
submit_results.append((idx, req_info))
except Exception as e:
print(f"[FAL RemoveBG] Submit failed for image {idx}: {e}")
if not submit_results:
self._raise("All FAL submission requests failed")
print(f"[FAL RemoveBG] {len(submit_results)} requests submitted. Polling for results...")
# PHASE 2: Poll all requests in parallel
results = {}
failed_count = 0
with ThreadPoolExecutor(max_workers=max_workers) as executor:
poll_futs = {
executor.submit(self._fal_poll_and_fetch, req_info): idx
for idx, req_info in submit_results
}
for fut in as_completed(poll_futs):
idx = poll_futs[fut]
try:
fg_tensor, mask_tensor = fut.result()
results[idx] = (fg_tensor, mask_tensor)
except Exception as e:
failed_count += 1
print(f"[FAL RemoveBG] Poll failed for image {idx}: {e}")
if not results:
self._raise(f"All FAL requests failed during polling ({failed_count} failures)")
if failed_count > 0:
print(f"[FAL RemoveBG WARNING] {failed_count}/{batch_size} requests failed, continuing with {len(results)} successful results")
# Stack results in original order (with None for failed images)
fg_list = []
mask_list = []
for i in range(batch_size):
if i in results:
fg_tensor, mask_tensor = results[i]
fg_list.append(fg_tensor)
mask_list.append(mask_tensor)
if not fg_list:
self._raise("No successful results to return")
# Stack into batched tensors
fg_batch = torch.stack(fg_list, dim=0) # (B,H,W,C)
mask_batch = torch.stack(mask_list, dim=0) # (B,H,W)
_t1 = time.time()
print(f"[FAL RemoveBG] Completed {len(fg_list)}/{batch_size} images in {(_t1 - _t0):.2f}s")
return (fg_batch, mask_batch)
# ---- ComfyUI discovery ----
NODE_CLASS_MAPPINGS = {
"PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai) [Batch]",
}
+441 -299
View File
@@ -1,21 +1,33 @@
# pvl_google_nano_banana.py
# Node: PVL Google Nano-Banana API (SDK + robust FAL fallback)
# Node: PVL Google Nano-Banana API (Gemini + FAL) — with delimiter support + parallel prompts
# Author: PVL
# License: MIT
#
# Requires:
# pip install google-genai
#
# Changes in this build:
# - FAL fallback now sends ONLY "image_urls" (list) to avoid duplicate outputs from some routes.
# - Enforces num_images: if num_images==1 returns the FIRST image only; otherwise caps to num_images.
# - Keeps sync_mode true and robust response parsing; prints raw FAL JSON when debug is ON.
# - Return order: ("STRING","IMAGE") => (text, images).
# pip install google-genai
import os, io, json, base64, typing as T, time
# Features:
# - Regex-based delimiter input for splitting prompts (default: [*])
# - Parallel API calls for multiple prompts
# - TRUE PARALLEL FAL execution: submit all requests first, then poll for results
# - If prompts < num_images, reuses the last prompt to fill remaining calls
# - Single optional ComfyUI IMAGE input (can be batched).
# - aspect_ratio supported for both Google and FAL.
# - use_fal_fallback default True (fallback to FAL when Google returns no images).
# - force_fal toggle to always use FAL (bypass Google).
# - Dual FAL routes: img2img (edit) vs txt2img (generate) chosen automatically.
# - Individual request error handling with partial results support.
import os
import io
import json
import base64
import typing as T
import time
import re
from concurrent.futures import ThreadPoolExecutor, as_completed
import requests
import numpy as np
from PIL import Image
@@ -24,23 +36,24 @@ import torch
NODE_NAME = "PVL Google Nano-Banana API"
NODE_CATEGORY = "PVL/Google"
DEFAULT_MODEL = "gemini-2.5-flash-image-preview"
_TOP_P = 0.95
_TOP_K = 64
_MAX_TOKENS = 4096
_VALID_ASPECTS = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"}
_VALID_ASPECTS = {"21:9", "1:1", "4:3", "3:2", "2:3", "5:4", "4:5", "3:4", "16:9", "9:16"}
# ----------------- image helpers -----------------
# --- image helpers ---
def pil_to_tensor(img: Image.Image) -> torch.Tensor:
if img.mode != "RGB":
img = img.convert("RGB")
arr = np.asarray(img, dtype=np.float32) / 255.0
return torch.from_numpy(arr)[None, ...]
t = torch.from_numpy(arr)[None, ...]
return t
def tensor_to_pil(t: torch.Tensor) -> Image.Image:
if t.ndim == 4:
t = t[0]
arr = (t.clamp(0,1).cpu().numpy() * 255).astype("uint8")
arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8")
return Image.fromarray(arr, "RGB")
def encode_pil_bytes(img: Image.Image, mime: str) -> bytes:
@@ -51,6 +64,24 @@ def encode_pil_bytes(img: Image.Image, mime: str) -> bytes:
img.save(buf, format="PNG")
return buf.getvalue()
def stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor:
"""Concatenate (B,H,W,C) batches along B. If shapes mismatch, resize to the first image size."""
if not tensors:
raise RuntimeError("No images to stack.")
try:
return torch.cat(tensors, dim=0)
except RuntimeError:
if debug:
print("[PVL NODE] Mismatched sizes, resizing to match first image.")
target_h, target_w = tensors[0].shape[1], tensors[0].shape[2]
fixed = []
for t in tensors:
pil = tensor_to_pil(t)
rp = pil.resize((target_w, target_h), Image.LANCZOS)
fixed.append(pil_to_tensor(rp))
return torch.cat(fixed, dim=0)
def _extract_image_bytes_from_part(part) -> T.Optional[bytes]:
try:
inline = getattr(part, "inline_data", None) or getattr(part, "inlineData", None)
@@ -65,13 +96,14 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]:
return None
except Exception:
pass
if isinstance(part, dict):
blob = None
if "inline_data" in part and isinstance(part["inline_data"], dict):
blob = part["inline_data"].get("data")
elif "inlineData" in part and isinstance(part["inlineData"], dict):
blob = part["inlineData"].get("data")
else:
blob = None
if isinstance(blob, (bytes, bytearray)):
return bytes(blob)
if isinstance(blob, str):
@@ -79,6 +111,7 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]:
return base64.b64decode(blob, validate=False)
except Exception:
return None
return None
def _extract_text_from_part(part) -> T.Optional[str]:
@@ -88,89 +121,99 @@ def _extract_text_from_part(part) -> T.Optional[str]:
return str(txt)
except Exception:
pass
if isinstance(part, dict) and "text" in part and part["text"] is not None:
return str(part["text"])
return None
def _data_url(mime: str, raw: bytes) -> str:
return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8")
class PVL_Google_NanoBanana_API:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "A tiny banana spaceship over a neon city."}),
"delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Regex delimiter e.g. [*], \\n, |"}),
},
"optional": {
"images": ("IMAGE",),
"aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2 — works for both Google & FAL"}),
"aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2"}),
"model": ("STRING", {"default": DEFAULT_MODEL}),
"endpoint_override": ("STRING", {"default": ""}),
"api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use GEMINI_API_KEY"}),
"api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use GEMINI_API_KEY"}),
"temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}),
"output_format": (["png","jpeg"], {"default": "png"}),
"output_format": (["png", "jpeg"], {"default": "png"}),
"capture_text_output": ("BOOLEAN", {"default": False}),
"num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}),
"timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}),
"request_id": ("STRING", {"default": ""}),
"debug_log": ("BOOLEAN", {"default": False}),
# --- FAL fallback ---
"use_fal_fallback": ("BOOLEAN", {"default": False}),
"fal_api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use FAL_KEY"}),
"fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}),
# FAL flags
"use_fal_fallback": ("BOOLEAN", {"default": True}),
"force_fal": ("BOOLEAN", {"default": False}),
"sync_mode": ("BOOLEAN", {"default": False}),
"fal_api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use FAL_KEY"}),
"fal_route_img2img": ("STRING", {"default": "fal-ai/nano-banana/edit"}),
"fal_route_txt2img": ("STRING", {"default": "fal-ai/nano-banana"}),
}
}
RETURN_TYPES = ("STRING","IMAGE",)
RETURN_NAMES = ("text","images")
RETURN_TYPES = ("STRING", "IMAGE",)
RETURN_NAMES = ("text", "images")
FUNCTION = "run"
CATEGORY = NODE_CATEGORY
# ---- helpers ----
# -------- helpers --------
def _make_client(self, api_key: str, endpoint_override: str):
try:
from google import genai
from google.genai import types
except Exception as e:
raise RuntimeError("Google GenAI SDK not installed. Run: pip install google-genai") from e
http_options = None
if endpoint_override.strip():
try:
http_options = types.HttpOptions(base_url=endpoint_override.strip())
except Exception:
http_options = None
if http_options is not None:
client = genai.Client(api_key=api_key, http_options=http_options)
else:
client = genai.Client(api_key=api_key)
return client
def _build_parts(self, prompt: str, images: T.Optional[torch.Tensor], mime: str):
parts: T.List[dict] = []
if prompt and prompt.strip():
parts.append({"text": prompt})
if images is not None and torch.is_tensor(images):
batch = images if images.ndim == 4 else images.unsqueeze(0)
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i+1])
parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}})
return parts
def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str = "1:1"):
def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str):
try:
from google.genai import types
cfg = types.GenerateContentConfig(
temperature=float(temperature),
top_p=float(_TOP_P),
top_k=int(_TOP_K),
max_output_tokens=int(_MAX_TOKENS),
response_modalities=["IMAGE","TEXT"] if want_text else ["IMAGE"],
response_modalities=["IMAGE", "TEXT"] if want_text else ["IMAGE"],
image_config=types.ImageConfig(aspect_ratio=aspect_ratio),
safety_settings=[
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"},
@@ -179,14 +222,17 @@ class PVL_Google_NanoBanana_API:
{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"},
],
)
return cfg
except Exception:
# Fallback dict config
return {
"temperature": float(temperature),
"top_p": float(_TOP_P),
"top_k": int(_TOP_K),
"max_output_tokens": int(_MAX_TOKENS),
"response_modalities": ["IMAGE","TEXT"] if want_text else ["IMAGE"],
"response_modalities": ["IMAGE", "TEXT"] if want_text else ["IMAGE"],
"image_config": {"aspect_ratio": aspect_ratio},
"safety_settings": [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"},
@@ -194,7 +240,35 @@ class PVL_Google_NanoBanana_API:
{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"},
],
}
def _build_call_prompts(self, base_prompts: T.List[str], num_images: int, debug: bool) -> T.List[str]:
"""
Maps prompts to calls according to the agreed rule:
- If len(prompts) >= num_images: take first num_images
- If len(prompts) < num_images: repeat the last prompt to fill N
"""
N = max(1, int(num_images))
if not base_prompts:
return [""]
if len(base_prompts) >= N:
call_prompts = base_prompts[:N]
else:
if debug:
print(f"[PVL NODE] Provided {len(base_prompts)} prompts but num_images={N}. "
f"Reusing the last prompt for remaining calls.")
print(f"[PVL WARNING] prompt list shorter than num_images ({len(base_prompts)} < {N}). "
f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.")
call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts))
if debug:
for i, cp in enumerate(call_prompts, 1):
show = cp if len(cp) <= 160 else (cp[:157] + "...")
print(f"[PVL NODE] Call {i} prompt: {show}")
return call_prompts
def _single_google_call(self, client, model: str, parts: list, cfg, request_id: str, debug: bool):
kwargs = {}
try:
@@ -203,19 +277,21 @@ class PVL_Google_NanoBanana_API:
kwargs["request_options"] = types.RequestOptions(request_id=request_id.strip())
except Exception:
pass
resp = client.models.generate_content(
model=model,
contents=[{"role": "user", "parts": parts}],
config=cfg,
**kwargs
)
imgs, texts = [], []
cands = getattr(resp, "candidates", None) or []
for cand in cands:
content = getattr(cand, "content", None)
finish_reason = getattr(cand, "finish_reason", None)
if debug:
try:
um = getattr(resp, "usage_metadata", None)
@@ -224,10 +300,12 @@ class PVL_Google_NanoBanana_API:
print(f"[PVL Debug] finish_reason: {finish_reason}")
except Exception:
pass
if content is None:
continue
parts = getattr(content, "parts", []) or []
for p in parts:
pparts = getattr(content, "parts", []) or []
for p in pparts:
blob = _extract_image_bytes_from_part(p)
if blob:
try:
@@ -240,298 +318,362 @@ class PVL_Google_NanoBanana_API:
t = _extract_text_from_part(p)
if t:
texts.append(t)
return imgs, texts, resp
# ---- FAL fallback via Queue API ----
def _fal_queue_call(self, route: str, prompt: str, image_tensor: T.Optional[torch.Tensor],
mime: str, fal_key: str, timeout: int, debug: bool, num_images: int, output_format: str):
# -------- FAL Queue API - TWO PHASE EXECUTION --------
def _fal_submit_only(self, route: str, prompt: str, image_tensors: T.Optional[torch.Tensor],
mime: str, fal_key: str, timeout: int, debug: bool,
output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False):
"""
Phase 1: Submit request to FAL queue and return request info immediately.
Does NOT poll for completion.
"""
if not fal_key:
raise RuntimeError("FAL fallback requested but FAL_KEY is missing. Provide fal_api_key or set env FAL_KEY.")
if image_tensor is None or not torch.is_tensor(image_tensor):
raise RuntimeError("FAL fallback requires an input image tensor.")
batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0)
# Prepare data URIs (all frames)
raise RuntimeError("FAL requested but FAL_KEY is missing.")
# Build data URLs only if image provided
data_urls: T.List[str] = []
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i+1])
raw = encode_pil_bytes(pil, mime)
data_urls.append(_data_url(mime, raw))
if image_tensors is not None and torch.is_tensor(image_tensors):
batch = image_tensors if image_tensors.ndim == 4 else image_tensors.unsqueeze(0)
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i + 1])
raw = encode_pil_bytes(pil, mime)
data_urls.append(_data_url(mime, raw))
base = "https://queue.fal.run"
submit_url = f"{base}/{route.strip()}"
headers = {"Authorization": f"Key {fal_key}"}
payload = {
"prompt": prompt or "",
# Only plural form to avoid duplicate outputs
"image_urls": data_urls,
"num_images": int(max(1, num_images)),
"output_format": ("png" if str(output_format).lower()=="png" else "jpeg"),
"sync_mode": True,
"num_images": 1,
"output_format": "png" if str(output_format).lower() == "png" else "jpeg",
"aspect_ratio": aspect_ratio,
"sync_mode": sync_mode,
}
# Only add image_urls for img2img route
if data_urls:
payload["image_urls"] = data_urls
if debug:
print(f"[PVL FAL] QUEUE SUBMIT {submit_url} with {len(data_urls)} image(s) and num_images={payload['num_images']}")
r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout)
if r.status_code >= 400:
raise RuntimeError(f"FAL queue submit error {r.status_code}: {r.text}")
try:
sub = r.json()
except Exception:
raise RuntimeError("FAL queue submit returned non-JSON.")
print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}")
# Submit request
rr = requests.post(submit_url, headers=headers, json=payload, timeout=timeout)
if rr.status_code != 200:
raise RuntimeError(f"FAL submit error {rr.status_code}: {rr.text}")
sub = rr.json()
req_id = sub.get("request_id")
status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None)
resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None)
if not req_id or not status_url or not resp_url:
raise RuntimeError("FAL queue submit missing request_id/status_url/response_url.")
# Poll
deadline = time.time() + max(5, int(timeout))
last_status = None
while time.time() < deadline:
sr = requests.get(status_url, headers=headers, timeout=10)
if not sr.ok:
time.sleep(0.5); continue
sdata = sr.json()
last_status = sdata.get("status")
if last_status == "COMPLETED":
break
time.sleep(0.6)
if last_status != "COMPLETED":
raise RuntimeError(f"FAL queue did not complete in time (last status={last_status})")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if rr.status_code >= 400:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
try:
rdata = rr.json()
except Exception:
raise RuntimeError("FAL result returned non-JSON.")
if not req_id:
raise RuntimeError("FAL did not return a request_id")
# Get status and result URLs
status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
"prompt": prompt
}
def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns (image_tensor, description_text).
"""
headers = {"Authorization": f"Key {fal_key}"}
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
req_id = request_info["request_id"]
if debug:
print(f"[FAL POLL] request_id={req_id[:16]}...")
# Poll for completion with timeout check
deadline = time.time() + timeout
completed = False
while time.time() < deadline:
try:
s = json.dumps(rdata)[:1200]
except Exception:
s = str(rdata)[:1200]
print("[PVL FAL] raw response:", s)
# Normalize containers
resp = rdata.get("response") if isinstance(rdata, dict) else None
if not isinstance(resp, dict):
resp = rdata if isinstance(rdata, dict) else {}
description = resp.get("description") or rdata.get("description") or resp.get("output_text") or ""
# Collect potential image items
sr = requests.get(status_url, headers=headers, timeout=min(10, timeout))
if sr.ok and sr.json().get("status") == "COMPLETED":
completed = True
break
except Exception as e:
if debug:
print(f"[FAL POLL] Status check error: {e}")
time.sleep(0.6)
# Check if we timed out
if not completed:
raise RuntimeError(f"FAL request {req_id[:16]} timed out after {timeout}s")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=min(15, timeout))
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
data = rr.json()
if debug:
print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED")
# Extract response data
resp = data.get("response") if isinstance(data, dict) else None
if resp is None and isinstance(data, dict):
resp = data
# Parse images from various possible locations
images_out: T.List[torch.Tensor] = []
buckets: T.List[T.Union[str, dict]] = []
for key in ("images", "outputs", "artifacts"):
val = resp.get(key)
if isinstance(val, list):
buckets.extend(val)
for key in ("image", "output", "result"):
val = resp.get(key)
if isinstance(val, (str, dict)):
buckets.append(val)
for key in ("images", "image", "output", "outputs", "artifacts"):
val = rdata.get(key) if isinstance(rdata, dict) else None
if isinstance(val, list):
buckets.extend(val)
elif isinstance(val, (str, dict)):
buckets.append(val)
# Decode images, but respect num_images cap
out = []
def add_image_from_item(item):
if isinstance(resp, dict):
for key in ("images", "outputs", "artifacts"):
val = resp.get(key)
if isinstance(val, list):
buckets.extend(val)
for key in ("image", "output", "result"):
val = resp.get(key)
if isinstance(val, (str, dict)):
buckets.append(val)
for item in buckets:
try:
if isinstance(item, str):
url_or_data = item
elif isinstance(item, dict):
url_or_data = item.get("url") or item.get("data") or item.get("image") or item.get("content")
else:
return
url_or_data = item if isinstance(item, str) else (item.get("url") or item.get("data") or item.get("image"))
if not isinstance(url_or_data, str):
return
continue
if url_or_data.startswith("data:image/"):
b64 = url_or_data.split(",", 1)[1]
blob = base64.b64decode(b64)
else:
ir = requests.get(url_or_data, timeout=timeout)
ir = requests.get(url_or_data, timeout=min(15, timeout))
if not ir.ok:
return
continue
blob = ir.content
pil = Image.open(io.BytesIO(blob)).convert("RGB")
out.append(pil_to_tensor(pil))
images_out.append(pil_to_tensor(pil))
except Exception as ex:
if debug:
print("[PVL FAL] image decode failed:", ex)
for item in buckets:
if len(out) >= int(max(1, num_images)):
break
add_image_from_item(item)
if not out:
raise RuntimeError("FAL API returned no images.")
if len(out) == 1:
images_tensor = out[0]
else:
images_tensor = torch.cat(out, dim=0)
if debug:
print(f"[PVL FAL] returning {len(out)} image(s) (capped to num_images={int(max(1, num_images))})")
return images_tensor, description
# ---- main ----
def run(self, prompt: str, images: T.Optional[torch.Tensor] = None,
aspect_ratio: str = "1:1",
model: str = DEFAULT_MODEL, endpoint_override: str = "",
api_key: str = "",
temperature: float = 0.6, output_format: str = "png",
capture_text_output: bool = False, num_images: int = 1,
timeout_sec: int = 120, request_id: str = "",
debug_log: bool = False,
use_fal_fallback: bool = False, fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"):
print("[FAL] image decode failed:", ex)
description = resp.get("description", "") if isinstance(resp, dict) else ""
if not images_out:
raise RuntimeError(f"FAL returned no images for request_id={req_id}")
return images_out[0], description
# -------- main --------
def run(
self,
prompt: str,
delimiter: str = "[*]",
images: T.Optional[torch.Tensor] = None,
aspect_ratio: str = "1:1",
model: str = DEFAULT_MODEL,
endpoint_override: str = "",
api_key: str = "",
temperature: float = 0.6,
output_format: str = "png",
capture_text_output: bool = False,
num_images: int = 1,
timeout_sec: int = 120,
request_id: str = "",
debug_log: bool = False,
use_fal_fallback: bool = True,
force_fal: bool = False,
sync_mode: bool = False,
fal_api_key: str = "",
fal_route_img2img: str = "fal-ai/nano-banana/edit",
fal_route_txt2img: str = "fal-ai/nano-banana",
):
# Validate aspect_ratio
if aspect_ratio.strip() not in _VALID_ASPECTS:
print(f"[PVL WARNING] Invalid or missing aspect_ratio '{aspect_ratio}', defaulting to 1:1.")
aspect_ratio = "1:1"
key = (api_key or os.getenv("GEMINI_API_KEY","")).strip()
# Split prompts using regex delimiter
try:
base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()]
except re.error:
print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.")
base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()]
if not base_prompts:
raise RuntimeError("No valid prompts provided.")
# Map prompts to num_images calls
call_prompts = self._build_call_prompts(base_prompts, num_images, debug_log)
key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip()
input_mime = "image/png" if str(output_format).lower() == "png" else "image/jpeg"
want_text = bool(capture_text_output)
# If no Google key but fallback is enabled, attempt FAL directly
# Decide FAL route based on presence of input images
route_to_use = fal_route_img2img if (images is not None and torch.is_tensor(images) and images.numel() > 0) else fal_route_txt2img
# Helper function for parallel FAL submission + polling
def parallel_fal_execution(prompts_list, fal_key_str, route, debug):
"""Submit all FAL requests in parallel, then poll all in parallel"""
if debug:
print(f"[FAL] Submitting {len(prompts_list)} requests in parallel...")
# PHASE 1: Submit all requests IN PARALLEL
submit_results = []
with ThreadPoolExecutor(max_workers=min(len(prompts_list), 6)) as ex:
submit_futs = {
ex.submit(self._fal_submit_only, route, p, images, input_mime,
fal_key_str, timeout_sec, debug, output_format,
aspect_ratio, sync_mode): p
for p in prompts_list
}
for fut in as_completed(submit_futs):
try:
req_info = fut.result()
submit_results.append(req_info)
except Exception as e:
if debug:
print(f"[FAL SUBMIT ERROR] {e}")
if not submit_results:
raise RuntimeError("All FAL submission requests failed")
if debug:
print(f"[FAL] {len(submit_results)} requests submitted successfully. Polling for results...")
# PHASE 2: Poll all requests IN PARALLEL
results, texts = [], []
failed_count = 0
with ThreadPoolExecutor(max_workers=min(len(submit_results), 6)) as ex:
poll_futs = {
ex.submit(self._fal_poll_and_fetch, req_info, fal_key_str,
timeout_sec, debug): req_info
for req_info in submit_results
}
for fut in as_completed(poll_futs):
try:
img, t = fut.result()
results.append(img)
if t:
texts.append(t)
except Exception as e:
failed_count += 1
if debug:
print(f"[FAL POLL ERROR] {e}")
if not results:
raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)")
if failed_count > 0:
print(f"[PVL WARNING] {failed_count}/{len(submit_results)} FAL requests failed, continuing with {len(results)} successful results")
return results, texts
# ---- CASE: FAL only (TRUE PARALLEL) ----
if force_fal:
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("force_fal=True but FAL_KEY missing.")
results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log)
images_tensor = stack_images_same_size(results, debug_log)
text_out = "\n".join(texts) if want_text else ""
return text_out, images_tensor
# ---- Google path ----
if not key:
if use_fal_fallback:
fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip()
try:
img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format)
text_out = (fal_text or "") if want_text else ""
if text_out:
print("[PVL FAL Text]:\n" + text_out)
return (text_out, img_tensor)
except Exception as fe:
print(f"[PVL FAL Fallback] FAL call failed without Google key: {fe}")
raise RuntimeError(f"Gemini image generation failed and FAL fallback also failed: {fe}")
# no Google key, try FAL directly (TRUE PARALLEL)
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("GEMINI_API_KEY missing and FAL_KEY missing.")
results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log)
images_tensor = stack_images_same_size(results, debug_log)
text_out = "\n".join(texts) if want_text else ""
return text_out, images_tensor
raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.")
# Build client & request
client = self._make_client(key, endpoint_override)
parts = self._build_parts(prompt, images, input_mime)
cfg = self._build_config(temperature, want_text, aspect_ratio)
if debug_log:
p_preview = (prompt or "")[:180].replace("\n"," ")
img_count = (images.shape[0] if (isinstance(images, torch.Tensor) and images.ndim==4) else (1 if isinstance(images, torch.Tensor) else 0))
print(f"[PVL Debug] prompt chars={len(prompt or '')} preview='{p_preview}...'")
print(f"[PVL Debug] parts: text={1 if (prompt and prompt.strip()) else 0}, images={img_count}")
safe_parts = []
for pr in parts:
if "text" in pr:
safe_parts.append({"text": pr["text"][:120]})
elif "inline_data" in pr:
di = pr["inline_data"]
safe_parts.append({"inline_data": {"mime_type": di.get("mime_type","image/*"), "data": "<bytes>"}})
try:
temp = getattr(cfg,'temperature',None) if hasattr(cfg,'temperature') else cfg.get('temperature')
top_p = getattr(cfg,'top_p',None) if hasattr(cfg,'top_p') else cfg.get('top_p')
top_k = getattr(cfg,'top_k',None) if hasattr(cfg,'top_k') else cfg.get('top_k')
mot = getattr(cfg,'max_output_tokens',None) if hasattr(cfg,'max_output_tokens') else cfg.get('max_output_tokens')
mods = getattr(cfg,'response_modalities',None) if hasattr(cfg,'response_modalities') else cfg.get('response_modalities')
except Exception:
temp=top_p=top_k=mot=mods=None
print("[PVL Debug] config:", {"temperature": temp, "top_p": top_p, "top_k": top_k, "max_output_tokens": mot, "modalities": mods})
print("[PVL Debug] contents:", [{"role":"user","parts": safe_parts}])
# Parallel Google calls
N = max(1, int(num_images))
results = [None] * N
errors = []
def call_i(i: int):
rid = (request_id.strip() + f"-{i}") if request_id.strip() else f"pvl-nb-{int(time.time()*1000)}-{i}"
return self._single_google_call(client, model, parts, cfg, rid, debug_log)
max_workers = min(N, 6)
if N == 1:
try:
results[0] = call_i(0)
except Exception as e:
errors.append(f"google call 0 failed: {e}")
else:
with ThreadPoolExecutor(max_workers=max_workers) as ex:
futmap = {ex.submit(call_i, i): i for i in range(N)}
for fut in as_completed(futmap):
i = futmap[fut]
try:
results[i] = fut.result()
except Exception as e:
errors.append(f"google call {i} failed: {e}")
# Parse all Google results
out_imgs, out_texts = [], []
for idx, tup in enumerate(results):
if tup is None:
errors.append(f"google call {idx} returned no response")
continue
imgs_i, texts_i, resp = tup
if not imgs_i:
print(f"[PVL Debug] Processing {len(call_prompts)} prompts in parallel")
def google_call(p: str, debug: bool):
parts = self._build_parts(p, images, input_mime)
g_imgs, g_texts, _resp = self._single_google_call(client, model, parts, cfg, request_id, debug)
return g_imgs, g_texts
out_imgs: T.List[torch.Tensor] = []
out_texts: T.List[str] = []
failed_google = 0
# Google API calls with error handling
with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex:
futs = {ex.submit(google_call, p, debug_log): p for p in call_prompts}
for fut in as_completed(futs):
try:
cands = getattr(resp, "candidates", None) or []
fin = getattr(cands[0], "finish_reason", None) if cands else None
except Exception:
fin = None
errors.append(f"google call {idx} returned no images (finish_reason={fin})")
else:
out_imgs.extend(imgs_i)
if texts_i:
out_texts.append("\n".join(texts_i))
# If Google failed and fallback enabled -> try FAL
if errors and use_fal_fallback:
print("[PVL Fallback] Google call failed; attempting FAL.ai fallback...")
for e in errors:
print("[PVL Google Error]", e)
try:
fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip()
img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format)
final_text = "" # Start with empty; then merge Google texts if requested
if bool(capture_text_output):
pieces = []
if out_texts:
pieces.append("\n\n--- Google ---\n\n" + ("\n".join(out_texts)))
if fal_text:
pieces.append("\n\n--- FAL ---\n\n" + fal_text)
final_text = "".join(pieces)
if final_text:
print("[PVL Fallback Note] Combined text:")
print(final_text)
return (final_text, img_tensor)
except Exception as fe:
print("[PVL FAL Error]", fe)
raise RuntimeError("Both Google and FAL failed. See console for details.")
# If Google produced errors and fallback not used -> raise
if errors:
if out_texts:
print("[PVL Google Text]:\n" + ("\n\n---\n\n".join(out_texts)))
raise RuntimeError("Gemini image generation failed: " + " | ".join(errors[:5]))
imgs, texts = fut.result()
if imgs:
out_imgs.append(imgs[0])
out_texts.extend(texts)
except Exception as e:
failed_google += 1
if debug_log:
print(f"[GOOGLE ERROR] {e}")
if failed_google > 0:
print(f"[PVL WARNING] {failed_google}/{len(call_prompts)} Google requests failed")
# If Google failed to produce images, optionally fallback to FAL (TRUE PARALLEL)
if not out_imgs:
raise RuntimeError("Gemini image generation failed: no images across all Google calls.")
images_tensor = torch.cat(out_imgs, dim=0) if len(out_imgs) > 1 else out_imgs[0]
final_text = ("\n\n---\n\n".join(out_texts)) if (bool(capture_text_output) and out_texts) else ""
if final_text:
print(f"[PVL Google NanoBanana Output]:\n{final_text}\n")
return (final_text, images_tensor,)
if use_fal_fallback:
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("Google returned no images and FAL_KEY is missing for fallback.")
results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log)
images_tensor = stack_images_same_size(results, debug_log)
if want_text:
combined_text = ""
if out_texts:
combined_text += "\n\n--- Google ---\n\n" + "\n".join(out_texts)
if texts:
combined_text += "\n\n--- FAL ---\n\n" + "\n".join(texts)
text_out = combined_text
else:
text_out = ""
return text_out, images_tensor
else:
raise RuntimeError("Gemini returned no images")
# Merge Google images to a 4D tensor
images_tensor = stack_images_same_size(out_imgs, debug_log)
if images_tensor.ndim == 3:
images_tensor = images_tensor.unsqueeze(0)
text_out = "\n".join(out_texts) if (want_text and out_texts) else ""
if text_out and debug_log:
print("[PVL Google NanoBanana Output]:\n" + text_out)
return text_out, images_tensor
NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_API": PVL_Google_NanoBanana_API}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_API": NODE_NAME}
+408 -406
View File
@@ -1,46 +1,47 @@
# pvl_google_nano_banana_mandatory_img.py
# Node: PVL Google Nano-Banana API (Mandatory IMG + Delimiter + Aspect Ratio + Parallel + FAL)
# pvl_google_nano_banana.py
# Node: PVL Google Nano-Banana API (SDK + robust FAL fallback)
# Author: PVL
# License: MIT
#
# Requires:
# pip install google-genai
#
# Changes in this build:
# - FAL fallback now sends ONLY "image_urls" (list) to avoid duplicate outputs from some routes.
# - Enforces num_images: if num_images==1 returns the FIRST image only; otherwise caps to num_images.
# - Keeps sync_mode true and robust response parsing; prints raw FAL JSON when debug is ON.
# - Return order: ("STRING","IMAGE") => (text, images).
# Features in this build:
# - Mandatory image input.
# - Literal delimiter input (default "[*]") to split a long prompt into multiple prompts.
# - Prompt reuse rule to match num_images
# - Aspect ratio support (STRING) via new SDK
# - TRUE PARALLEL execution for both Google and FAL
# - FAL fallback and Force-FAL modes with sync_mode toggle
# - Error handling for individual requests with partial results support
# - Execution time is printed at the end
import os, io, json, base64, typing as T, time
from concurrent.futures import ThreadPoolExecutor, as_completed
import requests
import numpy as np
from PIL import Image
import torch
from google.genai import types
NODE_NAME = "PVL Google Nano-Banana API mandatory IMG"
NODE_CATEGORY = "PVL/Google"
DEFAULT_MODEL = "gemini-2.5-flash-image-preview"
_TOP_P = 0.95
_TOP_K = 64
_MAX_TOKENS = 4096
_VALID_ASPECTS = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"}
# --- image helpers ---
# --------------------------- Image/Tensor helpers ---------------------------
def pil_to_tensor(img: Image.Image) -> torch.Tensor:
if img.mode != "RGB":
img = img.convert("RGB")
arr = np.asarray(img, dtype=np.float32) / 255.0
return torch.from_numpy(arr)[None, ...]
return torch.from_numpy(arr)[None, ...] # (1,H,W,3)
def tensor_to_pil(t: torch.Tensor) -> Image.Image:
if t.ndim == 4:
t = t[0]
arr = (t.clamp(0,1).cpu().numpy() * 255).astype("uint8")
arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8")
return Image.fromarray(arr, "RGB")
def encode_pil_bytes(img: Image.Image, mime: str) -> bytes:
@@ -65,475 +66,476 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]:
return None
except Exception:
pass
if isinstance(part, dict):
if "inline_data" in part and isinstance(part["inline_data"], dict):
blob = part["inline_data"].get("data")
elif "inlineData" in part and isinstance(part["inlineData"], dict):
blob = part["inlineData"].get("data")
else:
blob = None
if isinstance(blob, (bytes, bytearray)):
return bytes(blob)
if isinstance(blob, str):
try:
return base64.b64decode(blob, validate=False)
except Exception:
return None
inline = part.get("inline_data") or part.get("inlineData")
if isinstance(inline, dict):
data = inline.get("data")
if isinstance(data, (bytes, bytearray)):
return bytes(data)
if isinstance(data, str):
try:
return base64.b64decode(data, validate=False)
except Exception:
return None
return None
def _extract_text_from_part(part) -> T.Optional[str]:
try:
txt = getattr(part, "text", None)
if txt is not None:
return str(txt)
except Exception:
pass
if isinstance(part, dict) and "text" in part and part["text"] is not None:
return str(part["text"])
return None
if isinstance(part, dict):
if "text" in part and part["text"] is not None:
return str(part["text"])
txt = getattr(part, "text", None)
return str(txt) if txt is not None else None
def _data_url(mime: str, raw: bytes) -> str:
return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8")
def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor:
if not tensors:
raise RuntimeError("No images to stack.")
try:
return torch.cat(tensors, dim=0)
except RuntimeError:
if debug:
print("[PVL NODE] Mismatched sizes, resizing to match first image.")
target_h, target_w = tensors[0].shape[1], tensors[0].shape[2]
fixed = []
for t in tensors:
pil = tensor_to_pil(t)
rp = pil.resize((target_w, target_h), Image.LANCZOS)
fixed.append(pil_to_tensor(rp))
return torch.cat(fixed, dim=0)
# --------------------------- Main Node ---------------------------
class PVL_Google_NanoBanana_API_mandatory_IMG:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "A tiny banana spaceship over a neon city."}),
"prompt": ("STRING", {"multiline": True, "default": "", "placeholder": "Enter prompts separated by delimiter"}),
"images": ("IMAGE",),
},
"optional": {
"delimiter": ("STRING", {"default": "[*]", "placeholder": "Delimiter string"}),
"aspect_ratio": ("STRING", {"default": "1:1"}),
"model": ("STRING", {"default": DEFAULT_MODEL}),
"endpoint_override": ("STRING", {"default": ""}),
"api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use GEMINI_API_KEY"}),
"api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use GEMINI_API_KEY"}),
"temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}),
"output_format": (["png","jpeg"], {"default": "png"}),
"output_format": (["png", "jpeg"], {"default": "png"}),
"capture_text_output": ("BOOLEAN", {"default": False}),
"num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}),
"timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}),
"request_id": ("STRING", {"default": ""}),
"debug_log": ("BOOLEAN", {"default": False}),
"aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2 — works for both Google & FAL"}),
# --- FAL fallback ---
"use_fal_fallback": ("BOOLEAN", {"default": False}),
"fal_api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use FAL_KEY"}),
"use_fal_fallback": ("BOOLEAN", {"default": True}),
"force_fal": ("BOOLEAN", {"default": False}),
"sync_mode": ("BOOLEAN", {"default": False}),
"fal_api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use FAL_KEY"}),
"fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}),
}
}
RETURN_TYPES = ("STRING","IMAGE",)
RETURN_NAMES = ("text","images")
RETURN_TYPES = ("STRING", "IMAGE",)
RETURN_NAMES = ("text", "images")
FUNCTION = "run"
CATEGORY = NODE_CATEGORY
# ---- helpers ----
# ------------------------- INTERNAL HELPERS -------------------------
def _make_client(self, api_key: str, endpoint_override: str):
try:
from google import genai
from google.genai import types
except Exception as e:
raise RuntimeError("Google GenAI SDK not installed. Run: pip install google-genai") from e
from google import genai
from google.genai import types as gtypes
http_options = None
if endpoint_override.strip():
try:
http_options = types.HttpOptions(base_url=endpoint_override.strip())
http_options = gtypes.HttpOptions(base_url=endpoint_override.strip())
except Exception:
http_options = None
if http_options is not None:
client = genai.Client(api_key=api_key, http_options=http_options)
else:
client = genai.Client(api_key=api_key)
return client
def _build_parts(self, prompt: str, images: T.Optional[torch.Tensor], mime: str):
return genai.Client(api_key=api_key, http_options=http_options)
return genai.Client(api_key=api_key)
def _build_parts(self, prompt: str, images: torch.Tensor, mime: str):
parts: T.List[dict] = []
if prompt and prompt.strip():
parts.append({"text": prompt})
if images is not None and torch.is_tensor(images):
batch = images if images.ndim == 4 else images.unsqueeze(0)
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i+1])
parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}})
batch = images if images.ndim == 4 else images.unsqueeze(0)
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i+1])
parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}})
return parts
def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str = "1:1"):
try:
from google.genai import types
cfg = types.GenerateContentConfig(
temperature=float(temperature),
top_p=float(_TOP_P),
top_k=int(_TOP_K),
max_output_tokens=int(_MAX_TOKENS),
response_modalities=["IMAGE","TEXT"] if want_text else ["IMAGE"],
image_config=types.ImageConfig(aspect_ratio=aspect_ratio),
safety_settings=[
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"},
],
)
return cfg
except Exception:
return {
"temperature": float(temperature),
"top_p": float(_TOP_P),
"top_k": int(_TOP_K),
"max_output_tokens": int(_MAX_TOKENS),
"response_modalities": ["IMAGE","TEXT"] if want_text else ["IMAGE"],
"safety_settings": [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"},
],
}
def _single_google_call(self, client, model: str, parts: list, cfg, request_id: str, debug: bool):
kwargs = {}
try:
from google.genai import types
if request_id.strip():
kwargs["request_options"] = types.RequestOptions(request_id=request_id.strip())
except Exception:
pass
resp = client.models.generate_content(
model=model,
contents=[{"role": "user", "parts": parts}],
config=cfg,
**kwargs
)
imgs, texts = [], []
cands = getattr(resp, "candidates", None) or []
for cand in cands:
content = getattr(cand, "content", None)
finish_reason = getattr(cand, "finish_reason", None)
if debug:
try:
um = getattr(resp, "usage_metadata", None)
if um is not None:
print("[PVL Debug] usage_metadata:", getattr(um, "__dict__", str(um)))
print(f"[PVL Debug] finish_reason: {finish_reason}")
except Exception:
pass
if content is None:
continue
parts = getattr(content, "parts", []) or []
for p in parts:
blob = _extract_image_bytes_from_part(p)
if blob:
try:
pil = Image.open(io.BytesIO(blob)).convert("RGB")
imgs.append(pil_to_tensor(pil))
except Exception as ex:
if debug:
print("[PVL Debug] image decode error:", ex)
else:
t = _extract_text_from_part(p)
if t:
texts.append(t)
return imgs, texts, resp
# ---- FAL fallback via Queue API ----
def _fal_queue_call(self, route: str, prompt: str, image_tensor: T.Optional[torch.Tensor],
mime: str, fal_key: str, timeout: int, debug: bool, num_images: int, output_format: str):
# ---- FAL Queue API - TWO PHASE EXECUTION ----
def _fal_submit_only(self, route: str, prompt: str, image_tensor: torch.Tensor,
mime: str, fal_key: str, timeout: int, debug: bool,
output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False):
"""
Phase 1: Submit request to FAL queue and return request info immediately.
Does NOT poll for completion.
"""
if not fal_key:
raise RuntimeError("FAL fallback requested but FAL_KEY is missing. Provide fal_api_key or set env FAL_KEY.")
if image_tensor is None or not torch.is_tensor(image_tensor):
raise RuntimeError("FAL fallback requires an input image tensor.")
batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0)
# Prepare data URIs (all frames)
raise RuntimeError("FAL requested but FAL_KEY is missing.")
# Build data URLs for all frames
data_urls: T.List[str] = []
batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0)
for i in range(batch.shape[0]):
pil = tensor_to_pil(batch[i:i+1])
raw = encode_pil_bytes(pil, mime)
data_urls.append(_data_url(mime, raw))
base = "https://queue.fal.run"
submit_url = f"{base}/{route.strip()}"
headers = {"Authorization": f"Key {fal_key}"}
payload = {
"prompt": prompt or "",
# Only plural form to avoid duplicate outputs
"image_urls": data_urls,
"num_images": int(max(1, num_images)),
"output_format": ("png" if str(output_format).lower()=="png" else "jpeg"),
"sync_mode": True,
"num_images": 1,
"output_format": ("png" if str(output_format).lower() == "png" else "jpeg"),
"aspect_ratio": aspect_ratio,
"sync_mode": sync_mode,
}
if debug:
print(f"[PVL FAL] QUEUE SUBMIT {submit_url} with {len(data_urls)} image(s) and num_images={payload['num_images']}")
print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}")
# Submit request
r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout)
if r.status_code >= 400:
raise RuntimeError(f"FAL queue submit error {r.status_code}: {r.text}")
try:
sub = r.json()
except Exception:
raise RuntimeError("FAL queue submit returned non-JSON.")
if not r.ok:
raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}")
sub = r.json()
req_id = sub.get("request_id")
status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None)
resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None)
if not req_id or not status_url or not resp_url:
raise RuntimeError("FAL queue submit missing request_id/status_url/response_url.")
# Poll
deadline = time.time() + max(5, int(timeout))
last_status = None
if not req_id:
raise RuntimeError("FAL did not return a request_id")
# Get status and result URLs
status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
"prompt": prompt
}
def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns (image_tensor, description_text).
"""
headers = {"Authorization": f"Key {fal_key}"}
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
req_id = request_info["request_id"]
if debug:
print(f"[FAL POLL] request_id={req_id[:16]}...")
# Poll for completion with timeout check
deadline = time.time() + timeout
completed = False
while time.time() < deadline:
sr = requests.get(status_url, headers=headers, timeout=10)
if not sr.ok:
time.sleep(0.5); continue
sdata = sr.json()
last_status = sdata.get("status")
if last_status == "COMPLETED":
break
try:
sr = requests.get(status_url, headers=headers, timeout=10)
if sr.ok and sr.json().get("status") == "COMPLETED":
completed = True
break
except Exception as e:
if debug:
print(f"[FAL POLL] Status check error: {e}")
time.sleep(0.6)
if last_status != "COMPLETED":
raise RuntimeError(f"FAL queue did not complete in time (last status={last_status})")
# Check if we timed out
if not completed:
raise RuntimeError(f"FAL request {req_id[:16]} timed out after {timeout}s")
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if rr.status_code >= 400:
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
try:
rdata = rr.json()
except Exception:
raise RuntimeError("FAL result returned non-JSON.")
rdata = rr.json()
if debug:
print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED")
# Extract response data
resp = rdata.get("response", rdata) if isinstance(rdata, dict) else rdata
# Parse images from various possible locations
buckets = []
if isinstance(resp, dict):
for key in ("images", "outputs", "artifacts"):
val = resp.get(key)
if isinstance(val, list):
buckets.extend(val)
for key in ("image", "output", "result"):
val = resp.get(key)
if isinstance(val, (str, dict)):
buckets.append(val)
out: T.List[torch.Tensor] = []
for item in buckets:
try:
s = json.dumps(rdata)[:1200]
except Exception:
s = str(rdata)[:1200]
print("[PVL FAL] raw response:", s)
# Normalize containers
resp = rdata.get("response") if isinstance(rdata, dict) else None
if not isinstance(resp, dict):
resp = rdata if isinstance(rdata, dict) else {}
description = resp.get("description") or rdata.get("description") or resp.get("output_text") or ""
# Collect potential image items
buckets: T.List[T.Union[str, dict]] = []
for key in ("images", "outputs", "artifacts"):
val = resp.get(key)
if isinstance(val, list):
buckets.extend(val)
for key in ("image", "output", "result"):
val = resp.get(key)
if isinstance(val, (str, dict)):
buckets.append(val)
for key in ("images", "image", "output", "outputs", "artifacts"):
val = rdata.get(key) if isinstance(rdata, dict) else None
if isinstance(val, list):
buckets.extend(val)
elif isinstance(val, (str, dict)):
buckets.append(val)
# Decode images, but respect num_images cap
out = []
def add_image_from_item(item):
try:
if isinstance(item, str):
url_or_data = item
elif isinstance(item, dict):
url_or_data = item.get("url") or item.get("data") or item.get("image") or item.get("content")
url = item if isinstance(item, str) else (item.get("url") or item.get("data") or item.get("image"))
if not url:
continue
if url.startswith("data:image/"):
blob = base64.b64decode(url.split(",", 1)[1])
else:
return
if not isinstance(url_or_data, str):
return
if url_or_data.startswith("data:image/"):
b64 = url_or_data.split(",", 1)[1]
blob = base64.b64decode(b64)
else:
ir = requests.get(url_or_data, timeout=timeout)
ir = requests.get(url, timeout=timeout)
if not ir.ok:
return
continue
blob = ir.content
pil = Image.open(io.BytesIO(blob)).convert("RGB")
out.append(pil_to_tensor(pil))
except Exception as ex:
if debug:
print("[PVL FAL] image decode failed:", ex)
for item in buckets:
if len(out) >= int(max(1, num_images)):
break
add_image_from_item(item)
print("[FAL decode fail]", ex)
if not out:
raise RuntimeError("FAL API returned no images.")
if len(out) == 1:
images_tensor = out[0]
else:
images_tensor = torch.cat(out, dim=0)
if debug:
print(f"[PVL FAL] returning {len(out)} image(s) (capped to num_images={int(max(1, num_images))})")
raise RuntimeError(f"FAL returned no images for request_id={req_id}")
description = ""
if isinstance(resp, dict):
description = resp.get("description") or resp.get("output_text") or ""
images_tensor = _stack_images_same_size(out, debug)
return images_tensor, description
# ---- main ----
def run(self, prompt: str, images: T.Optional[torch.Tensor] = None,
aspect_ratio: str = "1:1",
# --------------------------- RUN MAIN -----------------------------
def run(self, prompt: str, images: torch.Tensor,
delimiter: str = "[*]", aspect_ratio: str = "1:1",
model: str = DEFAULT_MODEL, endpoint_override: str = "",
api_key: str = "",
temperature: float = 0.6, output_format: str = "png",
capture_text_output: bool = False, num_images: int = 1,
timeout_sec: int = 120, request_id: str = "",
debug_log: bool = False,
use_fal_fallback: bool = False, fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"):
if aspect_ratio.strip() not in _VALID_ASPECTS:
print(f"[PVL WARNING] Invalid or missing aspect_ratio '{aspect_ratio}', defaulting to 1:1.")
aspect_ratio = "1:1"
key = (api_key or os.getenv("GEMINI_API_KEY","")).strip()
input_mime = "image/png" if str(output_format).lower() == "png" else "image/jpeg"
api_key: str = "", temperature: float = 0.6,
output_format: str = "png", capture_text_output: bool = False,
num_images: int = 1, timeout_sec: int = 120,
debug_log: bool = False, use_fal_fallback: bool = True,
force_fal: bool = False, sync_mode: bool = False,
fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"):
_t0 = time.time()
key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip()
input_mime = "image/png" if output_format.lower() == "png" else "image/jpeg"
want_text = bool(capture_text_output)
# If no Google key but fallback is enabled, attempt FAL directly
N = max(1, int(num_images))
# Split prompts by literal delimiter and expand to N using reuse rule
raw_parts = [p.strip() for p in str(prompt).split(delimiter) if p.strip()]
if not raw_parts:
raise RuntimeError("Prompt is empty after splitting by delimiter.")
if len(raw_parts) >= N:
prompts = raw_parts[:N]
else:
prompts = raw_parts + [raw_parts[-1]] * (N - len(raw_parts))
if debug_log:
print(f"[PVL Debug] delimiter='{delimiter}' aspect_ratio='{aspect_ratio}' num_images={N} sync_mode={sync_mode}")
for i, pr in enumerate(prompts, 1):
preview = pr if len(pr) <= 160 else (pr[:157] + "...")
print(f"[PVL Debug] Call #{i} prompt: {preview}")
# Helper function for parallel FAL submission + polling
def parallel_fal_execution(prompts_list, fal_key_str, debug):
"""Submit all FAL requests in parallel, then poll all in parallel"""
if debug:
print(f"[FAL] Submitting {len(prompts_list)} requests in parallel...")
# PHASE 1: Submit all requests IN PARALLEL
submit_results = []
with ThreadPoolExecutor(max_workers=min(len(prompts_list), 6)) as ex:
submit_futs = {
ex.submit(self._fal_submit_only, fal_route, p, images, input_mime,
fal_key_str, int(timeout_sec), debug, output_format,
aspect_ratio, sync_mode): p
for p in prompts_list
}
for fut in as_completed(submit_futs):
try:
req_info = fut.result()
submit_results.append(req_info)
except Exception as e:
if debug:
print(f"[FAL SUBMIT ERROR] {e}")
if not submit_results:
raise RuntimeError("All FAL submission requests failed")
if debug:
print(f"[FAL] {len(submit_results)} requests submitted successfully. Polling for results...")
# PHASE 2: Poll all requests IN PARALLEL
images_list, texts = [], []
failed_count = 0
with ThreadPoolExecutor(max_workers=min(len(submit_results), 6)) as ex:
poll_futs = {
ex.submit(self._fal_poll_and_fetch, req_info, fal_key_str,
int(timeout_sec), debug): req_info
for req_info in submit_results
}
for fut in as_completed(poll_futs):
try:
it, txt = fut.result()
images_list.append(it)
if want_text and txt:
texts.append(txt)
except Exception as e:
failed_count += 1
if debug:
print(f"[FAL POLL ERROR] {e}")
if not images_list:
raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)")
if failed_count > 0:
print(f"[PVL WARNING] {failed_count}/{len(submit_results)} FAL requests failed, continuing with {len(images_list)} successful results")
return images_list, texts
# Force-FAL path (TRUE PARALLEL)
if force_fal:
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("force_fal=True but FAL_KEY missing.")
images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log)
images_tensor = _stack_images_same_size(images_list, debug_log)
text_out = "\n".join(texts) if (want_text and texts) else ""
_t1 = time.time()
print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s")
return text_out, images_tensor
# No API key? Try fallback if enabled (TRUE PARALLEL)
if not key:
if use_fal_fallback:
fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip()
try:
img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format)
text_out = (fal_text or "") if want_text else ""
if text_out:
print("[PVL FAL Text]:\n" + text_out)
return (text_out, img_tensor)
except Exception as fe:
print(f"[PVL FAL Fallback] FAL call failed without Google key: {fe}")
raise RuntimeError(f"Gemini image generation failed and FAL fallback also failed: {fe}")
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("GEMINI_API_KEY missing and FAL_KEY missing.")
images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log)
images_tensor = _stack_images_same_size(images_list, debug_log)
text_out = "\n".join(texts) if (want_text and texts) else ""
_t1 = time.time()
print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s")
return text_out, images_tensor
raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.")
# Build client & request
# Google SDK path
from google import genai
client = self._make_client(key, endpoint_override)
parts = self._build_parts(prompt, images, input_mime)
cfg = self._build_config(temperature, want_text, aspect_ratio)
if debug_log:
p_preview = (prompt or "")[:180].replace("\n"," ")
img_count = (images.shape[0] if (isinstance(images, torch.Tensor) and images.ndim==4) else (1 if isinstance(images, torch.Tensor) else 0))
print(f"[PVL Debug] prompt chars={len(prompt or '')} preview='{p_preview}...'")
print(f"[PVL Debug] parts: text={1 if (prompt and prompt.strip()) else 0}, images={img_count}")
safe_parts = []
for pr in parts:
if "text" in pr:
safe_parts.append({"text": pr["text"][:120]})
elif "inline_data" in pr:
di = pr["inline_data"]
safe_parts.append({"inline_data": {"mime_type": di.get("mime_type","image/*"), "data": "<bytes>"}})
try:
temp = getattr(cfg,'temperature',None) if hasattr(cfg,'temperature') else cfg.get('temperature')
top_p = getattr(cfg,'top_p',None) if hasattr(cfg,'top_p') else cfg.get('top_p')
top_k = getattr(cfg,'top_k',None) if hasattr(cfg,'top_k') else cfg.get('top_k')
mot = getattr(cfg,'max_output_tokens',None) if hasattr(cfg,'max_output_tokens') else cfg.get('max_output_tokens')
mods = getattr(cfg,'response_modalities',None) if hasattr(cfg,'response_modalities') else cfg.get('response_modalities')
except Exception:
temp=top_p=top_k=mot=mods=None
print("[PVL Debug] config:", {"temperature": temp, "top_p": top_p, "top_k": top_k, "max_output_tokens": mot, "modalities": mods})
print("[PVL Debug] contents:", [{"role":"user","parts": safe_parts}])
# Parallel Google calls
N = max(1, int(num_images))
results = [None] * N
errors = []
def call_i(i: int):
rid = (request_id.strip() + f"-{i}") if request_id.strip() else f"pvl-nb-{int(time.time()*1000)}-{i}"
return self._single_google_call(client, model, parts, cfg, rid, debug_log)
max_workers = min(N, 6)
if N == 1:
try:
results[0] = call_i(0)
except Exception as e:
errors.append(f"google call 0 failed: {e}")
else:
with ThreadPoolExecutor(max_workers=max_workers) as ex:
futmap = {ex.submit(call_i, i): i for i in range(N)}
for fut in as_completed(futmap):
i = futmap[fut]
try:
results[i] = fut.result()
except Exception as e:
errors.append(f"google call {i} failed: {e}")
# Parse all Google results
out_imgs, out_texts = [], []
for idx, tup in enumerate(results):
if tup is None:
errors.append(f"google call {idx} returned no response")
continue
imgs_i, texts_i, resp = tup
if not imgs_i:
def google_call(p: str, idx: int):
parts = self._build_parts(p, images, input_mime)
cfg = types.GenerateContentConfig(
temperature=float(temperature),
top_p=_TOP_P,
top_k=_TOP_K,
max_output_tokens=_MAX_TOKENS,
response_modalities=["Image"],
image_config=types.ImageConfig(aspect_ratio=aspect_ratio),
)
if debug_log:
print(f"[GOOGLE SUBMIT] Call #{idx+1} model={model}")
resp = client.models.generate_content(
model=model,
contents=[{"role": "user", "parts": parts}],
config=cfg,
)
imgs, texts = [], []
for cand in getattr(resp, "candidates", []) or []:
content = getattr(cand, "content", None)
parts_out = getattr(content, "parts", []) if content else []
for prt in parts_out:
blob = _extract_image_bytes_from_part(prt)
if blob:
try:
pil = Image.open(io.BytesIO(blob)).convert("RGB")
imgs.append(pil_to_tensor(pil))
except Exception as e:
if debug_log:
print("[Decode fail]", e)
else:
t = _extract_text_from_part(prt)
if t:
texts.append(t)
return imgs, texts
out_imgs: T.List[torch.Tensor] = []
out_texts: T.List[str] = []
failed_google = 0
with ThreadPoolExecutor(max_workers=min(N, 6)) as ex:
futs = {ex.submit(google_call, p, i): i for i, p in enumerate(prompts)}
for fut in as_completed(futs):
try:
cands = getattr(resp, "candidates", None) or []
fin = getattr(cands[0], "finish_reason", None) if cands else None
except Exception:
fin = None
errors.append(f"google call {idx} returned no images (finish_reason={fin})")
else:
out_imgs.extend(imgs_i)
if texts_i:
out_texts.append("\n".join(texts_i))
# If Google failed and fallback enabled -> try FAL
if errors and use_fal_fallback:
print("[PVL Fallback] Google call failed; attempting FAL.ai fallback...")
for e in errors:
print("[PVL Google Error]", e)
try:
fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip()
img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format)
final_text = "" # Start with empty; then merge Google texts if requested
if bool(capture_text_output):
pieces = []
if out_texts:
pieces.append("\n\n--- Google ---\n\n" + ("\n".join(out_texts)))
if fal_text:
pieces.append("\n\n--- FAL ---\n\n" + fal_text)
final_text = "".join(pieces)
if final_text:
print("[PVL Fallback Note] Combined text:")
print(final_text)
return (final_text, img_tensor)
except Exception as fe:
print("[PVL FAL Error]", fe)
raise RuntimeError("Both Google and FAL failed. See console for details.")
# If Google produced errors and fallback not used -> raise
if errors:
if out_texts:
print("[PVL Google Text]:\n" + ("\n\n---\n\n".join(out_texts)))
raise RuntimeError("Gemini image generation failed: " + " | ".join(errors[:5]))
imgs, texts = fut.result()
if imgs:
out_imgs.extend(imgs)
if texts:
out_texts.extend(texts)
except Exception as e:
failed_google += 1
if debug_log:
print(f"[GOOGLE ERROR] {e}")
if failed_google > 0:
print(f"[PVL WARNING] {failed_google}/{N} Google requests failed")
if not out_imgs:
raise RuntimeError("Gemini image generation failed: no images across all Google calls.")
images_tensor = torch.cat(out_imgs, dim=0) if len(out_imgs) > 1 else out_imgs[0]
final_text = ("\n\n---\n\n".join(out_texts)) if (bool(capture_text_output) and out_texts) else ""
if final_text:
print(f"[PVL Google NanoBanana Output]:\n{final_text}\n")
return (final_text, images_tensor,)
if use_fal_fallback:
# per-prompt FAL fallback (TRUE PARALLEL)
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("Google returned no images and FAL_KEY missing for fallback.")
images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log)
images_tensor = _stack_images_same_size(images_list, debug_log)
text_out = "\n".join(out_texts + texts) if want_text else ""
_t1 = time.time()
print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s")
return text_out, images_tensor
raise RuntimeError("Gemini returned no images")
images_tensor = _stack_images_same_size(out_imgs, debug_log)
text_out = "\n".join(out_texts) if (want_text and out_texts) else ""
if text_out:
print(f"[PVL Google NanoBanana Output]:\n{text_out}\n")
_t1 = time.time()
print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s")
return text_out, images_tensor
NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_API_mandatory_IMG": PVL_Google_NanoBanana_API_mandatory_IMG}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_API_mandatory_IMG": NODE_NAME}
+194 -199
View File
@@ -1,16 +1,20 @@
# pvl_google_nano_banana_multi_img.py
# PVL Google Nano-Banana Multi API — with regex delimiter support
# Author: PVL
# License: MIT
#
# Features:
# - Text box for prompt input ("STRING").
# - Regex-based delimiter input (e.g., \n|\| or ;+).
# - num_images controls number of parallel API calls.
# - If prompts < num_images, reuses the *last* prompt to fill the rest, with a warning.
# - Parallel calls for Google; optional FAL-only mode and Google→FAL fallback.
# - TRUE PARALLEL FAL execution: submit all requests first, then poll for results.
# - Optionally capture text output and print to console.
# - Ensures IMAGE output is a 4D tensor (B, H, W, C).
# - sync_mode toggle for FAL API (default: False)
import os, io, json, base64, typing as T, time, re
from concurrent.futures import ThreadPoolExecutor, as_completed
@@ -23,30 +27,24 @@ import torch
NODE_NAME = "PVL Google Nano-Banana Multi API"
NODE_CATEGORY = "PVL/Google"
DEFAULT_MODEL = "gemini-2.5-flash-image-preview"
_TOP_P = 0.95
_TOP_K = 64
_MAX_TOKENS = 4096
# --------------------------- Image helpers ---------------------------
def pil_to_tensor(img: Image.Image) -> torch.Tensor:
if img.mode != "RGB":
img = img.convert("RGB")
arr = np.asarray(img, dtype=np.float32) / 255.0
# Shape (H,W,C) -> (1,H,W,C) as ComfyUI image tensor
return torch.from_numpy(arr)[None, ...]
def tensor_to_pil(t: torch.Tensor) -> Image.Image:
# Accept (B,H,W,C) or (H,W,C)
if t.ndim == 4:
t = t[0]
arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8")
return Image.fromarray(arr, "RGB")
def encode_pil_bytes(img: Image.Image, mime: str) -> bytes:
buf = io.BytesIO()
if mime == "image/jpeg":
@@ -55,118 +53,49 @@ def encode_pil_bytes(img: Image.Image, mime: str) -> bytes:
img.save(buf, format="PNG")
return buf.getvalue()
def _extract_image_bytes_from_part(part) -> T.Optional[bytes]:
"""
Enhanced to handle more image data formats from Gemini API response
"""
# Debug: Print the part structure
print(f"[EXTRACT DEBUG] Part type: {type(part)}")
"""Extract image bytes from Gemini API response part."""
try:
# Try to get inline_data in various formats
inline = None
if hasattr(part, 'inline_data'):
inline = getattr(part, 'inline_data')
print(f"[EXTRACT DEBUG] Found inline_data attribute")
elif hasattr(part, 'inlineData'):
inline = getattr(part, 'inlineData')
print(f"[EXTRACT DEBUG] Found inlineData attribute")
elif isinstance(part, dict) and 'inline_data' in part:
inline = part['inline_data']
print(f"[EXTRACT DEBUG] Found inline_data dict key")
elif isinstance(part, dict) and 'inlineData' in part:
inline = part['inlineData']
print(f"[EXTRACT DEBUG] Found inlineData dict key")
inline = getattr(part, "inline_data", None) or getattr(part, "inlineData", None)
if inline is not None:
# Try to get data from inline
data = None
if hasattr(inline, 'data'):
data = getattr(inline, 'data')
print(f"[EXTRACT DEBUG] Found data attribute, type: {type(data)}")
elif isinstance(inline, dict) and 'data' in inline:
data = inline['data']
print(f"[EXTRACT DEBUG] Found data dict key, type: {type(data)}")
if data is not None:
if isinstance(data, str):
try:
result = base64.b64decode(data, validate=False)
print(f"[EXTRACT DEBUG] Successfully decoded base64 string of length {len(data)}")
return result
except Exception as e:
print(f"[EXTRACT DEBUG] Failed to decode base64: {e}")
elif isinstance(data, bytes):
print(f"[EXTRACT DEBUG] Found raw bytes of length {len(data)}")
return data
except Exception as e:
print(f"[EXTRACT DEBUG] Exception during extraction: {e}")
# Try alternative approaches
try:
# Check if the part itself has image data
if hasattr(part, 'image_bytes'):
image_bytes = getattr(part, 'image_bytes')
print(f"[EXTRACT DEBUG] Found image_bytes attribute")
if isinstance(image_bytes, bytes):
return image_bytes
data = getattr(inline, "data", None)
if isinstance(data, (bytes, bytearray)):
return bytes(data)
if isinstance(data, str):
try:
return base64.b64decode(data, validate=False)
except Exception:
return None
except Exception:
pass
if isinstance(part, dict):
blob = None
if "inline_data" in part and isinstance(part["inline_data"], dict):
blob = part["inline_data"].get("data")
elif "inlineData" in part and isinstance(part["inlineData"], dict):
blob = part["inlineData"].get("data")
# Check if the part is a dict with image data
if isinstance(part, dict):
for key in ['image_bytes', 'image_data', 'binary_data']:
if key in part:
data = part[key]
print(f"[EXTRACT DEBUG] Found {key} dict key")
if isinstance(data, str):
try:
return base64.b64decode(data, validate=False)
except Exception:
pass
elif isinstance(data, bytes):
return data
except Exception as e:
print(f"[EXTRACT DEBUG] Exception during alternative extraction: {e}")
if isinstance(blob, (bytes, bytearray)):
return bytes(blob)
if isinstance(blob, str):
try:
return base64.b64decode(blob, validate=False)
except Exception:
return None
# Last resort: try to find any base64 string in the part
try:
if isinstance(part, dict):
for key, value in part.items():
if isinstance(value, str) and len(value) > 100: # Likely a base64 encoded image
try:
# Check if it looks like base64
if value.startswith('data:image/') or re.match(r'^[A-Za-z0-9+/]+={0,2}$', value):
# If it's a data URL, extract the base64 part
if value.startswith('data:image/'):
base64_data = value.split(',', 1)[1]
else:
base64_data = value
result = base64.b64decode(base64_data, validate=False)
print(f"[EXTRACT DEBUG] Found and decoded base64 in dict key '{key}'")
return result
except Exception:
pass
except Exception as e:
print(f"[EXTRACT DEBUG] Exception during base64 search: {e}")
print("[EXTRACT DEBUG] No image data found")
return None
def _extract_text_from_part(part) -> T.Optional[str]:
# Prefer dict, fallback to attribute getattr
if isinstance(part, dict):
if "text" in part and part["text"] is not None:
return str(part["text"])
txt = getattr(part, "text", None)
return str(txt) if txt is not None else None
def _data_url(mime: str, raw: bytes) -> str:
return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8")
def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor:
"""
Concatenate (B,H,W,C) batches along B. If shapes mismatch, resize to the first image size.
@@ -186,16 +115,16 @@ def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False)
fixed.append(pil_to_tensor(rp))
return torch.cat(fixed, dim=0)
# --------------------------- Main Node ---------------------------
class PVL_Google_NanoBanana_Multi_API:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "", "placeholder": "Enter prompts separated by regex delimiter"}),
"delimiter": ("STRING", {"default": "\\n-----\\n", "multiline": False, "placeholder": "Regex (e.g. \\n|\\| or ;+)"}),
"delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Regex (e.g. \\n|\\| or ;+)"}),
},
"optional": {
"image_1": ("IMAGE",),
@@ -210,45 +139,50 @@ class PVL_Google_NanoBanana_Multi_API:
"model": ("STRING", {"default": DEFAULT_MODEL}),
"endpoint_override": ("STRING", {"default": ""}),
"api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use GEMINI_API_KEY"}),
"placeholder": "Leave empty to use GEMINI_API_KEY"}),
"temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}),
"output_format": (["png", "jpeg"], {"default": "png"}),
"capture_text_output": ("BOOLEAN", {"default": False}),
"num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}),
"timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}),
"debug_log": ("BOOLEAN", {"default": False}),
"use_fal_fallback": ("BOOLEAN", {"default": False}),
"use_fal_fallback": ("BOOLEAN", {"default": True}),
"force_fal": ("BOOLEAN", {"default": False}),
"sync_mode": ("BOOLEAN", {"default": False}),
"fal_api_key": ("STRING", {"default": "", "multiline": False,
"placeholder": "Leave empty to use FAL_KEY"}),
"placeholder": "Leave empty to use FAL_KEY"}),
"fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}),
}
}
RETURN_TYPES = ("STRING", "IMAGE",)
RETURN_NAMES = ("text", "images")
FUNCTION = "run"
CATEGORY = NODE_CATEGORY
# ------------------------- INTERNAL HELPERS -------------------------
def _make_client(self, api_key: str, endpoint_override: str):
from google import genai
from google.genai import types
http_options = None
if endpoint_override.strip():
try:
http_options = types.HttpOptions(base_url=endpoint_override.strip())
except Exception:
http_options = None
if http_options is not None:
return genai.Client(api_key=api_key, http_options=http_options)
return genai.Client(api_key=api_key)
def _build_parts(self, prompt: str, image_tensors: T.List[torch.Tensor], mime: str):
parts: T.List[dict] = []
if prompt and prompt.strip():
parts.append({"text": prompt})
for img_tensor in image_tensors:
if img_tensor.ndim == 4:
for i in range(img_tensor.shape[0]):
@@ -258,41 +192,49 @@ class PVL_Google_NanoBanana_Multi_API:
else:
pil = tensor_to_pil(img_tensor)
parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}})
return parts
def _build_call_prompts(self, base_prompts: T.List[str], num_images: int, debug: bool) -> T.List[str]:
"""
Maps prompts to calls according to the agreed rule:
- If len(prompts) >= num_images → take first num_images
- If len(prompts) < num_images → repeat the *last* prompt to fill
- If len(prompts) >= num_images → take first num_images
- If len(prompts) < num_images → repeat the *last* prompt to fill
"""
N = max(1, int(num_images))
if not base_prompts:
return []
if len(base_prompts) >= N:
call_prompts = base_prompts[:N]
else:
if debug:
print(f"[PVL NODE] Provided {len(base_prompts)} prompts but num_images={N}. "
f"Reusing the last prompt for remaining calls.")
print(f"[PVL WARNING] prompt list shorter than num_images: {len(base_prompts)} < {N}. "
f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.")
print(f"[PVL WARNING] prompt list shorter than num_images: {len(base_prompts)} < {N}. "
f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.")
call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts))
if debug:
for i, cp in enumerate(call_prompts, 1):
show = cp if len(cp) <= 160 else (cp[:157] + "...")
print(f"[PVL NODE] Call #{i} prompt: {show}")
return call_prompts
def _fal_queue_call(self, route: str, prompt: str, image_tensors: T.List[torch.Tensor],
mime: str, fal_key: str, timeout: int, debug: bool,
output_format: str, aspect_ratio: str = "1:1"):
# -------- FAL API - TWO PHASE EXECUTION --------
def _fal_submit_only(self, route: str, prompt: str, image_tensors: T.List[torch.Tensor],
mime: str, fal_key: str, timeout: int, debug: bool,
output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False):
"""
Calls FAL queue endpoint in sync_mode, returns (image_tensor_batched1, description_text).
Phase 1: Submit request to FAL queue and return request info immediately.
Does NOT poll for completion.
"""
if not fal_key:
raise RuntimeError("FAL requested but FAL_KEY is missing.")
# Convert any input images to data URLs
data_urls = []
for t in image_tensors:
@@ -306,59 +248,89 @@ class PVL_Google_NanoBanana_Multi_API:
pil = tensor_to_pil(t)
raw = encode_pil_bytes(pil, mime)
data_urls.append(_data_url(mime, raw))
base = "https://queue.fal.run"
submit_url = f"{base}/{route.strip()}"
headers = {"Authorization": f"Key {fal_key}"}
payload = {
"prompt": prompt or "",
"image_urls": data_urls,
"num_images": 1,
"output_format": ("png" if str(output_format).lower() == "png" else "jpeg"),
"aspect_ratio": aspect_ratio,
"sync_mode": True,
"sync_mode": sync_mode,
}
if debug:
print("[FAL SUBMIT]", json.dumps(payload)[:1000])
print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}")
r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout)
if not r.ok:
raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}")
sub = r.json()
req_id = sub.get("request_id")
status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None)
resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None)
if not req_id or not status_url or not resp_url:
raise RuntimeError("FAL queue missing request_id/status/response")
if not req_id:
raise RuntimeError("FAL did not return a request_id")
status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status"
resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}"
return {
"request_id": req_id,
"status_url": status_url,
"response_url": resp_url,
"prompt": prompt
}
def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool):
"""
Phase 2: Poll a single FAL request until complete and fetch the result.
Returns (image_tensor, description_text).
"""
headers = {"Authorization": f"Key {fal_key}"}
status_url = request_info["status_url"]
resp_url = request_info["response_url"]
req_id = request_info["request_id"]
if debug:
print(f"[FAL POLL] request_id={req_id[:16]}...")
# Poll for completion
deadline = time.time() + timeout
while time.time() < deadline:
sr = requests.get(status_url, headers=headers, timeout=10)
if sr.ok and sr.json().get("status") == "COMPLETED":
break
time.sleep(0.6)
# Fetch result
rr = requests.get(resp_url, headers=headers, timeout=15)
if not rr.ok:
raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}")
rdata = rr.json()
if debug:
print("[FAL RAW]", json.dumps(rdata)[:2000])
print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED")
# Extract response data
buckets = []
resp = rdata.get("response") if isinstance(rdata, dict) else None
if resp is None and isinstance(rdata, dict):
resp = rdata
if isinstance(resp, dict):
for key in ("images", "outputs", "artifacts"):
val = resp.get(key)
if isinstance(val, list):
buckets.extend(val)
for key in ("image", "output", "result"):
val = resp.get(key)
if isinstance(val, (str, dict)):
buckets.append(val)
out = []
for item in buckets:
try:
@@ -366,6 +338,7 @@ class PVL_Google_NanoBanana_Multi_API:
else (item.get("url") or item.get("data") or item.get("image"))
if not url:
continue
if url.startswith("data:image/"):
blob = base64.b64decode(url.split(",", 1)[1])
else:
@@ -373,19 +346,21 @@ class PVL_Google_NanoBanana_Multi_API:
if not ir.ok:
continue
blob = ir.content
pil = Image.open(io.BytesIO(blob)).convert("RGB")
out.append(pil_to_tensor(pil))
except Exception as ex:
if debug:
print("[FAL decode fail]", ex)
print(f"[FAL decode fail] {ex}")
if not out:
raise RuntimeError("FAL returned no images")
return out[0], (resp.get("description", "") if isinstance(resp, dict) else "")
raise RuntimeError(f"FAL returned no images for request_id={req_id}")
description = resp.get("description", "") if isinstance(resp, dict) else ""
return out[0], description
# --------------------------- RUN MAIN -----------------------------
def run(self, prompt: str, delimiter: str,
image_1=None, image_2=None, image_3=None, image_4=None,
image_5=None, image_6=None, image_7=None, image_8=None,
@@ -395,7 +370,9 @@ class PVL_Google_NanoBanana_Multi_API:
capture_text_output: bool = False, num_images: int = 1,
timeout_sec: int = 120, debug_log: bool = False,
use_fal_fallback: bool = False, force_fal: bool = False,
sync_mode: bool = False,
fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"):
# --- Validate aspect ratio (for Google & FAL) ---
ar_original = aspect_ratio.strip()
valid_ratios = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"}
@@ -404,163 +381,181 @@ class PVL_Google_NanoBanana_Multi_API:
aspect_ratio = "1:1"
else:
aspect_ratio = ar_original
# --- Regex-based prompt splitting (with safe fallback to literal split) ---
try:
base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()]
except re.error:
print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.")
base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()]
if not base_prompts:
raise RuntimeError("No valid prompts provided.")
# Collect any provided images (tensors) in order
image_tensors: T.List[torch.Tensor] = []
for img in [image_1, image_2, image_3, image_4, image_5, image_6, image_7, image_8]:
if img is not None and torch.is_tensor(img):
image_tensors.append(img)
# Map provided prompts to num_images calls
call_prompts = self._build_call_prompts(base_prompts, num_images, debug_log)
input_mime = "image/png" if output_format.lower() == "png" else "image/jpeg"
want_text = bool(capture_text_output)
# ---------------- FAL-only path ----------------
# ---------------- FAL-only path (TRUE PARALLEL) ----------------
if force_fal:
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("force_fal=True but no FAL_KEY provided.")
if debug_log:
print(f"[FAL] Submitting {len(call_prompts)} requests...")
# PHASE 1: Submit all requests (fast, non-blocking)
request_infos = []
for p in call_prompts:
req_info = self._fal_submit_only(fal_route, p, image_tensors,
input_mime, fal_key, timeout_sec,
debug_log, output_format,
aspect_ratio, sync_mode)
request_infos.append(req_info)
if debug_log:
print(f"[FAL] All {len(request_infos)} requests submitted. Polling for results...")
# PHASE 2: Poll all requests in parallel
results, texts = [], []
with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex:
with ThreadPoolExecutor(max_workers=min(len(request_infos), 6)) as ex:
futs = {
ex.submit(self._fal_queue_call, fal_route, p, image_tensors,
input_mime, fal_key, timeout_sec, debug_log, output_format, aspect_ratio): p
for p in call_prompts
ex.submit(self._fal_poll_and_fetch, req_info, fal_key,
timeout_sec, debug_log): req_info
for req_info in request_infos
}
for fut in as_completed(futs):
img, t = fut.result()
results.append(img) # Fixed: removed .unsqueeze(0)
results.append(img)
if t:
texts.append(t)
images_tensor = _stack_images_same_size(results, debug_log)
text_out = "\n".join(texts) if want_text else ""
return text_out, images_tensor
# ---------------- Google path (with optional FAL fallback) ----------------
key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip()
if not key:
raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.")
client = self._make_client(key, endpoint_override)
def google_call(p: str, debug: bool):
# Prepend explicit image generation instruction
image_prompt = f"{p}"
if debug:
print(f"[GEMINI PROMPT] {image_prompt}")
parts = self._build_parts(p, image_tensors, input_mime)
parts = self._build_parts(image_prompt, image_tensors, input_mime)
# Correct configuration for image generation
cfg = types.GenerateContentConfig(
temperature=float(temperature),
top_p=_TOP_P,
top_k=_TOP_K,
max_output_tokens=_MAX_TOKENS,
response_modalities=["Image"], # Request image generation
# No response_mime_type - model returns PNG by default
response_modalities=["Image"],
image_config=types.ImageConfig(aspect_ratio=aspect_ratio),
)
if debug:
print(f"[GOOGLE SUBMIT] prompt: {p[:100]}...")
resp = client.models.generate_content(
model=model,
contents=[{"role": "user", "parts": parts}],
config=cfg,
)
if debug:
print("[GEMINI RESPONSE STRUCTURE]")
print(f" Candidates: {len(getattr(resp, 'candidates', []))}")
for i, cand in enumerate(getattr(resp, 'candidates', []) or []):
content = getattr(cand, 'content', None)
parts_out = getattr(content, 'parts', []) if content else []
print(f" Candidate {i}: {len(parts_out)} parts")
for j, prt in enumerate(parts_out):
# Enhanced debugging
print(f" Part {j}: Type = {type(prt)}")
if hasattr(prt, '__dict__'):
print(f" Attributes: {list(prt.__dict__.keys())}")
elif isinstance(prt, dict):
print(f" Keys: {list(prt.keys())}")
imgs, texts = [], []
for cand in getattr(resp, "candidates", []) or []:
content = getattr(cand, "content", None)
parts_out = getattr(content, "parts", []) if content else []
for prt in parts_out:
blob = _extract_image_bytes_from_part(prt)
if blob:
if debug:
print(f" Extracted image blob of size {len(blob)} bytes")
try:
pil = Image.open(io.BytesIO(blob)).convert("RGB")
imgs.append(pil_to_tensor(pil))
except Exception as e:
print(f" Failed to convert blob to image: {e}")
if debug:
print(f"[Decode fail] {e}")
else:
t = _extract_text_from_part(prt)
if t:
if debug:
print(f" Extracted text: {t[:100]}...")
texts.append(t)
if debug:
print(f"[GEMINI RESULT] Found {len(imgs)} images and {len(texts)} text parts")
print(f"[GOOGLE RESULT] Found {len(imgs)} images")
return imgs, texts
out_imgs: T.List[torch.Tensor] = []
out_texts: T.List[str] = []
with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex:
futs = {ex.submit(google_call, p, debug_log): p for p in call_prompts}
for fut in as_completed(futs):
imgs, texts = fut.result()
if imgs:
# Fixed: removed .unsqueeze(0) since pil_to_tensor already adds batch dimension
out_imgs.append(imgs[0])
if texts:
out_texts.extend(texts)
if not out_imgs:
if use_fal_fallback:
fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip()
if not fal_key:
raise RuntimeError("Google returned no images and FAL fallback has no FAL_KEY.")
if debug_log:
print(f"[FAL FALLBACK] Submitting {len(call_prompts)} requests...")
# PHASE 1: Submit all requests
request_infos = []
for p in call_prompts:
req_info = self._fal_submit_only(fal_route, p, image_tensors,
input_mime, fal_key, timeout_sec,
debug_log, output_format,
aspect_ratio, sync_mode)
request_infos.append(req_info)
if debug_log:
print(f"[FAL FALLBACK] All requests submitted. Polling for results...")
# PHASE 2: Poll all requests in parallel
results, texts = [], []
with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex:
with ThreadPoolExecutor(max_workers=min(len(request_infos), 6)) as ex:
futs = {
ex.submit(self._fal_queue_call, fal_route, p, image_tensors,
input_mime, fal_key, timeout_sec, debug_log, output_format, aspect_ratio): p
for p in call_prompts
ex.submit(self._fal_poll_and_fetch, req_info, fal_key,
timeout_sec, debug_log): req_info
for req_info in request_infos
}
for fut in as_completed(futs):
img, t = fut.result()
results.append(img) # Fixed: removed .unsqueeze(0)
results.append(img)
if t:
texts.append(t)
images_tensor = _stack_images_same_size(results, debug_log)
combo_texts = out_texts + texts
text_out = "\n".join(combo_texts) if want_text else ""
return text_out, images_tensor
raise RuntimeError("Gemini returned no images")
images_tensor = _stack_images_same_size(out_imgs, debug_log)
text_out = "\n".join(out_texts) if (want_text and out_texts) else ""
if text_out:
if text_out and debug_log:
print(f"[PVL Google NanoBanana Output]:\n{text_out}\n")
return text_out, images_tensor
NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": PVL_Google_NanoBanana_Multi_API}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": NODE_NAME}
NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": NODE_NAME}