Add files via upload
This commit is contained in:
@@ -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
@@ -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"}
|
||||
|
||||
@@ -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
@@ -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
@@ -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}
|
||||
|
||||
@@ -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
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user